Skip to main content

user_driver_emulated_mock/
lib.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4//! This crate provides a collection of wrapper structs around things like devices and memory. Through the wrappers, it provides functionality to emulate devices such
5//! as Nvme and Mana and gives some additional control over things like [`GuestMemory`] to make testing devices easier.
6//! Everything in this crate is meant for TESTING PURPOSES ONLY and it should only ever be added as a dev-dependency (Few expceptions like using this for fuzzing)
7
8mod guest_memory_access_wrapper;
9
10use crate::guest_memory_access_wrapper::GuestMemoryAccessWrapper;
11
12use anyhow::Context;
13use chipset_device::mmio::MmioIntercept;
14use chipset_device::pci::ByteEnabledDwordRead;
15use chipset_device::pci::ByteEnabledDwordWrite;
16use chipset_device::pci::PciConfigSpace;
17use guestmem::GuestMemory;
18use inspect::Inspect;
19use inspect::InspectMut;
20use memory_range::MemoryRange;
21use page_pool_alloc::PagePool;
22use page_pool_alloc::PagePoolAllocator;
23use page_pool_alloc::TestMapper;
24use parking_lot::Mutex;
25use pci_core::chipset_device_ext::PciChipsetDeviceExt;
26use pci_core::msi::MsiConnection;
27use pci_core::msi::SignalMsi;
28use std::sync::Arc;
29use user_driver::DeviceBacking;
30use user_driver::DeviceRegisterIo;
31use user_driver::DmaClient;
32use user_driver::interrupt::DeviceInterrupt;
33use user_driver::interrupt::DeviceInterruptSource;
34use user_driver::memory::PAGE_SIZE64;
35
36/// A wrapper around any user_driver device T. It provides device emulation by providing access to the memory shared with the device and thus
37/// allowing the user to control device behaviour to a certain extent. Can be used with devices such as the `NvmeController`
38pub struct EmulatedDevice<T, U> {
39    device: Arc<Mutex<T>>,
40    controller: Arc<MsiController>,
41    dma_client: Arc<U>,
42    bar0_len: usize,
43}
44
45impl<T: InspectMut, U> Inspect for EmulatedDevice<T, U> {
46    fn inspect(&self, req: inspect::Request<'_>) {
47        self.device.lock().inspect_mut(req);
48    }
49}
50
51struct MsiController {
52    events: Arc<[DeviceInterruptSource]>,
53}
54
55impl MsiController {
56    fn new(n: usize) -> Self {
57        Self {
58            events: (0..n).map(|_| DeviceInterruptSource::new()).collect(),
59        }
60    }
61}
62
63impl SignalMsi for MsiController {
64    fn signal_msi(&self, _devid: Option<u32>, address: u64, _data: u32) {
65        let index = address as usize;
66        if let Some(event) = self.events.get(index) {
67            tracing::debug!(index, "signaling interrupt");
68            event.signal_uncached();
69        } else {
70            tracing::info!("interrupt ignored");
71        }
72    }
73}
74
75impl<T: PciConfigSpace + MmioIntercept, U: DmaClient> Clone for EmulatedDevice<T, U> {
76    fn clone(&self) -> Self {
77        Self {
78            device: self.device.clone(),
79            controller: self.controller.clone(),
80            dma_client: self.dma_client.clone(),
81            bar0_len: self.bar0_len,
82        }
83    }
84}
85
86impl<T: PciConfigSpace + MmioIntercept, U: DmaClient> EmulatedDevice<T, U> {
87    /// Creates a new emulated device, wrapping `device` of type T, using the provided MSI Interrupt Set. Dma_client should point to memory
88    /// shared with the device.
89    pub fn new(mut device: T, msi_conn: MsiConnection, dma_client: Arc<U>) -> Self {
90        let bars = device.probe_bar_masks();
91        let bar0_len = !(bars[0] & !0xf) as usize + 1;
92
93        // Enable BAR0 at 0, BAR4 at X.
94        device
95            .pci_cfg_write(0x20, ByteEnabledDwordWrite::with_all_bytes_enabled(0))
96            .unwrap();
97        device
98            .pci_cfg_write(0x24, ByteEnabledDwordWrite::with_all_bytes_enabled(0x1))
99            .unwrap();
100        device
101            .pci_cfg_write(
102                0x4,
103                ByteEnabledDwordWrite::with_all_bytes_enabled(
104                    pci_core::spec::cfg_space::Command::new()
105                        .with_mmio_enabled(true)
106                        .into_bits() as u32,
107                ),
108            )
109            .unwrap();
110
111        // Determine the number of MSI-X vectors.
112        let msix_table_size = {
113            let mut n = 0;
114            device
115                .pci_cfg_read(0x40, ByteEnabledDwordRead::with_all_bytes_enabled(&mut n))
116                .unwrap();
117            ((n >> 16) & 0x7ff) + 1
118        } as usize;
119
120        // Connect an interrupt controller.
121        let controller = Arc::new(MsiController::new(msix_table_size));
122        msi_conn.connect(controller.clone());
123
124        // Enable MSIX.
125        for i in 0u64..64 {
126            device
127                .mmio_write((0x1 << 32) + i * 16, &i.to_ne_bytes())
128                .unwrap();
129            device
130                .mmio_write((0x1 << 32) + i * 16 + 12, &0u32.to_ne_bytes())
131                .unwrap();
132        }
133        device
134            .pci_cfg_write(
135                0x40,
136                ByteEnabledDwordWrite::with_all_bytes_enabled(0x80000000),
137            )
138            .unwrap();
139
140        Self {
141            device: Arc::new(Mutex::new(device)),
142            controller,
143            dma_client,
144            bar0_len,
145        }
146    }
147}
148
149/// A memory mapping for an [`EmulatedDevice`].
150#[derive(Inspect)]
151pub struct Mapping<T> {
152    #[inspect(skip)]
153    device: Arc<Mutex<T>>,
154    addr: u64,
155    len: usize,
156}
157
158impl<T: 'static + Send + InspectMut + MmioIntercept, U: 'static + Send + DmaClient> DeviceBacking
159    for EmulatedDevice<T, U>
160{
161    type Registers = Mapping<T>;
162
163    fn id(&self) -> &str {
164        "emulated"
165    }
166
167    fn map_bar(&mut self, n: u8) -> anyhow::Result<Self::Registers> {
168        if n != 0 {
169            anyhow::bail!("invalid bar {n}");
170        }
171        Ok(Mapping {
172            device: self.device.clone(),
173            addr: (n as u64) << 32,
174            len: self.bar0_len,
175        })
176    }
177
178    fn dma_client(&self) -> Arc<dyn DmaClient> {
179        self.dma_client.clone()
180    }
181
182    fn dma_client_for(&self, _pool: user_driver::DmaPool) -> anyhow::Result<Arc<dyn DmaClient>> {
183        // In the emulated device, we only have one dma client.
184        Ok(self.dma_client.clone())
185    }
186
187    fn max_interrupt_count(&self) -> u32 {
188        self.controller.events.len() as u32
189    }
190
191    fn map_interrupt(&mut self, msix: u32, _cpu: u32) -> anyhow::Result<DeviceInterrupt> {
192        Ok(self
193            .controller
194            .events
195            .get(msix as usize)
196            .with_context(|| format!("invalid msix index {msix}"))?
197            .new_target())
198    }
199}
200
201impl<T: MmioIntercept + Send> DeviceRegisterIo for Mapping<T> {
202    fn len(&self) -> usize {
203        self.len
204    }
205
206    fn read_u32(&self, offset: usize) -> u32 {
207        let mut n = [0; 4];
208        self.device
209            .lock()
210            .mmio_read(self.addr + offset as u64, &mut n)
211            .unwrap();
212        u32::from_ne_bytes(n)
213    }
214
215    fn read_u64(&self, offset: usize) -> u64 {
216        let mut n = [0; 8];
217        self.device
218            .lock()
219            .mmio_read(self.addr + offset as u64, &mut n)
220            .unwrap();
221        u64::from_ne_bytes(n)
222    }
223
224    fn write_u32(&self, offset: usize, data: u32) {
225        self.device
226            .lock()
227            .mmio_write(self.addr + offset as u64, &data.to_ne_bytes())
228            .unwrap();
229    }
230
231    fn write_u64(&self, offset: usize, data: u64) {
232        self.device
233            .lock()
234            .mmio_write(self.addr + offset as u64, &data.to_ne_bytes())
235            .unwrap();
236    }
237}
238
239/// A wrapper around the [`TestMapper`] that generates both [`GuestMemory`] and [`PagePoolAllocator`] backed
240/// by the same underlying memory. Meant to provide shared memory for testing devices.
241pub struct DeviceTestMemory {
242    guest_mem: GuestMemory,
243    payload_mem: GuestMemory,
244    _pool: PagePool,
245    allocator: Arc<PagePoolAllocator>,
246}
247
248impl DeviceTestMemory {
249    /// Creates test memory that leverages the [`TestMapper`] as the backing. It creates 3 accessors for the underlying memory:
250    /// guest_memory [`GuestMemory`] - Has access to the entire range.
251    /// payload_memory [`GuestMemory`] - Has access to the second half of the range.
252    /// dma_client [`PagePoolAllocator`] - Has access to the first half of the range.
253    /// If the `allow_dma` switch is enabled, both guest_memory and payload_memory will report a base_iova of 0.
254    pub fn new(num_pages: u64, allow_dma: bool, pool_name: &str) -> Self {
255        let test_mapper = TestMapper::new(num_pages).unwrap();
256        let sparse_mmap = test_mapper.sparse_mapping();
257        let guest_mem = GuestMemoryAccessWrapper::create_test_guest_memory(sparse_mmap, allow_dma);
258        let pool = PagePool::new(
259            &[MemoryRange::from_4k_gpn_range(0..num_pages / 2)],
260            test_mapper,
261        )
262        .unwrap();
263
264        // Save page pool so that it is not dropped.
265        let allocator = pool.allocator(pool_name.into()).unwrap();
266        let range_half = num_pages / 2 * PAGE_SIZE64;
267        Self {
268            guest_mem: guest_mem.clone(),
269            payload_mem: guest_mem.subrange(range_half, range_half, false).unwrap(),
270            _pool: pool,
271            allocator: Arc::new(allocator),
272        }
273    }
274
275    /// Returns [`GuestMemory`] accessor to the underlying memory. Reports base_iova as 0 if `allow_dma` switch is enabled.
276    pub fn guest_memory(&self) -> GuestMemory {
277        self.guest_mem.clone()
278    }
279
280    /// Returns [`GuestMemory`] accessor to the second half of underlying memory. Reports base_iova as 0 if `allow_dma` switch is enabled.
281    pub fn payload_mem(&self) -> GuestMemory {
282        self.payload_mem.clone()
283    }
284
285    /// Returns [`PagePoolAllocator`] with access to the first half of the underlying memory.
286    pub fn dma_client(&self) -> Arc<PagePoolAllocator> {
287        self.allocator.clone()
288    }
289}
290
291/// Callbacks for the [`DeviceTestDmaClient`]. Tests supply these to customize the behaviour of the dma client.
292pub trait DeviceTestDmaClientCallbacks: Sync + Send {
293    /// Called when the DMA client needs to allocate a new DMA buffer.
294    fn allocate_dma_buffer(
295        &self,
296        allocator: &PagePoolAllocator,
297        total_size: usize,
298    ) -> anyhow::Result<user_driver::memory::MemoryBlock>;
299
300    /// Called when the DMA client needs to attach pending buffers.
301    fn attach_pending_buffers(
302        &self,
303        inner: &PagePoolAllocator,
304    ) -> anyhow::Result<Vec<user_driver::memory::MemoryBlock>>;
305}
306
307/// A DMA client that uses a [`PagePoolAllocator`] as the backing. It can be customized through the use of
308/// [`DeviceTestDmaClientCallbacks`] to modify its behaviour for testing purposes.
309///
310/// # Example
311/// ```rust
312/// use std::sync::Arc;
313/// use user_driver::DmaClient;
314/// use user_driver_emulated_mock::DeviceTestDmaClient;
315/// use page_pool_alloc::PagePoolAllocator;
316///
317/// struct MyCallbacks;
318/// impl user_driver_emulated_mock::DeviceTestDmaClientCallbacks for MyCallbacks {
319///     fn allocate_dma_buffer(
320///         &self,
321///         allocator: &page_pool_alloc::PagePoolAllocator,
322///         total_size: usize,
323///     ) -> anyhow::Result<user_driver::memory::MemoryBlock> {
324///         // Custom test logic here, for example:
325///         anyhow::bail!("allocation failed for testing");
326///     }
327///
328///     fn attach_pending_buffers(
329///         &self,
330///         allocator: &page_pool_alloc::PagePoolAllocator,
331///     ) -> anyhow::Result<Vec<user_driver::memory::MemoryBlock>> {
332///         // Custom test logic here, for example:
333///         anyhow::bail!("attachment failed for testing");
334///     }
335/// }
336///
337/// // Use the above in a test ...
338/// fn test_dma_client() {
339///     let pages = 1000;
340///     let device_test_memory = user_driver_emulated_mock::DeviceTestMemory::new(
341///         pages,
342///         true,
343///         "test_dma_client",
344///     );
345///     let page_pool_allocator = device_test_memory.dma_client();
346///     let dma_client = DeviceTestDmaClient::new(page_pool_allocator, MyCallbacks);
347///
348///     // Use dma_client in tests...
349///     assert!(dma_client.allocate_dma_buffer(4096).is_err());
350/// }
351/// ```
352#[derive(Inspect)]
353#[inspect(transparent)]
354pub struct DeviceTestDmaClient<C>
355where
356    C: DeviceTestDmaClientCallbacks,
357{
358    inner: Arc<PagePoolAllocator>,
359    #[inspect(skip)]
360    callbacks: C,
361}
362
363impl<C: DeviceTestDmaClientCallbacks> DeviceTestDmaClient<C> {
364    /// Creates a new [`DeviceTestDmaClient`] with the given inner allocator.
365    pub fn new(inner: Arc<PagePoolAllocator>, callbacks: C) -> Self {
366        Self { inner, callbacks }
367    }
368}
369
370impl<C: DeviceTestDmaClientCallbacks> DmaClient for DeviceTestDmaClient<C> {
371    fn allocate_dma_buffer(
372        &self,
373        total_size: usize,
374    ) -> anyhow::Result<user_driver::memory::MemoryBlock> {
375        self.callbacks.allocate_dma_buffer(&self.inner, total_size)
376    }
377
378    fn attach_pending_buffers(&self) -> anyhow::Result<Vec<user_driver::memory::MemoryBlock>> {
379        self.callbacks.attach_pending_buffers(&self.inner)
380    }
381}