use crate::egress::EgressPolicy;
use crate::queues::WakePipe;
use crate::virtio_net_log;
use smoltcp::iface::{SocketHandle, SocketSet};
use smoltcp::socket::udp::{PacketBuffer, PacketMetadata, Socket as UdpSocket, UdpMetadata};
use std::collections::HashMap;
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr, UdpSocket as HostUdpSocket};
use std::os::fd::AsRawFd;
use std::sync::mpsc::{self, Receiver, SyncSender, TryRecvError, TrySendError};
use std::sync::Arc;
use std::thread;
use std::time::{Duration, Instant};
const CHANNEL_CAPACITY: usize = 256;
const MAX_FLOWS: usize = 1024;
const MAX_DST_SOCKETS: usize = 128;
const UDP_PACKET_SLOTS: usize = 16;
const UDP_BUFFER_BYTES: usize = 64 * 1024;
const MAX_DATAGRAM_BYTES: usize = 65_535;
const FLOW_IDLE_TIMEOUT: Duration = Duration::from_secs(60);
const DST_SOCKET_IDLE_TIMEOUT: Duration = Duration::from_secs(120);
const RELAY_POLL_MAX_MS: i32 = 1000;
pub struct UdpDatagram {
pub guest: SocketAddr,
pub destination: SocketAddr,
pub payload: Vec<u8>,
}
pub struct UdpRelayChannels {
pub to_relay: SyncSender<UdpDatagram>,
pub from_relay: Receiver<UdpDatagram>,
pub relay_thread_wake: WakePipe,
}
pub fn start_udp_relay(
reply_wake: Arc<WakePipe>,
shutdown: Arc<dyn Fn() -> bool + Send + Sync>,
) -> UdpRelayChannels {
let (to_relay_tx, to_relay_rx) = mpsc::sync_channel(CHANNEL_CAPACITY);
let (from_relay_tx, from_relay_rx) = mpsc::sync_channel(CHANNEL_CAPACITY);
let relay_thread_wake = WakePipe::new();
let thread_wake = relay_thread_wake.clone();
let _ = thread::Builder::new()
.name("smolvm-udp-relay".into())
.spawn(move || {
run_udp_relay(
to_relay_rx,
from_relay_tx,
thread_wake,
reply_wake,
shutdown,
);
});
UdpRelayChannels {
to_relay: to_relay_tx,
from_relay: from_relay_rx,
relay_thread_wake,
}
}
struct UdpFlow {
socket: HostUdpSocket,
guest: SocketAddr,
destination: SocketAddr,
last_active: Instant,
}
fn run_udp_relay(
outbound: Receiver<UdpDatagram>,
inbound: SyncSender<UdpDatagram>,
wake: WakePipe,
reply_wake: Arc<WakePipe>,
shutdown: Arc<dyn Fn() -> bool + Send + Sync>,
) {
let mut flows: HashMap<(SocketAddr, SocketAddr), UdpFlow> = HashMap::new();
let mut recv_buf = vec![0u8; MAX_DATAGRAM_BYTES];
loop {
if shutdown() {
return;
}
loop {
match outbound.try_recv() {
Ok(datagram) => {
let key = (datagram.guest, datagram.destination);
if !flows.contains_key(&key) {
if flows.len() >= MAX_FLOWS {
virtio_net_log!(
"virtio-net: dropping UDP flow {} -> {} (flow table full)",
datagram.guest,
datagram.destination
);
continue;
}
match create_flow_socket(datagram.destination) {
Ok(socket) => {
flows.insert(
key,
UdpFlow {
socket,
guest: datagram.guest,
destination: datagram.destination,
last_active: Instant::now(),
},
);
}
Err(err) => {
virtio_net_log!(
"virtio-net: failed to open host UDP socket for {} -> {}: {}",
datagram.guest,
datagram.destination,
err
);
continue;
}
}
}
let flow = flows.get_mut(&key).expect("flow inserted above");
flow.last_active = Instant::now();
let _ = flow.socket.send(&datagram.payload);
}
Err(TryRecvError::Empty) => break,
Err(TryRecvError::Disconnected) => return,
}
}
let mut poll_fds: Vec<libc::pollfd> = Vec::with_capacity(flows.len() + 1);
poll_fds.push(libc::pollfd {
fd: wake.as_raw_fd(),
events: libc::POLLIN,
revents: 0,
});
let keys: Vec<(SocketAddr, SocketAddr)> = flows.keys().copied().collect();
for key in &keys {
poll_fds.push(libc::pollfd {
fd: flows[key].socket.as_raw_fd(),
events: libc::POLLIN,
revents: 0,
});
}
unsafe {
libc::poll(
poll_fds.as_mut_ptr(),
poll_fds.len() as libc::nfds_t,
RELAY_POLL_MAX_MS,
);
}
if poll_fds[0].revents & libc::POLLIN != 0 {
wake.drain();
}
let mut woke_reply = false;
for (slot, key) in keys.iter().enumerate() {
if poll_fds[slot + 1].revents & libc::POLLIN == 0 {
continue;
}
let Some(flow) = flows.get_mut(key) else {
continue;
};
while let Ok(len) = flow.socket.recv(&mut recv_buf) {
flow.last_active = Instant::now();
let reply = UdpDatagram {
guest: flow.guest,
destination: flow.destination,
payload: recv_buf[..len].to_vec(),
};
match inbound.try_send(reply) {
Ok(()) => woke_reply = true,
Err(TrySendError::Full(_)) => {
virtio_net_log!(
"virtio-net: dropping UDP reply for {} (inbound queue full)",
flow.guest
);
}
Err(TrySendError::Disconnected(_)) => return,
}
}
}
if woke_reply {
reply_wake.wake();
}
let now = Instant::now();
flows.retain(|_, flow| now.duration_since(flow.last_active) < FLOW_IDLE_TIMEOUT);
}
}
fn create_flow_socket(destination: SocketAddr) -> std::io::Result<HostUdpSocket> {
let bind_addr: SocketAddr = if destination.is_ipv4() {
(Ipv4Addr::UNSPECIFIED, 0).into()
} else {
(Ipv6Addr::UNSPECIFIED, 0).into()
};
let socket = HostUdpSocket::bind(bind_addr)?;
socket.connect(destination)?;
socket.set_nonblocking(true)?;
Ok(socket)
}
pub struct UdpSocketTable {
sockets: HashMap<SocketAddr, SocketHandle>,
last_active: HashMap<SocketAddr, Instant>,
}
impl UdpSocketTable {
pub fn new() -> Self {
Self {
sockets: HashMap::new(),
last_active: HashMap::new(),
}
}
pub fn ensure_socket(&mut self, destination: SocketAddr, sockets: &mut SocketSet<'_>) -> bool {
if self.sockets.contains_key(&destination) {
self.last_active.insert(destination, Instant::now());
return true;
}
if self.sockets.len() >= MAX_DST_SOCKETS {
virtio_net_log!(
"virtio-net: dropping UDP datagram to {} (destination socket table full)",
destination
);
return false;
}
let rx = PacketBuffer::new(
vec![PacketMetadata::EMPTY; UDP_PACKET_SLOTS],
vec![0u8; UDP_BUFFER_BYTES],
);
let tx = PacketBuffer::new(
vec![PacketMetadata::EMPTY; UDP_PACKET_SLOTS],
vec![0u8; UDP_BUFFER_BYTES],
);
let mut socket = UdpSocket::new(rx, tx);
if socket
.bind(smoltcp::wire::IpListenEndpoint {
addr: Some(destination.ip().into()),
port: destination.port(),
})
.is_err()
{
return false;
}
let handle = sockets.add(socket);
self.sockets.insert(destination, handle);
self.last_active.insert(destination, Instant::now());
true
}
pub fn drain_to_relay(
&mut self,
sockets: &mut SocketSet<'_>,
to_relay: &SyncSender<UdpDatagram>,
) -> bool {
let mut forwarded = false;
for (&destination, &handle) in &self.sockets {
let socket = sockets.get_mut::<UdpSocket>(handle);
while socket.can_recv() {
let Ok((payload, metadata)) = socket.recv() else {
break;
};
self.last_active.insert(destination, Instant::now());
let guest = endpoint_to_socket_addr(metadata.endpoint);
let datagram = UdpDatagram {
guest,
destination,
payload: payload.to_vec(),
};
match to_relay.try_send(datagram) {
Ok(()) => forwarded = true,
Err(TrySendError::Full(_)) => {
virtio_net_log!(
"virtio-net: dropping guest UDP datagram to {} (relay queue full)",
destination
);
}
Err(TrySendError::Disconnected(_)) => return forwarded,
}
}
}
forwarded
}
pub fn deliver_replies(
&mut self,
sockets: &mut SocketSet<'_>,
from_relay: &Receiver<UdpDatagram>,
) {
while let Ok(reply) = from_relay.try_recv() {
let Some(&handle) = self.sockets.get(&reply.destination) else {
continue;
};
self.last_active.insert(reply.destination, Instant::now());
let socket = sockets.get_mut::<UdpSocket>(handle);
let metadata = UdpMetadata {
endpoint: smoltcp::wire::IpEndpoint {
addr: reply.guest.ip().into(),
port: reply.guest.port(),
},
local_address: Some(reply.destination.ip().into()),
meta: Default::default(),
};
if socket.send_slice(&reply.payload, metadata).is_err() {
virtio_net_log!(
"virtio-net: dropping UDP reply to {} (guest socket buffer full)",
reply.guest
);
}
}
}
pub fn expire_idle(&mut self, sockets: &mut SocketSet<'_>) {
let now = Instant::now();
let expired: Vec<SocketAddr> = self
.last_active
.iter()
.filter(|(_, last)| now.duration_since(**last) >= DST_SOCKET_IDLE_TIMEOUT)
.map(|(dst, _)| *dst)
.collect();
for destination in expired {
if let Some(handle) = self.sockets.remove(&destination) {
sockets.remove(handle);
}
self.last_active.remove(&destination);
}
}
}
impl Default for UdpSocketTable {
fn default() -> Self {
Self::new()
}
}
pub fn should_relay_udp(destination: SocketAddr, egress: &EgressPolicy) -> bool {
destination.port() != 53 && egress.allows(destination.ip())
}
fn endpoint_to_socket_addr(endpoint: smoltcp::wire::IpEndpoint) -> SocketAddr {
let ip: std::net::IpAddr = endpoint.addr.into();
SocketAddr::new(ip, endpoint.port)
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicBool, Ordering};
#[test]
fn flow_socket_round_trip() {
let server = HostUdpSocket::bind("127.0.0.1:0").unwrap();
let server_addr = server.local_addr().unwrap();
let flow = create_flow_socket(server_addr).unwrap();
flow.send(b"ping").unwrap();
let mut buf = [0u8; 16];
let (len, peer) = server.recv_from(&mut buf).unwrap();
assert_eq!(&buf[..len], b"ping");
server.send_to(b"pong", peer).unwrap();
let deadline = Instant::now() + Duration::from_secs(2);
loop {
match flow.recv(&mut buf) {
Ok(len) => {
assert_eq!(&buf[..len], b"pong");
break;
}
Err(_) if Instant::now() < deadline => thread::sleep(Duration::from_millis(10)),
Err(e) => panic!("no reply: {e}"),
}
}
}
#[test]
fn relay_thread_round_trips_a_datagram() {
let server = HostUdpSocket::bind("127.0.0.1:0").unwrap();
let server_addr = server.local_addr().unwrap();
let reply_wake = Arc::new(WakePipe::new());
let stop = Arc::new(AtomicBool::new(false));
let stop_flag = stop.clone();
let channels = start_udp_relay(
reply_wake.clone(),
Arc::new(move || stop_flag.load(Ordering::Relaxed)),
);
let guest: SocketAddr = "100.96.0.2:40000".parse().unwrap();
channels
.to_relay
.send(UdpDatagram {
guest,
destination: server_addr,
payload: b"hello".to_vec(),
})
.unwrap();
channels.relay_thread_wake.wake();
let mut buf = [0u8; 16];
server
.set_read_timeout(Some(Duration::from_secs(2)))
.unwrap();
let (len, peer) = server.recv_from(&mut buf).unwrap();
assert_eq!(&buf[..len], b"hello");
server.send_to(b"world", peer).unwrap();
let reply = channels
.from_relay
.recv_timeout(Duration::from_secs(2))
.unwrap();
assert_eq!(reply.guest, guest);
assert_eq!(reply.destination, server_addr);
assert_eq!(reply.payload, b"world");
stop.store(true, Ordering::Relaxed);
channels.relay_thread_wake.wake();
}
#[test]
fn should_relay_respects_dns_carveout_and_policy() {
let open = EgressPolicy::unrestricted();
assert!(should_relay_udp("1.2.3.4:123".parse().unwrap(), &open));
assert!(!should_relay_udp("1.2.3.4:53".parse().unwrap(), &open));
let restricted = EgressPolicy::from_allowed_cidrs(Some(&["8.8.8.0/24".into()]));
assert!(should_relay_udp(
"8.8.8.3:123".parse().unwrap(),
&restricted
));
assert!(!should_relay_udp(
"1.2.3.4:123".parse().unwrap(),
&restricted
));
}
}