Skip to main content

subetha_cxc/
locale_vsock.rs

1//! `locale_vsock`: host-VM byte streaming that bypasses the network
2//! stack (cross-platform).
3//!
4//! A "remote-but-same-machine" locale member, between QUIC (cross-host)
5//! and ShmFs (same-host shared memory): the hypervisor forwards bytes
6//! between guest and host with no TCP/IP, DNS, or NIC involved. Only the
7//! socket family is gated; the [`HostVmSocket`] stream API
8//! (`listen_loopback` / `accept` / `connect_loopback` / `send` / `recv`)
9//! is shared:
10//!
11//! - Linux (`#[cfg(target_os = "linux")]`): vsock(7) - `AF_VSOCK`
12//!   `SOCK_STREAM`, addressed by `(cid, port)`. The raw `VsockSocket`
13//!   exposes the explicit-CID primitive for real guest<->host links;
14//!   `HostVmSocket` wraps it for the loopback path (`VMADDR_CID_LOCAL`).
15//! - Windows (`#[cfg(windows)]`): Hyper-V sockets - `AF_HYPERV`
16//!   `SOCK_STREAM` / `HV_PROTOCOL_RAW`, addressed by `(VmId, ServiceId)`
17//!   GUIDs. The port maps to a ServiceId via the documented VSOCK
18//!   template GUID; loopback uses `HV_GUID_LOOPBACK`.
19//!
20//! Both expose loopback (same-partition) streaming for same-host IPC and
21//! self-test; the same code reaches a real guest/host by addressing a
22//! peer CID / VmId instead of loopback. Loopback needs the kernel
23//! `vsock_loopback` module (Linux) or the Hyper-V / Virtual Machine
24//! Platform feature registering the `AF_HYPERV` provider (Windows); when
25//! absent the constructors return `Err` so callers degrade gracefully.
26
27#![cfg(any(target_os = "linux", windows))]
28
29// ---------------------------------------------------------------------
30// Linux: vsock(7) over AF_VSOCK.
31// ---------------------------------------------------------------------
32
33#[cfg(target_os = "linux")]
34mod vsock_impl {
35    use std::io;
36    use std::os::unix::io::{AsRawFd, FromRawFd, RawFd};
37
38    /// CID values defined by vsock(7).
39    pub const VMADDR_CID_ANY: u32 = 0xFFFFFFFF;
40    pub const VMADDR_CID_HOST: u32 = 2;
41    /// Loopback CID (same-host, no hypervisor transport): requires the
42    /// `vsock_loopback` kernel module (Linux 5.6+).
43    pub const VMADDR_CID_LOCAL: u32 = 1;
44
45    /// A vsock SOCK_STREAM socket. Owns its fd; closes on Drop. This is
46    /// the explicit-CID primitive; for the loopback path use the
47    /// cross-platform [`HostVmSocket`].
48    pub struct VsockSocket {
49        fd: RawFd,
50    }
51
52    unsafe impl Send for VsockSocket {}
53    unsafe impl Sync for VsockSocket {}
54
55    impl VsockSocket {
56        /// Create a fresh vsock SOCK_STREAM socket (not yet bound /
57        /// connected).
58        pub fn new() -> io::Result<Self> {
59            let fd = unsafe { libc::socket(libc::AF_VSOCK, libc::SOCK_STREAM, 0) };
60            if fd < 0 {
61                return Err(io::Error::last_os_error());
62            }
63            Ok(Self { fd })
64        }
65
66        /// Bind to (cid, port). Use VMADDR_CID_ANY to bind to any CID.
67        pub fn bind(&self, cid: u32, port: u32) -> io::Result<()> {
68            let mut addr: libc::sockaddr_vm = unsafe { std::mem::zeroed() };
69            addr.svm_family = libc::AF_VSOCK as _;
70            addr.svm_cid = cid;
71            addr.svm_port = port;
72            let rc = unsafe {
73                libc::bind(
74                    self.fd,
75                    &addr as *const _ as *const libc::sockaddr,
76                    std::mem::size_of::<libc::sockaddr_vm>() as _,
77                )
78            };
79            if rc < 0 {
80                Err(io::Error::last_os_error())
81            } else {
82                Ok(())
83            }
84        }
85
86        /// Listen on a bound socket.
87        pub fn listen(&self, backlog: i32) -> io::Result<()> {
88            let rc = unsafe { libc::listen(self.fd, backlog) };
89            if rc < 0 {
90                Err(io::Error::last_os_error())
91            } else {
92                Ok(())
93            }
94        }
95
96        /// Accept one incoming vsock connection. Returns the connected
97        /// socket; the listener stays alive.
98        pub fn accept(&self) -> io::Result<VsockSocket> {
99            let fd =
100                unsafe { libc::accept(self.fd, std::ptr::null_mut(), std::ptr::null_mut()) };
101            if fd < 0 {
102                return Err(io::Error::last_os_error());
103            }
104            Ok(VsockSocket { fd })
105        }
106
107        /// Connect this socket to (cid, port).
108        pub fn connect(&self, cid: u32, port: u32) -> io::Result<()> {
109            let mut addr: libc::sockaddr_vm = unsafe { std::mem::zeroed() };
110            addr.svm_family = libc::AF_VSOCK as _;
111            addr.svm_cid = cid;
112            addr.svm_port = port;
113            let rc = unsafe {
114                libc::connect(
115                    self.fd,
116                    &addr as *const _ as *const libc::sockaddr,
117                    std::mem::size_of::<libc::sockaddr_vm>() as _,
118                )
119            };
120            if rc < 0 {
121                Err(io::Error::last_os_error())
122            } else {
123                Ok(())
124            }
125        }
126
127        /// Blocking send. Returns bytes written.
128        pub fn send(&self, buf: &[u8]) -> io::Result<usize> {
129            let n = unsafe { libc::send(self.fd, buf.as_ptr() as *const _, buf.len(), 0) };
130            if n < 0 {
131                Err(io::Error::last_os_error())
132            } else {
133                Ok(n as usize)
134            }
135        }
136
137        /// Blocking recv. Returns bytes read.
138        pub fn recv(&self, buf: &mut [u8]) -> io::Result<usize> {
139            let n = unsafe { libc::recv(self.fd, buf.as_mut_ptr() as *mut _, buf.len(), 0) };
140            if n < 0 {
141                Err(io::Error::last_os_error())
142            } else {
143                Ok(n as usize)
144            }
145        }
146    }
147
148    impl AsRawFd for VsockSocket {
149        fn as_raw_fd(&self) -> RawFd {
150            self.fd
151        }
152    }
153
154    impl FromRawFd for VsockSocket {
155        unsafe fn from_raw_fd(fd: RawFd) -> Self {
156            Self { fd }
157        }
158    }
159
160    impl Drop for VsockSocket {
161        fn drop(&mut self) {
162            if self.fd >= 0 {
163                unsafe { libc::close(self.fd) };
164            }
165        }
166    }
167
168    /// Cross-platform host-VM stream socket (Linux side). Wraps
169    /// [`VsockSocket`] with the loopback-oriented port API shared with
170    /// the Windows `AF_HYPERV` implementation.
171    pub struct HostVmSocket {
172        inner: VsockSocket,
173    }
174
175    impl HostVmSocket {
176        /// Bind + listen on `port` for loopback connections.
177        pub fn listen_loopback(port: u32) -> io::Result<Self> {
178            let s = VsockSocket::new()?;
179            // Bind to ANY so loopback (CID_LOCAL) connectors reach us.
180            s.bind(VMADDR_CID_ANY, port)?;
181            s.listen(16)?;
182            Ok(Self { inner: s })
183        }
184
185        /// Accept one loopback connection.
186        pub fn accept(&self) -> io::Result<Self> {
187            Ok(Self { inner: self.inner.accept()? })
188        }
189
190        /// Connect to a loopback listener on `port` (same host).
191        pub fn connect_loopback(port: u32) -> io::Result<Self> {
192            let s = VsockSocket::new()?;
193            s.connect(VMADDR_CID_LOCAL, port)?;
194            Ok(Self { inner: s })
195        }
196
197        pub fn send(&self, buf: &[u8]) -> io::Result<usize> {
198            self.inner.send(buf)
199        }
200        pub fn recv(&self, buf: &mut [u8]) -> io::Result<usize> {
201            self.inner.recv(buf)
202        }
203    }
204}
205
206#[cfg(target_os = "linux")]
207pub use vsock_impl::{
208    HostVmSocket, VsockSocket, VMADDR_CID_ANY, VMADDR_CID_HOST, VMADDR_CID_LOCAL,
209};
210
211// ---------------------------------------------------------------------
212// Windows: Hyper-V sockets over AF_HYPERV.
213// ---------------------------------------------------------------------
214
215#[cfg(windows)]
216mod hyperv_impl {
217    use std::io;
218    use std::sync::Once;
219    use windows_sys::core::GUID;
220    use windows_sys::Win32::Networking::WinSock::{
221        accept, bind, closesocket, connect, listen, recv, send, socket, WSAStartup,
222        INVALID_SOCKET, SOCKADDR, SOCKET, SOCKET_ERROR, SOCK_STREAM, WSADATA,
223    };
224
225    const AF_HYPERV: i32 = 34;
226    const HV_PROTOCOL_RAW: i32 = 1;
227
228    /// Loopback VmId - connecting here reaches the same partition.
229    const HV_GUID_LOOPBACK: GUID = GUID {
230        data1: 0xe0e1_6197,
231        data2: 0xdd56,
232        data3: 0x4a10,
233        data4: [0x91, 0x95, 0x5e, 0xe7, 0xa1, 0x55, 0xa8, 0x38],
234    };
235    /// Wildcard VmId - listeners bind here to accept from all partitions.
236    const HV_GUID_WILDCARD: GUID = GUID {
237        data1: 0,
238        data2: 0,
239        data3: 0,
240        data4: [0; 8],
241    };
242
243    /// The documented Linux-guest VSOCK service-ID template; `Data1` is
244    /// the port. Mapping a port through this keeps the address model the
245    /// same as the Linux `(cid, port)` side.
246    fn service_id_for_port(port: u32) -> GUID {
247        GUID {
248            data1: port,
249            data2: 0xfacb,
250            data3: 0x11e6,
251            data4: [0xbd, 0x58, 0x64, 0x00, 0x6a, 0x79, 0x86, 0xd3],
252        }
253    }
254
255    #[repr(C)]
256    #[derive(Clone, Copy)]
257    struct SOCKADDR_HV {
258        family: u16,
259        reserved: u16,
260        vm_id: GUID,
261        service_id: GUID,
262    }
263
264    fn ensure_winsock() {
265        static START: Once = Once::new();
266        START.call_once(|| {
267            let mut data: WSADATA = unsafe { std::mem::zeroed() };
268            // MAKEWORD(2, 2) = 0x0202.
269            unsafe { WSAStartup(0x0202, &mut data) };
270        });
271    }
272
273    fn last_err() -> io::Error {
274        io::Error::last_os_error()
275    }
276
277    /// Cross-platform host-VM stream socket (Windows side): a Hyper-V
278    /// socket addressed by VmId + ServiceId GUIDs. Same surface as the
279    /// Linux `AF_VSOCK` implementation.
280    pub struct HostVmSocket {
281        sock: SOCKET,
282    }
283
284    unsafe impl Send for HostVmSocket {}
285    unsafe impl Sync for HostVmSocket {}
286
287    impl HostVmSocket {
288        fn raw_socket() -> io::Result<SOCKET> {
289            ensure_winsock();
290            let s = unsafe { socket(AF_HYPERV, SOCK_STREAM, HV_PROTOCOL_RAW) };
291            if s == INVALID_SOCKET {
292                Err(io::Error::other(format!(
293                    "AF_HYPERV socket unavailable ({}); enable the Hyper-V / \
294                     Virtual Machine Platform feature",
295                    last_err()
296                )))
297            } else {
298                Ok(s)
299            }
300        }
301
302        fn addr(vm_id: GUID, port: u32) -> SOCKADDR_HV {
303            SOCKADDR_HV {
304                family: AF_HYPERV as u16,
305                reserved: 0,
306                vm_id,
307                service_id: service_id_for_port(port),
308            }
309        }
310
311        /// Bind + listen on `port` for loopback connections.
312        pub fn listen_loopback(port: u32) -> io::Result<Self> {
313            let sock = Self::raw_socket()?;
314            let addr = Self::addr(HV_GUID_WILDCARD, port);
315            let rc = unsafe {
316                bind(
317                    sock,
318                    &addr as *const _ as *const SOCKADDR,
319                    std::mem::size_of::<SOCKADDR_HV>() as i32,
320                )
321            };
322            if rc == SOCKET_ERROR {
323                let e = last_err();
324                unsafe { closesocket(sock) };
325                return Err(e);
326            }
327            if unsafe { listen(sock, 16) } == SOCKET_ERROR {
328                let e = last_err();
329                unsafe { closesocket(sock) };
330                return Err(e);
331            }
332            Ok(Self { sock })
333        }
334
335        /// Accept one loopback connection.
336        pub fn accept(&self) -> io::Result<Self> {
337            let s = unsafe { accept(self.sock, std::ptr::null_mut(), std::ptr::null_mut()) };
338            if s == INVALID_SOCKET {
339                Err(last_err())
340            } else {
341                Ok(Self { sock: s })
342            }
343        }
344
345        /// Connect to a loopback listener on `port` (same partition).
346        pub fn connect_loopback(port: u32) -> io::Result<Self> {
347            let sock = Self::raw_socket()?;
348            let addr = Self::addr(HV_GUID_LOOPBACK, port);
349            let rc = unsafe {
350                connect(
351                    sock,
352                    &addr as *const _ as *const SOCKADDR,
353                    std::mem::size_of::<SOCKADDR_HV>() as i32,
354                )
355            };
356            if rc == SOCKET_ERROR {
357                let e = last_err();
358                unsafe { closesocket(sock) };
359                return Err(e);
360            }
361            Ok(Self { sock })
362        }
363
364        pub fn send(&self, buf: &[u8]) -> io::Result<usize> {
365            let n = unsafe { send(self.sock, buf.as_ptr(), buf.len() as i32, 0) };
366            if n == SOCKET_ERROR {
367                Err(last_err())
368            } else {
369                Ok(n as usize)
370            }
371        }
372
373        pub fn recv(&self, buf: &mut [u8]) -> io::Result<usize> {
374            let n = unsafe { recv(self.sock, buf.as_mut_ptr(), buf.len() as i32, 0) };
375            if n == SOCKET_ERROR {
376                Err(last_err())
377            } else {
378                Ok(n as usize)
379            }
380        }
381    }
382
383    impl Drop for HostVmSocket {
384        fn drop(&mut self) {
385            unsafe { closesocket(self.sock) };
386        }
387    }
388}
389
390#[cfg(windows)]
391pub use hyperv_impl::HostVmSocket;
392
393#[cfg(test)]
394mod tests {
395    use super::*;
396    use std::thread;
397
398    /// In-process loopback round trip through the host-VM socket: a
399    /// listener accepts a loopback connection and echoes. Skips cleanly
400    /// when the loopback transport is unavailable (no `vsock_loopback`
401    /// module / no Hyper-V provider).
402    #[test]
403    fn host_vm_loopback_round_trip() {
404        // Per-process port so parallel test runs do not collide.
405        let port: u32 = 0x4000 + (std::process::id() & 0x3fff);
406
407        let listener = match HostVmSocket::listen_loopback(port) {
408            Ok(l) => l,
409            Err(e) => {
410                eprintln!("skipping: host-vm loopback unavailable ({e})");
411                return;
412            }
413        };
414
415        let client = thread::spawn(move || {
416            let c = match HostVmSocket::connect_loopback(port) {
417                Ok(c) => c,
418                Err(e) => {
419                    eprintln!("connect failed: {e}");
420                    return false;
421                }
422            };
423            if c.send(b"ping").is_err() {
424                return false;
425            }
426            let mut buf = [0u8; 4];
427            match c.recv(&mut buf) {
428                Ok(n) => &buf[..n] == b"pong",
429                Err(_) => false,
430            }
431        });
432
433        let conn = listener.accept().expect("accept");
434        let mut buf = [0u8; 4];
435        let n = conn.recv(&mut buf).expect("recv ping");
436        assert_eq!(&buf[..n], b"ping", "server received ping");
437        conn.send(b"pong").expect("send pong");
438
439        assert!(client.join().unwrap(), "client completed the round trip");
440    }
441}