Skip to main content

arcbox_virtio_vsock/
tcp.rs

1//! TCP-based vsock backend — maps vsock ports to TCP ports for host-side handling.
2
3use std::collections::HashMap;
4use std::io::{Read, Write};
5use std::net::{TcpListener, TcpStream};
6use std::sync::RwLock;
7
8use arcbox_virtio_core::error::{Result, VirtioError};
9
10use crate::addr::{HOST_CID, VsockAddr};
11use crate::backend::VsockBackend;
12
13/// TCP-based vsock backend.
14///
15/// Maps vsock ports to TCP ports for host-side handling.
16pub struct TcpBackend {
17    /// Guest CID.
18    guest_cid: u64,
19    /// Base TCP port (vsock port N maps to TCP port base + N).
20    base_port: u16,
21    /// Active connections.
22    connections: RwLock<HashMap<VsockAddr, TcpStream>>,
23    /// Listeners for incoming connections.
24    listeners: RwLock<HashMap<u32, TcpListener>>,
25}
26
27impl TcpBackend {
28    /// Creates a new TCP backend.
29    #[must_use]
30    pub fn new(guest_cid: u64, base_port: u16) -> Self {
31        Self {
32            guest_cid,
33            base_port,
34            connections: RwLock::new(HashMap::new()),
35            listeners: RwLock::new(HashMap::new()),
36        }
37    }
38
39    /// Listens on a vsock port.
40    pub fn listen(&self, port: u32) -> Result<()> {
41        let tcp_port = self.base_port + port as u16;
42        let listener = TcpListener::bind(format!("127.0.0.1:{tcp_port}"))
43            .map_err(|e| VirtioError::Io(format!("Failed to bind: {e}")))?;
44
45        listener
46            .set_nonblocking(true)
47            .map_err(|e| VirtioError::Io(format!("Failed to set nonblocking: {e}")))?;
48
49        self.listeners.write().unwrap().insert(port, listener);
50        tracing::info!("Vsock listening on port {} (TCP {})", port, tcp_port);
51        Ok(())
52    }
53
54    /// Accepts a pending connection.
55    pub fn accept(&self, port: u32) -> Result<Option<VsockAddr>> {
56        let listeners = self.listeners.read().unwrap();
57        if let Some(listener) = listeners.get(&port) {
58            match listener.accept() {
59                Ok((stream, _addr)) => {
60                    stream
61                        .set_nonblocking(true)
62                        .map_err(|e| VirtioError::Io(format!("Failed to set nonblocking: {e}")))?;
63
64                    let local = VsockAddr::new(HOST_CID, port);
65                    let remote = VsockAddr::new(self.guest_cid, port);
66
67                    self.connections.write().unwrap().insert(remote, stream);
68                    Ok(Some(local))
69                }
70                Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => Ok(None),
71                Err(e) => Err(VirtioError::Io(format!("Accept failed: {e}"))),
72            }
73        } else {
74            Ok(None)
75        }
76    }
77}
78
79impl VsockBackend for TcpBackend {
80    fn on_connect(&mut self, addr: VsockAddr) -> Result<()> {
81        let tcp_port = self.base_port + addr.port as u16;
82        let stream = TcpStream::connect(format!("127.0.0.1:{tcp_port}"))
83            .map_err(|e| VirtioError::Io(format!("Connect failed: {e}")))?;
84
85        stream
86            .set_nonblocking(true)
87            .map_err(|e| VirtioError::Io(format!("Failed to set nonblocking: {e}")))?;
88
89        self.connections.write().unwrap().insert(addr, stream);
90        Ok(())
91    }
92
93    fn on_send(&mut self, addr: VsockAddr, data: &[u8]) -> Result<usize> {
94        let mut connections = self.connections.write().unwrap();
95        if let Some(stream) = connections.get_mut(&addr) {
96            stream
97                .write(data)
98                .map_err(|e| VirtioError::Io(format!("Send failed: {e}")))
99        } else {
100            Err(VirtioError::InvalidOperation("Connection not found".into()))
101        }
102    }
103
104    fn on_recv(&mut self, addr: VsockAddr, buf: &mut [u8]) -> Result<usize> {
105        let mut connections = self.connections.write().unwrap();
106        if let Some(stream) = connections.get_mut(&addr) {
107            match stream.read(buf) {
108                Ok(n) => Ok(n),
109                Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => Ok(0),
110                Err(e) => Err(VirtioError::Io(format!("Recv failed: {e}"))),
111            }
112        } else {
113            Err(VirtioError::InvalidOperation("Connection not found".into()))
114        }
115    }
116
117    fn on_close(&mut self, addr: VsockAddr) -> Result<()> {
118        self.connections.write().unwrap().remove(&addr);
119        Ok(())
120    }
121
122    fn has_pending_data(&self, addr: VsockAddr) -> bool {
123        // TCP streams don't have a simple way to check pending data; existence
124        // of the connection entry stands in as a coarse signal.
125        self.connections.read().unwrap().contains_key(&addr)
126    }
127}
128
129impl std::fmt::Debug for TcpBackend {
130    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
131        f.debug_struct("TcpBackend")
132            .field("guest_cid", &self.guest_cid)
133            .field("base_port", &self.base_port)
134            .finish()
135    }
136}
137
138#[cfg(test)]
139mod tests {
140    use super::*;
141
142    #[test]
143    fn test_tcp_backend_creation() {
144        let backend = TcpBackend::new(3, 10000);
145        assert!(!backend.has_pending_data(VsockAddr::new(3, 1234)));
146    }
147}