arcbox_virtio_vsock/
backend.rs1use std::collections::HashMap;
4
5use arcbox_virtio_core::error::Result;
6
7use crate::addr::VsockAddr;
8
9pub trait VsockBackend: Send + Sync {
11 fn on_connect(&mut self, addr: VsockAddr) -> Result<()>;
13
14 fn on_send(&mut self, addr: VsockAddr, data: &[u8]) -> Result<usize>;
16
17 fn on_recv(&mut self, addr: VsockAddr, buf: &mut [u8]) -> Result<usize>;
19
20 fn on_close(&mut self, addr: VsockAddr) -> Result<()>;
22
23 fn has_pending_data(&self, addr: VsockAddr) -> bool;
25}
26
27#[derive(Debug, Default)]
29pub struct LoopbackBackend {
30 pending: HashMap<VsockAddr, Vec<u8>>,
32}
33
34impl LoopbackBackend {
35 #[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 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}