Skip to main content

tmk_vmm/
run.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4//! Support for running a VM's VPs.
5
6use crate::Options;
7use crate::load;
8use anyhow::Context as _;
9use futures::StreamExt as _;
10use guestmem::GuestMemory;
11use hvdef::Vtl;
12use pal_async::DefaultDriver;
13use std::sync::Arc;
14#[cfg(target_os = "linux")]
15use user_driver::DmaClient;
16use virt::PartitionCapabilities;
17use virt::Processor;
18use virt::StopVpSource;
19use virt::VpIndex;
20use virt::io::CpuIo;
21use virt::vp::AccessVpState as _;
22use vm_topology::memory::MemoryLayout;
23use vm_topology::processor::ProcessorTopology;
24use vm_topology::processor::TopologyBuilder;
25use vmcore::vmtime::VmTime;
26use vmcore::vmtime::VmTimeKeeper;
27use vmcore::vmtime::VmTimeSource;
28use zerocopy::TryFromBytes as _;
29
30pub const COMMAND_ADDRESS: u64 = 0xffff_0000;
31
32#[cfg(all(target_os = "linux", guest_arch = "aarch64"))]
33mod cca {
34    use super::DmaClient;
35    use super::MemoryLayout;
36    use super::Options;
37    use crate::HypervisorOpt;
38    use anyhow::Context as _;
39    use core::ops::Range;
40    use memory_range::MemoryRange;
41    use std::sync::Arc;
42    use underhill_mem::MemoryAcceptor;
43    use user_driver::lockmem::LockedMemorySpawner;
44    use user_driver::memory::MemoryBlock;
45    use user_driver::memory::PAGE_SIZE;
46    use user_driver::memory::PAGE_SIZE64;
47    use virt::IsolationType;
48    use vm_topology::memory::MemoryRangeWithNode;
49
50    pub(super) struct CcaState {
51        pub(super) private_dma_client: Arc<dyn DmaClient>,
52        pub(super) _guest_ram_backing: MemoryBlock,
53    }
54
55    pub(super) fn build(
56        opts: &Options,
57        memory_layout: &mut MemoryLayout,
58        ram_size: u64,
59    ) -> anyhow::Result<Option<CcaState>> {
60        let hv = opts.hv.expect("hv must have a finalized value");
61        match hv {
62            HypervisorOpt::Cca => {
63                let mut map_size = ram_size as usize;
64                let private_dma_client: Arc<dyn DmaClient> = Arc::new(LockedMemorySpawner);
65
66                let (private_memory, private_ram_pfn) = {
67                    const BITMAP_ALIGNMENT: u64 = PAGE_SIZE64 * 8;
68                    const MAX_ALLOC_ATTEMPTS: usize = 4;
69                    let mut selected = None;
70
71                    for _attempt in 0..MAX_ALLOC_ATTEMPTS {
72                        let private_memory = private_dma_client
73                            .allocate_dma_buffer(map_size)
74                            .with_context(|| {
75                                format!(
76                                    "failed to allocate private CCA RAM buffer of size {map_size}"
77                                )
78                            })?;
79
80                        let asking_size = ram_size
81                            .checked_add(BITMAP_ALIGNMENT - PAGE_SIZE64)
82                            .context("private CCA RAM search size overflowed")?;
83                        if let Some(pfns) =
84                            contiguous_subpfns(&private_memory, asking_size as usize)
85                        {
86                            let page_count = (ram_size as usize).div_ceil(PAGE_SIZE);
87                            if let Some(start_index) = pfns
88                                .iter()
89                                .position(|pfn| {
90                                    (pfn * PAGE_SIZE64).is_multiple_of(BITMAP_ALIGNMENT)
91                                })
92                                .filter(|&start_index| pfns.len() - start_index >= page_count)
93                            {
94                                selected = Some((private_memory, pfns[start_index]));
95                                break;
96                            }
97                        }
98
99                        map_size = map_size
100                            .checked_mul(2)
101                            .context("private CCA RAM allocation size overflowed while retrying")?;
102                    }
103
104                    selected.with_context(|| {
105                        format!(
106                            "failed to allocate private CCA RAM with {ram_size} contiguous bytes after {MAX_ALLOC_ATTEMPTS} attempts"
107                        )
108                    })?
109                };
110
111                private_memory.write_zeros(0, private_memory.len());
112
113                let pa = private_ram_pfn * PAGE_SIZE64;
114                let start = pa;
115                let end = pa
116                    .checked_add(ram_size)
117                    .context("private CCA RAM range overflowed")?;
118
119                *memory_layout = MemoryLayout::new_from_ranges(
120                    &[MemoryRangeWithNode {
121                        range: MemoryRange::new(Range { start, end }),
122                        vnode: 0,
123                    }],
124                    &[],
125                )
126                .context("bad memory layout")?;
127
128                // Grant GPA to Plane1 (eqv. VTL0)
129                let ram = memory_layout.ram().iter().map(|r| r.range);
130                let acceptor = MemoryAcceptor::new(IsolationType::Cca)?;
131                for range in ram {
132                    acceptor.apply_initial_lower_vtl_protections(range)?;
133                }
134
135                Ok(Some(CcaState {
136                    private_dma_client,
137                    _guest_ram_backing: private_memory,
138                }))
139            }
140            _ => Ok(None),
141        }
142    }
143
144    /// Returns a sorted contiguous subset of PFNs large enough for `asking_size` bytes.
145    fn contiguous_subpfns(memory: &MemoryBlock, asking_size: usize) -> Option<Vec<u64>> {
146        let page_count = asking_size.div_ceil(PAGE_SIZE);
147        if page_count == 0 {
148            return Some(Vec::new());
149        }
150
151        let mut pfns = memory.pfns().to_vec();
152        pfns.sort_unstable();
153
154        let mut run_start = 0;
155        for i in 1..=pfns.len() {
156            let run_ended = i == pfns.len() || pfns[i - 1] + 1 != pfns[i];
157            if run_ended {
158                if i - run_start >= page_count {
159                    pfns.truncate(run_start + page_count);
160                    pfns.drain(..run_start);
161                    return Some(pfns);
162                }
163                run_start = i;
164            }
165        }
166
167        None
168    }
169}
170
171pub struct CommonState {
172    pub driver: DefaultDriver,
173    pub opts: Options,
174    pub processor_topology: ProcessorTopology,
175    pub memory_layout: MemoryLayout,
176    #[cfg(all(target_os = "linux", guest_arch = "aarch64"))]
177    cca: Option<cca::CcaState>,
178}
179
180pub struct RunContext<'a> {
181    pub state: &'a CommonState,
182    pub vmtime_source: &'a VmTimeSource,
183}
184
185#[derive(Debug, Clone)]
186pub enum TestResult {
187    Passed,
188    Failed,
189    Faulted {
190        vp_index: VpIndex,
191        reason: String,
192        regs: Option<Box<virt::vp::Registers>>,
193    },
194}
195
196impl CommonState {
197    #[cfg(all(target_os = "linux", guest_arch = "aarch64"))]
198    pub fn cca_private_dma_client(&self) -> Arc<dyn DmaClient> {
199        self.cca
200            .as_ref()
201            .expect("CCA private DMA client is only available when running with --hv cca")
202            .private_dma_client
203            .clone()
204    }
205
206    #[cfg(all(target_os = "linux", not(guest_arch = "aarch64")))]
207    pub fn cca_private_dma_client(&self) -> Arc<dyn DmaClient> {
208        panic!("CCA private DMA client is only available on aarch64")
209    }
210
211    pub async fn new(driver: DefaultDriver, opts: Options) -> anyhow::Result<Self> {
212        #[cfg(guest_arch = "x86_64")]
213        let processor_topology = TopologyBuilder::new_x86()
214            .x2apic(vm_topology::processor::x86::X2ApicState::Supported)
215            .build(1)
216            .context("failed to build processor topology")?;
217
218        #[cfg(guest_arch = "aarch64")]
219        let processor_topology =
220            TopologyBuilder::new_aarch64(vm_topology::processor::arch::Aarch64PlatformConfig {
221                gic_distributor_base: 0xff000000,
222                gic_version: vm_topology::processor::aarch64::GicVersion::V3 {
223                    redistributors_base: 0xff020000,
224                },
225                gic_msi: vm_topology::processor::aarch64::GicMsiController::None,
226                pmu_gsiv: None,
227                virt_timer_ppi: 20, // DEFAULT_VIRT_TIMER_PPI
228                gic_nr_irqs: 256,
229            })
230            .build(1)
231            .context("failed to build processor topology")?;
232
233        let ram_size = 0x400000;
234
235        #[cfg_attr(
236            not(all(target_os = "linux", guest_arch = "aarch64")),
237            expect(unused_mut)
238        )]
239        let mut memory_layout =
240            MemoryLayout::new(ram_size, &[], &[], &[], None).context("bad memory layout")?;
241        #[cfg(all(target_os = "linux", guest_arch = "aarch64"))]
242        let cca = cca::build(&opts, &mut memory_layout, ram_size)?;
243
244        Ok(Self {
245            driver,
246            opts,
247            processor_topology,
248            memory_layout,
249            #[cfg(all(target_os = "linux", guest_arch = "aarch64"))]
250            cca,
251        })
252    }
253
254    pub async fn for_each_test(
255        &mut self,
256        mut f: impl AsyncFnMut(&mut RunContext<'_>, &load::TestInfo) -> anyhow::Result<TestResult>,
257    ) -> anyhow::Result<()> {
258        let tmk = fs_err::File::open(&self.opts.tmk).context("failed to open tmk")?;
259        let available_tests = load::enumerate_tests(&tmk)?;
260        let tests = if self.opts.tests.is_empty() {
261            available_tests
262        } else {
263            self.opts
264                .tests
265                .iter()
266                .map(|name| {
267                    available_tests
268                        .iter()
269                        .find(|test| test.name == *name)
270                        .cloned()
271                        .with_context(|| format!("test {} not found", name))
272                })
273                .collect::<anyhow::Result<Vec<_>>>()?
274        };
275        let mut success = true;
276        for test in &tests {
277            tracing::info!(target: "test", name = test.name, "test started");
278
279            if test.linux_only && !cfg!(target_os = "linux") {
280                tracing::info!(target: "test", name = test.name, "test skipped, incompatible os");
281                continue;
282            }
283
284            let mut vmtime_keeper = VmTimeKeeper::new(&self.driver, VmTime::from_100ns(0));
285            let vmtime_source = vmtime_keeper.builder().build(&self.driver).await.unwrap();
286            let mut ctx = RunContext {
287                state: self,
288                vmtime_source: &vmtime_source,
289            };
290
291            vmtime_keeper.start().await;
292
293            let r = f(&mut ctx, test)
294                .await
295                .with_context(|| format!("failed to run test {}", test.name))?;
296
297            vmtime_keeper.stop().await;
298
299            match (r, test.expected_failure) {
300                (TestResult::Passed, false) => {
301                    tracing::info!(target: "test", name = test.name, "test passed");
302                }
303                (TestResult::Passed, true) => {
304                    tracing::error!(
305                        target: "test",
306                        name = test.name,
307                        expected_failure = true,
308                        "test unexpectedly passed"
309                    );
310                    success = false;
311                }
312                (TestResult::Failed, false) => {
313                    tracing::error!(target: "test", name = test.name, reason = "explicit failure", "test failed");
314                    success = false;
315                }
316                (TestResult::Failed, true) => {
317                    tracing::info!(
318                        target: "test",
319                        name = test.name,
320                        expected_failure = true,
321                        reason = "explicit failure",
322                        "test passed"
323                    );
324                }
325                (
326                    TestResult::Faulted {
327                        vp_index,
328                        reason,
329                        regs,
330                    },
331                    false,
332                ) => {
333                    tracing::error!(
334                        target: "test",
335                        name = test.name,
336                        vp_index = vp_index.index(),
337                        reason,
338                        regs = format_args!("{:#x?}", regs),
339                        "test failed"
340                    );
341                    success = false;
342                }
343                (
344                    TestResult::Faulted {
345                        vp_index,
346                        reason,
347                        regs,
348                    },
349                    true,
350                ) => {
351                    tracing::info!(
352                        target: "test",
353                        name = test.name,
354                        expected_failure = true,
355                        vp_index = vp_index.index(),
356                        reason,
357                        regs = format_args!("{:#x?}", regs),
358                        "test passed"
359                    );
360                }
361            }
362        }
363        if !success {
364            anyhow::bail!("some tests failed");
365        }
366        Ok(())
367    }
368}
369
370impl RunContext<'_> {
371    pub async fn run(
372        &mut self,
373        guest_memory: &GuestMemory,
374        caps: &PartitionCapabilities,
375        test: &load::TestInfo,
376        start_vp: impl AsyncFnOnce(&mut Self, RunnerBuilder) -> anyhow::Result<()>,
377    ) -> anyhow::Result<TestResult> {
378        let (event_send, mut event_recv) = mesh::channel();
379
380        // Load the TMK.
381        let tmk = fs_err::File::open(&self.state.opts.tmk).context("failed to open tmk")?;
382        let regs = {
383            #[cfg(guest_arch = "x86_64")]
384            {
385                load::load_x86(
386                    &self.state.memory_layout,
387                    guest_memory,
388                    &self.state.processor_topology,
389                    caps,
390                    &tmk,
391                    test,
392                )?
393            }
394            #[cfg(guest_arch = "aarch64")]
395            {
396                load::load_aarch64(
397                    &self.state.memory_layout,
398                    guest_memory,
399                    &self.state.processor_topology,
400                    caps,
401                    &tmk,
402                    test,
403                )?
404            }
405        };
406
407        start_vp(
408            self,
409            RunnerBuilder::new(
410                VpIndex::BSP,
411                Arc::clone(&regs),
412                guest_memory.clone(),
413                event_send.clone(),
414            ),
415        )
416        .await?;
417
418        let event = event_recv.next().await.unwrap();
419        let r = match event {
420            VpEvent::TestComplete { success } => {
421                if success {
422                    TestResult::Passed
423                } else {
424                    TestResult::Failed
425                }
426            }
427            VpEvent::Halt {
428                vp_index,
429                reason,
430                regs,
431            } => TestResult::Faulted {
432                vp_index,
433                reason,
434                regs,
435            },
436        };
437
438        Ok(r)
439    }
440}
441
442enum VpEvent {
443    TestComplete {
444        success: bool,
445    },
446    Halt {
447        vp_index: VpIndex,
448        reason: String,
449        regs: Option<Box<virt::vp::Registers>>,
450    },
451}
452
453struct IoHandler<'a> {
454    guest_memory: &'a GuestMemory,
455    event_send: &'a mesh::Sender<VpEvent>,
456    stop: &'a StopVpSource,
457}
458
459fn widen(d: &[u8]) -> u64 {
460    let mut v = [0; 8];
461    v[..d.len()].copy_from_slice(d);
462    u64::from_ne_bytes(v)
463}
464
465impl CpuIo for IoHandler<'_> {
466    fn is_mmio(&self, _address: u64) -> bool {
467        false
468    }
469
470    fn acknowledge_pic_interrupt(&self) -> Option<u8> {
471        None
472    }
473
474    fn handle_eoi(&self, irq: u32) {
475        tracing::info!(irq, "eoi");
476    }
477
478    async fn read_mmio(&self, vp: VpIndex, address: u64, data: &mut [u8]) {
479        tracing::info!(vp = vp.index(), address, "read mmio");
480        data.fill(!0);
481    }
482
483    async fn write_mmio(&self, vp: VpIndex, address: u64, data: &[u8]) {
484        if address == COMMAND_ADDRESS {
485            let p = widen(data);
486            let r = self.handle_command(p);
487            if let Err(e) = r {
488                tracing::error!(
489                    error = e.as_ref() as &dyn std::error::Error,
490                    p,
491                    "failed to handle command"
492                );
493            }
494        } else {
495            tracing::info!(vp = vp.index(), address, data = widen(data), "write mmio");
496        }
497    }
498
499    async fn read_io(&self, vp: VpIndex, port: u16, data: &mut [u8]) {
500        tracing::info!(vp = vp.index(), port, "read io");
501        data.fill(!0);
502    }
503
504    async fn write_io(&self, vp: VpIndex, port: u16, data: &[u8]) {
505        tracing::info!(vp = vp.index(), port, data = widen(data), "write io");
506    }
507
508    #[track_caller]
509    fn fatal_error(&self, error: Box<dyn std::error::Error + Send + Sync>) -> virt::VpHaltReason {
510        tracing::error!(
511            err = error.as_ref() as &dyn std::error::Error,
512            "fatal error"
513        );
514        virt::VpHaltReason::TripleFault { vtl: Vtl::Vtl0 }
515    }
516}
517
518impl IoHandler<'_> {
519    fn read_str(&self, s: tmk_protocol::StrDescriptor) -> anyhow::Result<String> {
520        let mut buf = vec![0; s.len as usize];
521        self.guest_memory
522            .read_at(s.gpa, &mut buf)
523            .context("failed to read string")?;
524        String::from_utf8(buf).context("string not utf-8")
525    }
526
527    fn handle_command(&self, gpa: u64) -> anyhow::Result<()> {
528        let buf = self
529            .guest_memory
530            .read_plain::<[u8; size_of::<tmk_protocol::Command>()]>(gpa)
531            .context("failed to read command")?;
532        let cmd = tmk_protocol::Command::try_read_from_bytes(&buf)
533            .ok()
534            .context("bad command")?;
535        match cmd {
536            tmk_protocol::Command::Log(s) => {
537                let message = self.read_str(s)?;
538                tracing::info!(target: "tmk", message);
539            }
540            tmk_protocol::Command::Panic {
541                message,
542                filename,
543                line,
544            } => {
545                let message = self.read_str(message)?;
546                let location = if filename.len > 0 {
547                    Some(format!("{}:{}", self.read_str(filename)?, line))
548                } else {
549                    None
550                };
551                tracing::error!(target: "tmk", location, panic = message);
552                self.event_send
553                    .send(VpEvent::TestComplete { success: false });
554                self.stop.stop();
555            }
556            tmk_protocol::Command::Complete { success } => {
557                self.event_send.send(VpEvent::TestComplete { success });
558                self.stop.stop();
559            }
560        }
561        Ok(())
562    }
563}
564
565pub struct RunnerBuilder {
566    vp_index: VpIndex,
567    regs: Arc<virt::InitialRegs>,
568    guest_memory: GuestMemory,
569    event_send: mesh::Sender<VpEvent>,
570}
571
572impl RunnerBuilder {
573    fn new(
574        vp_index: VpIndex,
575        regs: Arc<virt::InitialRegs>,
576        guest_memory: GuestMemory,
577        event_send: mesh::Sender<VpEvent>,
578    ) -> Self {
579        Self {
580            vp_index,
581            regs,
582            guest_memory,
583            event_send,
584        }
585    }
586
587    pub fn build<P: Processor>(&mut self, mut vp: P) -> anyhow::Result<Runner<'_, P>> {
588        {
589            let mut state = vp.access_state(Vtl::Vtl0);
590            #[cfg(guest_arch = "x86_64")]
591            {
592                let virt::x86::X86InitialRegs {
593                    registers,
594                    mtrrs,
595                    pat,
596                } = self.regs.as_ref();
597                state.set_registers(registers)?;
598                state.set_mtrrs(mtrrs)?;
599                state.set_pat(pat)?;
600            }
601            #[cfg(guest_arch = "aarch64")]
602            {
603                let virt::aarch64::Aarch64InitialRegs {
604                    registers,
605                    system_registers,
606                } = self.regs.as_ref();
607                state.set_registers(registers)?;
608                state.set_system_registers(system_registers)?;
609            }
610            state.commit()?;
611        }
612        Ok(Runner {
613            vp,
614            vp_index: self.vp_index,
615            guest_memory: &self.guest_memory,
616            event_send: &self.event_send,
617        })
618    }
619}
620
621pub struct Runner<'a, P> {
622    vp: P,
623    vp_index: VpIndex,
624    guest_memory: &'a GuestMemory,
625    event_send: &'a mesh::Sender<VpEvent>,
626}
627
628impl<P: Processor> Runner<'_, P> {
629    pub async fn run_vp(&mut self) {
630        let stop = StopVpSource::new();
631        let Err(err) = self
632            .vp
633            .run_vp(
634                stop.checker(),
635                &IoHandler {
636                    guest_memory: self.guest_memory,
637                    event_send: self.event_send,
638                    stop: &stop,
639                },
640            )
641            .await;
642        let regs = self
643            .vp
644            .access_state(Vtl::Vtl0)
645            .registers()
646            .map(Box::new)
647            .ok();
648        self.event_send.send(VpEvent::Halt {
649            vp_index: self.vp_index,
650            reason: format!("{:?}", err),
651            regs,
652        });
653    }
654}