Skip to main content

arcbox_virtio_vsock/
backend.rs

1//! `VsockBackend` trait + `LoopbackBackend` (in-process echo).
2
3use std::collections::HashMap;
4
5use arcbox_virtio_core::error::Result;
6
7use crate::addr::VsockAddr;
8
9/// Vsock backend trait for handling host-side socket operations.
10pub trait VsockBackend: Send + Sync {
11    /// Called when guest requests a connection.
12    fn on_connect(&mut self, addr: VsockAddr) -> Result<()>;
13
14    /// Called when guest sends data.
15    fn on_send(&mut self, addr: VsockAddr, data: &[u8]) -> Result<usize>;
16
17    /// Called when guest requests data.
18    fn on_recv(&mut self, addr: VsockAddr, buf: &mut [u8]) -> Result<usize>;
19
20    /// Called when guest closes connection.
21    fn on_close(&mut self, addr: VsockAddr) -> Result<()>;
22
23    /// Checks if there's pending data for a connection.
24    fn has_pending_data(&self, addr: VsockAddr) -> bool;
25}
26
27/// Loopback vsock backend for testing.
28#[derive(Debug, Default)]
29pub struct LoopbackBackend {
30    /// Pending data per connection.
31    pending: HashMap<VsockAddr, Vec<u8>>,
32}
33
34impl LoopbackBackend {
35    /// Creates a new loopback backend.
36    #[must_use]
37    pub fn new() -> Self {
38        Self::default()
39    }
40}
41
42impl VsockBackend for LoopbackBackend {
43    fn on_connect(&mut self, addr: VsockAddr) -> Result<()> {
44        self.pending.insert(addr, Vec::new());
45        tracing::debug!("Loopback: connection from {:?}", addr);
46        Ok(())
47    }
48
49    fn on_send(&mut self, addr: VsockAddr, data: &[u8]) -> Result<usize> {
50        // Echo back the data
51        if let Some(buf) = self.pending.get_mut(&addr) {
52            buf.extend_from_slice(data);
53        }
54        Ok(data.len())
55    }
56
57    fn on_recv(&mut self, addr: VsockAddr, buf: &mut [u8]) -> Result<usize> {
58        if let Some(pending) = self.pending.get_mut(&addr) {
59            let len = buf.len().min(pending.len());
60            buf[..len].copy_from_slice(&pending[..len]);
61            pending.drain(..len);
62            Ok(len)
63        } else {
64            Ok(0)
65        }
66    }
67
68    fn on_close(&mut self, addr: VsockAddr) -> Result<()> {
69        self.pending.remove(&addr);
70        tracing::debug!("Loopback: connection closed {:?}", addr);
71        Ok(())
72    }
73
74    fn has_pending_data(&self, addr: VsockAddr) -> bool {
75        self.pending.get(&addr).map_or(false, |b| !b.is_empty())
76    }
77}
78
79#[cfg(test)]
80mod tests {
81    use super::*;
82
83    #[test]
84    fn test_loopback_backend() {
85        let mut backend = LoopbackBackend::new();
86        let addr = VsockAddr::new(3, 1234);
87
88        backend.on_connect(addr).unwrap();
89
90        let data = b"hello world";
91        let sent = backend.on_send(addr, data).unwrap();
92        assert_eq!(sent, data.len());
93
94        assert!(backend.has_pending_data(addr));
95
96        let mut buf = [0u8; 64];
97        let received = backend.on_recv(addr, &mut buf).unwrap();
98        assert_eq!(received, data.len());
99        assert_eq!(&buf[..received], data);
100
101        backend.on_close(addr).unwrap();
102        assert!(!backend.has_pending_data(addr));
103    }
104}