arcbox_virtio_vsock/
tcp.rs1use 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
13pub struct TcpBackend {
17 guest_cid: u64,
19 base_port: u16,
21 connections: RwLock<HashMap<VsockAddr, TcpStream>>,
23 listeners: RwLock<HashMap<u32, TcpListener>>,
25}
26
27impl TcpBackend {
28 #[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 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 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 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}