Skip to main content

arcbox_virtio_console/
device.rs

1//! `VirtioConsole` device — config, ports, queue handling, `VirtioDevice` impl.
2
3use std::collections::VecDeque;
4use std::sync::{Arc, Mutex};
5
6use tokio::sync::mpsc;
7
8use arcbox_virtio_core::error::{Result, VirtioError};
9use arcbox_virtio_core::queue::VirtQueue;
10use arcbox_virtio_core::{QueueConfig, VirtioDevice, VirtioDeviceId, virtio_bindings};
11
12use crate::{ConsoleIo, StdioConsole};
13
14/// Console device configuration.
15#[derive(Debug, Clone)]
16pub struct ConsoleConfig {
17    /// Number of columns.
18    pub cols: u16,
19    /// Number of rows.
20    pub rows: u16,
21    /// Maximum number of ports.
22    pub max_ports: u32,
23    /// Enable multiport.
24    pub multiport: bool,
25}
26
27impl Default for ConsoleConfig {
28    fn default() -> Self {
29        Self {
30            cols: 80,
31            rows: 25,
32            max_ports: 1,
33            multiport: false,
34        }
35    }
36}
37
38/// Console port state.
39#[derive(Debug)]
40#[allow(dead_code)]
41struct ConsolePort {
42    /// Port number.
43    id: u32,
44    /// Whether the port is open.
45    open: bool,
46    /// Input buffer.
47    input_buffer: VecDeque<u8>,
48    /// Output buffer.
49    output_buffer: VecDeque<u8>,
50}
51
52impl ConsolePort {
53    fn new(id: u32) -> Self {
54        Self {
55            id,
56            open: false,
57            input_buffer: VecDeque::with_capacity(4096),
58            output_buffer: VecDeque::with_capacity(4096),
59        }
60    }
61}
62
63/// `VirtIO` console device.
64#[allow(dead_code)]
65pub struct VirtioConsole {
66    config: ConsoleConfig,
67    features: u64,
68    acked_features: u64,
69    /// Console ports.
70    ports: Vec<ConsolePort>,
71    /// Receive queue (host -> guest).
72    rx_queue: Option<VirtQueue>,
73    /// Transmit queue (guest -> host).
74    tx_queue: Option<VirtQueue>,
75    /// Console I/O handler.
76    io: Option<Arc<Mutex<dyn ConsoleIo>>>,
77    /// Event sender for console input.
78    input_tx: Option<mpsc::UnboundedSender<Vec<u8>>>,
79}
80
81impl VirtioConsole {
82    /// Feature: Console size.
83    pub const FEATURE_SIZE: u64 = 1 << 0;
84    /// Feature: Multiport.
85    pub const FEATURE_MULTIPORT: u64 = 1 << 1;
86    /// Feature: Emergency write.
87    pub const FEATURE_EMERG_WRITE: u64 = 1 << 2;
88    /// `VirtIO` 1.0 feature.
89    pub const FEATURE_VERSION_1: u64 = 1 << virtio_bindings::virtio_config::VIRTIO_F_VERSION_1;
90
91    /// Creates a new console device.
92    #[must_use]
93    pub fn new(config: ConsoleConfig) -> Self {
94        // EVENT_IDX is not advertised for console because activate() does not
95        // propagate it to queues. Add it when console queue setup is updated to
96        // call set_event_idx().
97        let mut features = Self::FEATURE_SIZE | Self::FEATURE_EMERG_WRITE | Self::FEATURE_VERSION_1;
98
99        if config.multiport {
100            features |= Self::FEATURE_MULTIPORT;
101        }
102
103        let mut ports = Vec::with_capacity(config.max_ports as usize);
104        ports.push(ConsolePort::new(0)); // Port 0 is always present
105
106        Self {
107            config,
108            features,
109            acked_features: 0,
110            ports,
111            rx_queue: None,
112            tx_queue: None,
113            io: None,
114            input_tx: None,
115        }
116    }
117
118    /// Creates a console with standard I/O.
119    #[must_use]
120    pub fn with_stdio() -> Self {
121        let mut console = Self::new(ConsoleConfig::default());
122        console.io = Some(Arc::new(Mutex::new(StdioConsole)));
123        console
124    }
125
126    /// Sets the console I/O handler.
127    pub fn set_io(&mut self, io: Arc<Mutex<dyn ConsoleIo>>) {
128        self.io = Some(io);
129    }
130
131    /// Queues input data to be read by the guest.
132    ///
133    /// # Errors
134    ///
135    /// Returns an error if the console is not active.
136    pub fn queue_input(&mut self, data: &[u8]) -> Result<()> {
137        if let Some(port) = self.ports.first_mut() {
138            port.input_buffer.extend(data);
139            Ok(())
140        } else {
141            Err(VirtioError::NotReady("No console port".into()))
142        }
143    }
144
145    /// Reads output data written by the guest.
146    #[must_use]
147    pub fn read_output(&mut self) -> Vec<u8> {
148        if let Some(port) = self.ports.first_mut() {
149            port.output_buffer.drain(..).collect()
150        } else {
151            Vec::new()
152        }
153    }
154
155    /// Handles data from the guest (TX).
156    fn handle_tx(&mut self, data: &[u8]) -> Result<()> {
157        if let Some(port) = self.ports.first_mut() {
158            port.output_buffer.extend(data);
159        }
160
161        if let Some(io) = &self.io {
162            let mut io = io
163                .lock()
164                .map_err(|e| VirtioError::Io(format!("Failed to lock I/O: {e}")))?;
165            io.write(data)
166                .map_err(|e| VirtioError::Io(format!("Write failed: {e}")))?;
167            io.flush()
168                .map_err(|e| VirtioError::Io(format!("Flush failed: {e}")))?;
169        }
170
171        tracing::trace!("Console TX: {} bytes", data.len());
172        Ok(())
173    }
174
175    /// Handles data to the guest (RX).
176    #[allow(dead_code)]
177    fn handle_rx(&mut self, buf: &mut [u8]) -> Result<usize> {
178        if let Some(port) = self.ports.first_mut() {
179            if !port.input_buffer.is_empty() {
180                let len = buf.len().min(port.input_buffer.len());
181                for (i, byte) in port.input_buffer.drain(..len).enumerate() {
182                    buf[i] = byte;
183                }
184                return Ok(len);
185            }
186        }
187
188        if let Some(io) = &self.io {
189            let mut io = io
190                .lock()
191                .map_err(|e| VirtioError::Io(format!("Failed to lock I/O: {e}")))?;
192            let n = io
193                .read(buf)
194                .map_err(|e| VirtioError::Io(format!("Read failed: {e}")))?;
195            tracing::trace!("Console RX: {} bytes", n);
196            return Ok(n);
197        }
198
199        Ok(0)
200    }
201
202    /// Processes the transmit queue.
203    ///
204    /// # Errors
205    ///
206    /// Returns an error if processing fails.
207    pub fn process_tx_queue(&mut self, memory: &[u8]) -> Result<Vec<(u16, u32)>> {
208        let mut tx_data: Vec<(u16, Vec<u8>)> = Vec::new();
209
210        {
211            let queue = self
212                .tx_queue
213                .as_mut()
214                .ok_or_else(|| VirtioError::NotReady("TX queue not ready".into()))?;
215
216            while let Some((head_idx, chain)) = queue.pop_avail() {
217                let mut data = Vec::new();
218
219                for desc in chain {
220                    if !desc.is_write_only() {
221                        let start = desc.addr as usize;
222                        let end = start + desc.len as usize;
223                        if end <= memory.len() {
224                            data.extend_from_slice(&memory[start..end]);
225                        }
226                    }
227                }
228
229                tx_data.push((head_idx, data));
230            }
231        }
232
233        let mut completed = Vec::new();
234        for (head_idx, data) in tx_data {
235            let len = data.len() as u32;
236            self.handle_tx(&data)?;
237            completed.push((head_idx, len));
238        }
239
240        Ok(completed)
241    }
242
243    /// Gets the number of bytes available for RX.
244    #[must_use]
245    pub fn rx_available(&self) -> usize {
246        self.ports
247            .first()
248            .map(|p| p.input_buffer.len())
249            .unwrap_or(0)
250    }
251}
252
253impl VirtioDevice for VirtioConsole {
254    fn device_id(&self) -> VirtioDeviceId {
255        VirtioDeviceId::Console
256    }
257
258    fn features(&self) -> u64 {
259        self.features
260    }
261
262    fn ack_features(&mut self, features: u64) {
263        self.acked_features = self.features & features;
264    }
265
266    fn read_config(&self, offset: u64, data: &mut [u8]) {
267        // Configuration space layout (VirtIO 1.1):
268        // offset 0: cols (u16)
269        // offset 2: rows (u16)
270        // offset 4: max_nr_ports (u32)
271        // offset 8: emerg_wr (u32)
272        let config_data = [
273            self.config.cols.to_le_bytes().as_slice(),
274            &self.config.rows.to_le_bytes(),
275            &self.config.max_ports.to_le_bytes(),
276            &0u32.to_le_bytes(), // emerg_wr
277        ]
278        .concat();
279
280        let offset = offset as usize;
281        let len = data.len().min(config_data.len().saturating_sub(offset));
282        if len > 0 {
283            data[..len].copy_from_slice(&config_data[offset..offset + len]);
284        }
285    }
286
287    fn write_config(&mut self, offset: u64, data: &[u8]) {
288        // Handle emergency write at offset 8
289        if offset == 8 && data.len() >= 4 {
290            let ch = u32::from_le_bytes([data[0], data[1], data[2], data[3]]);
291            if ch != 0 {
292                if let Some(c) = char::from_u32(ch) {
293                    eprint!("{c}");
294                }
295            }
296        }
297    }
298
299    fn activate(&mut self) -> Result<()> {
300        self.rx_queue = Some(VirtQueue::new(256)?);
301        self.tx_queue = Some(VirtQueue::new(256)?);
302
303        if let Some(port) = self.ports.first_mut() {
304            port.open = true;
305        }
306
307        tracing::info!(
308            "VirtIO console activated: {}x{}, {} ports",
309            self.config.cols,
310            self.config.rows,
311            self.config.max_ports
312        );
313
314        Ok(())
315    }
316
317    fn reset(&mut self) {
318        self.acked_features = 0;
319        self.rx_queue = None;
320        self.tx_queue = None;
321
322        for port in &mut self.ports {
323            port.open = false;
324            port.input_buffer.clear();
325            port.output_buffer.clear();
326        }
327    }
328
329    fn process_queue(
330        &mut self,
331        queue_idx: u16,
332        memory: &mut [u8],
333        queue_config: &QueueConfig,
334    ) -> Result<Vec<(u16, u32)>> {
335        // Queue 0 = RX (host→guest), Queue 1 = TX (guest→host).
336        // We only handle TX here — extract guest output from descriptors.
337        if queue_idx != 1 {
338            return Ok(Vec::new());
339        }
340
341        if !queue_config.ready || queue_config.size == 0 {
342            return Ok(Vec::new());
343        }
344
345        // Translate GPAs to slice offsets by subtracting gpa_base (checked to
346        // guard against a malicious guest providing a GPA below the RAM base).
347        let gpa_base = queue_config.gpa_base as usize;
348        let desc_addr = (queue_config.desc_addr as usize)
349            .checked_sub(gpa_base)
350            .ok_or_else(|| {
351                tracing::warn!(
352                    "invalid desc GPA {:#x} below ram base {:#x}",
353                    queue_config.desc_addr,
354                    gpa_base
355                );
356                VirtioError::InvalidQueue("desc GPA below ram base".into())
357            })?;
358        let avail_addr = (queue_config.avail_addr as usize)
359            .checked_sub(gpa_base)
360            .ok_or_else(|| {
361                tracing::warn!(
362                    "invalid avail GPA {:#x} below ram base {:#x}",
363                    queue_config.avail_addr,
364                    gpa_base
365                );
366                VirtioError::InvalidQueue("avail GPA below ram base".into())
367            })?;
368        let used_addr = (queue_config.used_addr as usize)
369            .checked_sub(gpa_base)
370            .ok_or_else(|| {
371                tracing::warn!(
372                    "invalid used GPA {:#x} below ram base {:#x}",
373                    queue_config.used_addr,
374                    gpa_base
375                );
376                VirtioError::InvalidQueue("used GPA below ram base".into())
377            })?;
378        let queue_size = queue_config.size as usize;
379
380        if avail_addr + 4 > memory.len() {
381            return Ok(Vec::new());
382        }
383        let avail_idx = u16::from_le_bytes([memory[avail_addr + 2], memory[avail_addr + 3]]);
384
385        if used_addr + 4 > memory.len() {
386            return Ok(Vec::new());
387        }
388        let used_idx_ref = &memory[used_addr + 2..used_addr + 4];
389        let mut used_idx = u16::from_le_bytes([used_idx_ref[0], used_idx_ref[1]]);
390
391        let mut completions = Vec::new();
392
393        while used_idx != avail_idx {
394            let avail_ring_off = avail_addr + 4 + (used_idx as usize % queue_size) * 2;
395            if avail_ring_off + 2 > memory.len() {
396                break;
397            }
398            let head_idx = u16::from_le_bytes([memory[avail_ring_off], memory[avail_ring_off + 1]]);
399
400            // Walk descriptor chain, extract TX data.
401            let mut idx = head_idx as usize;
402            let mut total_len = 0u32;
403            for _ in 0..queue_size {
404                let d_off = desc_addr + idx * 16;
405                if d_off + 16 > memory.len() {
406                    break;
407                }
408                let addr = match (u64::from_le_bytes(memory[d_off..d_off + 8].try_into().unwrap())
409                    as usize)
410                    .checked_sub(gpa_base)
411                {
412                    Some(a) => a,
413                    None => continue,
414                };
415                let len = u32::from_le_bytes(memory[d_off + 8..d_off + 12].try_into().unwrap());
416                let flags = u16::from_le_bytes(memory[d_off + 12..d_off + 14].try_into().unwrap());
417                let next = u16::from_le_bytes(memory[d_off + 14..d_off + 16].try_into().unwrap());
418
419                let is_write = flags & 2 != 0; // VIRTQ_DESC_F_WRITE
420                if !is_write {
421                    // Read-only descriptor = data FROM guest (TX output).
422                    let start = addr;
423                    let end = start + len as usize;
424                    if end <= memory.len() {
425                        let data = &memory[start..end];
426                        if let Some(port) = self.ports.first_mut() {
427                            port.output_buffer.extend(data.iter().copied());
428                            // Flush on newline.
429                            while let Some(pos) =
430                                port.output_buffer.iter().position(|&b| b == b'\n')
431                            {
432                                let line: Vec<u8> = port.output_buffer.drain(..=pos).collect();
433                                if let Ok(s) = std::str::from_utf8(&line) {
434                                    tracing::info!(target: "guest_console", "{}", s.trim_end());
435                                }
436                            }
437                        }
438                        total_len += len;
439                    }
440                }
441
442                if flags & 1 == 0 {
443                    break; // No NEXT
444                }
445                idx = next as usize;
446            }
447
448            let used_ring_off = used_addr + 4 + (used_idx as usize % queue_size) * 8;
449            if used_ring_off + 8 <= memory.len() {
450                memory[used_ring_off..used_ring_off + 4]
451                    .copy_from_slice(&(head_idx as u32).to_le_bytes());
452                memory[used_ring_off + 4..used_ring_off + 8]
453                    .copy_from_slice(&total_len.to_le_bytes());
454            }
455
456            used_idx = used_idx.wrapping_add(1);
457            completions.push((head_idx, total_len));
458        }
459
460        if !completions.is_empty() {
461            std::sync::atomic::fence(std::sync::atomic::Ordering::Release);
462            let new_used = used_idx.to_le_bytes();
463            memory[used_addr + 2] = new_used[0];
464            memory[used_addr + 3] = new_used[1];
465
466            // Only write avail_event when VIRTIO_F_EVENT_IDX was negotiated;
467            // without it the field lives past the used ring and the write
468            // would corrupt guest memory. Console does not advertise EVENT_IDX
469            // so this branch is currently unreachable.
470            if (self.acked_features & arcbox_virtio_core::queue::VIRTIO_F_EVENT_IDX) != 0 {
471                let avail_event_off = used_addr + 4 + 8 * queue_size;
472                if avail_event_off + 2 <= memory.len() {
473                    let ae = avail_idx.to_le_bytes();
474                    memory[avail_event_off] = ae[0];
475                    memory[avail_event_off + 1] = ae[1];
476                }
477            }
478        }
479
480        Ok(completions)
481    }
482}
483
484#[cfg(test)]
485mod tests {
486    use super::*;
487    use crate::BufferConsole;
488
489    #[test]
490    fn test_console_creation() {
491        let console = VirtioConsole::new(ConsoleConfig::default());
492        assert_eq!(console.device_id(), VirtioDeviceId::Console);
493        assert!(console.features() & VirtioConsole::FEATURE_SIZE != 0);
494    }
495
496    #[test]
497    fn test_console_config_read() {
498        let config = ConsoleConfig {
499            cols: 120,
500            rows: 40,
501            max_ports: 4,
502            multiport: false,
503        };
504        let console = VirtioConsole::new(config);
505
506        let mut data = [0u8; 8];
507        console.read_config(0, &mut data);
508
509        assert_eq!(u16::from_le_bytes([data[0], data[1]]), 120); // cols
510        assert_eq!(u16::from_le_bytes([data[2], data[3]]), 40); // rows
511        assert_eq!(u32::from_le_bytes([data[4], data[5], data[6], data[7]]), 4);
512        // max_ports
513    }
514
515    #[test]
516    fn test_console_input_queue() {
517        let mut console = VirtioConsole::new(ConsoleConfig::default());
518        console.activate().unwrap();
519
520        console.queue_input(b"test input").unwrap();
521        assert_eq!(console.rx_available(), 10);
522    }
523
524    #[test]
525    fn test_console_output() {
526        let buffer = Arc::new(Mutex::new(BufferConsole::new()));
527        let mut console = VirtioConsole::new(ConsoleConfig::default());
528        console.set_io(buffer.clone());
529        console.activate().unwrap();
530
531        console.handle_tx(b"Hello, World!").unwrap();
532
533        let output = buffer.lock().unwrap().take_output();
534        assert_eq!(&output, b"Hello, World!");
535    }
536
537    #[test]
538    fn test_console_multiport_feature() {
539        let config = ConsoleConfig {
540            multiport: true,
541            ..Default::default()
542        };
543        let console = VirtioConsole::new(config);
544        assert!(console.features() & VirtioConsole::FEATURE_MULTIPORT != 0);
545    }
546
547    #[test]
548    fn test_console_activate_and_reset() {
549        let mut console = VirtioConsole::new(ConsoleConfig::default());
550
551        console.activate().unwrap();
552        assert!(console.rx_queue.is_some());
553        assert!(console.tx_queue.is_some());
554
555        console.reset();
556        assert!(console.rx_queue.is_none());
557        assert!(console.tx_queue.is_none());
558        assert_eq!(console.acked_features, 0);
559    }
560
561    #[test]
562    fn test_console_read_output() {
563        let mut console = VirtioConsole::new(ConsoleConfig::default());
564        console.activate().unwrap();
565
566        let output = console.read_output();
567        assert!(output.is_empty());
568
569        console.handle_tx(b"test output").unwrap();
570        let output = console.read_output();
571        assert_eq!(&output, b"test output");
572
573        let output2 = console.read_output();
574        assert!(output2.is_empty());
575    }
576
577    #[test]
578    fn test_console_queue_input_not_ready() {
579        let mut console = VirtioConsole::new(ConsoleConfig::default());
580
581        console.ports.clear();
582
583        let result = console.queue_input(b"test");
584        assert!(result.is_err());
585    }
586
587    #[test]
588    fn test_console_config_write() {
589        let mut console = VirtioConsole::new(ConsoleConfig::default());
590
591        let emergency_char = 'X' as u32;
592        console.write_config(8, &emergency_char.to_le_bytes());
593
594        // Should not crash — emergency write goes to stderr.
595    }
596
597    #[test]
598    fn test_console_feature_negotiation() {
599        let mut console = VirtioConsole::new(ConsoleConfig::default());
600
601        let offered = console.features();
602        assert!(offered & VirtioConsole::FEATURE_VERSION_1 != 0);
603
604        console.ack_features(VirtioConsole::FEATURE_SIZE | VirtioConsole::FEATURE_VERSION_1);
605        assert!(console.acked_features & VirtioConsole::FEATURE_SIZE != 0);
606    }
607
608    #[test]
609    fn test_console_with_stdio() {
610        let console = VirtioConsole::with_stdio();
611        assert!(console.io.is_some());
612    }
613
614    #[test]
615    fn test_console_config_partial_read() {
616        let console = VirtioConsole::new(ConsoleConfig {
617            cols: 80,
618            rows: 25,
619            max_ports: 1,
620            multiport: false,
621        });
622
623        let mut data = [0u8; 2];
624        console.read_config(0, &mut data);
625        assert_eq!(u16::from_le_bytes(data), 80);
626
627        let mut data2 = [0u8; 2];
628        console.read_config(2, &mut data2);
629        assert_eq!(u16::from_le_bytes(data2), 25);
630    }
631
632    #[test]
633    fn test_console_rx_available_empty() {
634        let console = VirtioConsole::new(ConsoleConfig::default());
635        assert_eq!(console.rx_available(), 0);
636    }
637}