Skip to main content

virtio_p9/
lib.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4#![expect(missing_docs)]
5#![forbid(unsafe_code)]
6#![cfg(any(windows, target_os = "linux"))]
7
8#[cfg(test)]
9mod integration_tests;
10pub mod resolver;
11
12use anyhow::Context as _;
13use futures::StreamExt;
14use guestmem::GuestMemory;
15use inspect::InspectMut;
16use pal_async::wait::PolledWait;
17use plan9::Plan9FileSystem;
18use task_control::AsyncRun;
19use task_control::Cancelled;
20use task_control::InspectTaskMut;
21use task_control::StopTask;
22use task_control::TaskControl;
23use virtio::DeviceTraits;
24use virtio::QueueResources;
25use virtio::VirtioDevice;
26use virtio::VirtioQueue;
27use virtio::VirtioQueueCallbackWork;
28use virtio::queue::QueueState;
29use virtio::spec::VirtioDeviceFeatures;
30use vmcore::vm_task::VmTaskDriver;
31use vmcore::vm_task::VmTaskDriverSource;
32
33const VIRTIO_9P_F_MOUNT_TAG: u32 = 1;
34
35#[derive(InspectMut)]
36pub struct VirtioPlan9Device {
37    tag: Vec<u8>,
38    driver: VmTaskDriver,
39    #[inspect(mut)]
40    worker: TaskControl<Plan9Worker, Plan9Queue>,
41}
42
43impl VirtioPlan9Device {
44    pub fn new(
45        driver_source: &VmTaskDriverSource,
46        tag: &str,
47        fs: Plan9FileSystem,
48    ) -> VirtioPlan9Device {
49        // The tag uses the same format as 9p protocol strings (2 byte length followed by string).
50        let length = tag.len() + size_of::<u16>();
51
52        // Round the length up to a multiple of 4 to make the read function simpler.
53        let length = (length + 3) & !3;
54        let mut tag_buffer = vec![0u8; length];
55
56        // Write a string preceded by a two byte length.
57        {
58            use std::io::Write;
59            let mut cursor = std::io::Cursor::new(&mut tag_buffer);
60            cursor.write_all(&(tag.len() as u16).to_le_bytes()).unwrap();
61            cursor.write_all(tag.as_bytes()).unwrap();
62        }
63
64        VirtioPlan9Device {
65            tag: tag_buffer,
66            driver: driver_source.simple(),
67            worker: TaskControl::new(Plan9Worker { fs }),
68        }
69    }
70}
71
72impl VirtioDevice for VirtioPlan9Device {
73    fn traits(&self) -> DeviceTraits {
74        DeviceTraits {
75            device_id: virtio::spec::VirtioDeviceType::P9,
76            device_features: VirtioDeviceFeatures::new()
77                .with_device_specific_low(VIRTIO_9P_F_MOUNT_TAG)
78                .with_ring_event_idx(true)
79                .with_ring_indirect_desc(true)
80                .with_ring_packed(true),
81            max_queues: 1,
82            device_register_length: self.tag.len() as u32,
83            ..Default::default()
84        }
85    }
86
87    async fn read_registers_u32(&mut self, offset: u16) -> u32 {
88        assert!(self.tag.len().is_multiple_of(4));
89        assert!(offset.is_multiple_of(4));
90
91        let offset = offset as usize;
92        if offset < self.tag.len() {
93            u32::from_le_bytes(
94                self.tag[offset..offset + 4]
95                    .try_into()
96                    .expect("Incorrect length"),
97            )
98        } else {
99            0
100        }
101    }
102
103    async fn write_registers_u32(&mut self, offset: u16, val: u32) {
104        tracing::warn!(offset, val, "[VIRTIO 9P] Unknown write",);
105    }
106
107    async fn start_queue(
108        &mut self,
109        idx: u16,
110        resources: QueueResources,
111        features: &VirtioDeviceFeatures,
112        initial_state: Option<QueueState>,
113    ) -> anyhow::Result<()> {
114        assert_eq!(idx, 0);
115
116        let queue_event = PolledWait::new(&self.driver, resources.event)
117            .context("failed to create polled wait")?;
118        let queue = VirtioQueue::new(
119            *features,
120            resources.params,
121            resources.guest_memory.clone(),
122            resources.notify,
123            queue_event,
124            initial_state,
125        )
126        .context("failed to create virtio queue")?;
127
128        self.worker.insert(
129            self.driver.clone(),
130            "virtio-9p-queue",
131            Plan9Queue {
132                queue,
133                mem: resources.guest_memory,
134            },
135        );
136        self.worker.start();
137        Ok(())
138    }
139
140    async fn stop_queue(&mut self, idx: u16) -> Option<QueueState> {
141        assert_eq!(idx, 0);
142        if !self.worker.has_state() {
143            return None;
144        }
145        self.worker.stop().await;
146        let state = self.worker.remove().queue.queue_state();
147        Some(state)
148    }
149
150    async fn reset(&mut self) {
151        self.worker.task().fs.reset();
152    }
153}
154
155#[derive(InspectMut)]
156struct Plan9Worker {
157    #[inspect(skip)]
158    fs: Plan9FileSystem,
159}
160
161#[derive(InspectMut)]
162struct Plan9Queue {
163    queue: VirtioQueue,
164    mem: GuestMemory,
165}
166
167impl InspectTaskMut<Plan9Queue> for Plan9Worker {
168    fn inspect_mut(&mut self, req: inspect::Request<'_>, state: Option<&mut Plan9Queue>) {
169        req.respond().merge(self).merge(state);
170    }
171}
172
173impl AsyncRun<Plan9Queue> for Plan9Worker {
174    async fn run(
175        &mut self,
176        stop: &mut StopTask<'_>,
177        state: &mut Plan9Queue,
178    ) -> Result<(), Cancelled> {
179        loop {
180            let work = stop.until_stopped(state.queue.next()).await?;
181            let Some(work) = work else { break };
182            match work {
183                Ok(work) => {
184                    let bytes = process_9p_request(&state.mem, &self.fs, &work);
185                    state.queue.complete(work, bytes);
186                }
187                Err(err) => {
188                    tracing::error!(error = &err as &dyn std::error::Error, "queue error");
189                    break;
190                }
191            }
192        }
193        Ok(())
194    }
195}
196
197fn process_9p_request(
198    mem: &GuestMemory,
199    fs: &Plan9FileSystem,
200    work: &VirtioQueueCallbackWork,
201) -> u32 {
202    // Make a copy of the incoming message.
203    let mut message = vec![0; work.get_payload_length(false) as usize];
204    if let Err(e) = work.read(mem, &mut message) {
205        tracing::error!(
206            error = &e as &dyn std::error::Error,
207            "[VIRTIO 9P] Failed to read guest memory"
208        );
209        return 0;
210    }
211
212    // Allocate a temporary buffer for the response.
213    let mut response = vec![9; work.get_payload_length(true) as usize];
214    let Ok(size) = fs.process_message(&message, &mut response) else {
215        return 0;
216    };
217
218    // Write out the response.
219    if let Err(e) = work.write(mem, &response[0..size]) {
220        tracing::error!(
221            error = &e as &dyn std::error::Error,
222            "[VIRTIO 9P] Failed to write guest memory"
223        );
224        return 0;
225    }
226
227    size as u32
228}