1#![cfg(any(target_os = "linux", windows))]
28
29#[cfg(target_os = "linux")]
34mod vsock_impl {
35 use std::io;
36 use std::os::unix::io::{AsRawFd, FromRawFd, RawFd};
37
38 pub const VMADDR_CID_ANY: u32 = 0xFFFFFFFF;
40 pub const VMADDR_CID_HOST: u32 = 2;
41 pub const VMADDR_CID_LOCAL: u32 = 1;
44
45 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 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 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 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 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 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 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 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 pub struct HostVmSocket {
172 inner: VsockSocket,
173 }
174
175 impl HostVmSocket {
176 pub fn listen_loopback(port: u32) -> io::Result<Self> {
178 let s = VsockSocket::new()?;
179 s.bind(VMADDR_CID_ANY, port)?;
181 s.listen(16)?;
182 Ok(Self { inner: s })
183 }
184
185 pub fn accept(&self) -> io::Result<Self> {
187 Ok(Self { inner: self.inner.accept()? })
188 }
189
190 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#[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 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 const HV_GUID_WILDCARD: GUID = GUID {
237 data1: 0,
238 data2: 0,
239 data3: 0,
240 data4: [0; 8],
241 };
242
243 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 unsafe { WSAStartup(0x0202, &mut data) };
270 });
271 }
272
273 fn last_err() -> io::Error {
274 io::Error::last_os_error()
275 }
276
277 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 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 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 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 #[test]
403 fn host_vm_loopback_round_trip() {
404 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}