use crate::device::VirtioNetworkDevice;
use crate::dns;
use crate::dns_relay::{self, DnsQuery, DnsResponse, DnsTransport};
use crate::egress::EgressPolicy;
use crate::icmp_relay;
use crate::queues::NetworkFrameQueues;
use crate::tcp_listeners::AcceptedTcpConnection;
use crate::tcp_relay::{spawn_tcp_relay, TcpRelayTable};
use crate::udp_relay;
use crate::virtio_net_log;
use smoltcp::iface::{
Config, Interface, PollIngressSingleResult, PollResult, SocketHandle, SocketSet,
};
use smoltcp::socket::raw::{
PacketBuffer as RawPacketBuffer, PacketMetadata as RawPacketMetadata, Socket as RawSocket,
};
use smoltcp::socket::tcp;
use smoltcp::socket::udp::{PacketBuffer, PacketMetadata, Socket as UdpSocket, UdpMetadata};
use smoltcp::time::Instant;
use smoltcp::wire::{
EthernetAddress, EthernetFrame, EthernetProtocol, HardwareAddress, IpAddress, IpCidr,
IpListenEndpoint, IpProtocol, IpVersion, Ipv4Packet, Ipv6Packet, TcpPacket, UdpPacket,
};
use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::atomic::Ordering;
use std::sync::mpsc::{Receiver, SyncSender, TryRecvError, TrySendError};
use std::sync::Arc;
use std::thread::{self, JoinHandle};
use std::time::{Duration, Instant as StdInstant};
const DNS_SOCKET_PORT: u16 = 53;
const DNS_PACKET_SLOTS: usize = 8;
const DNS_BUFFER_BYTES: usize = 2048;
const DNS_TCP_LISTENERS: usize = 4;
const DNS_TCP_RX_BYTES: usize = 4096;
const DNS_TCP_TX_BYTES: usize = 8192;
const DNS_TCP_MAX_MSG: usize = 4096;
const DEFAULT_IDLE_TIMEOUT_MS: i32 = 100;
const ICMP_PACKET_SLOTS: usize = 16;
const ICMP_BUFFER_BYTES: usize = 32 * 1024;
#[derive(Debug, Clone, Copy)]
pub struct VirtioPollConfig {
pub gateway_mac: [u8; 6],
pub guest_mac: [u8; 6],
pub gateway_ipv4: Ipv4Addr,
pub guest_ipv4: Ipv4Addr,
pub gateway_ipv6: Ipv6Addr,
pub guest_ipv6: Ipv6Addr,
pub prefix_len6: u8,
pub upstream_dns: Ipv4Addr,
pub mtu: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum FrameAction {
TcpSyn {
source: SocketAddr,
destination: SocketAddr,
},
DnsQuery,
UdpFlow {
destination: SocketAddr,
},
Passthrough,
}
pub fn start_network_stack(
queues: Arc<NetworkFrameQueues>,
config: VirtioPollConfig,
tcp_receiver: Option<Receiver<AcceptedTcpConnection>>,
egress: EgressPolicy,
) -> std::io::Result<JoinHandle<()>> {
virtio_net_log!(
"virtio-net: spawning poll thread guest_ip={} gateway_ip={} mtu={}",
config.guest_ipv4,
config.gateway_ipv4,
config.mtu
);
thread::Builder::new()
.name("smolvm-net-poll".into())
.spawn(move || run_network_stack(queues, config, tcp_receiver, egress))
}
fn run_network_stack(
queues: Arc<NetworkFrameQueues>,
config: VirtioPollConfig,
mut tcp_receiver: Option<Receiver<AcceptedTcpConnection>>,
egress: EgressPolicy,
) {
virtio_net_log!(
"virtio-net: poll loop started guest_ip={} gateway_ip={}",
config.guest_ipv4,
config.gateway_ipv4
);
let clock = StdInstant::now();
let mut device = VirtioNetworkDevice::new(queues.clone(), config.mtu);
let mut interface = create_interface(&mut device, &config);
let mut sockets = SocketSet::new(vec![]);
let dns_socket_handle = add_dns_socket(&mut sockets);
let dns_tcp_handles = add_dns_tcp_sockets(&mut sockets);
let mut dns_tcp_conns: Vec<DnsTcpConn> = (0..dns_tcp_handles.len())
.map(|_| DnsTcpConn::default())
.collect();
let (icmp4_handle, icmp6_handle) = add_icmp_raw_sockets(&mut sockets);
let gateway_addrs = [
IpAddr::V4(config.gateway_ipv4),
IpAddr::V6(config.gateway_ipv6),
IpAddr::V6(link_local_from_mac(config.gateway_mac)),
];
let relay_wake = Arc::new(queues.relay_wake.clone());
let mut relays = TcpRelayTable::new(None, egress.clone());
let mut udp_sockets = udp_relay::UdpSocketTable::new();
let udp_channels = {
let shutdown_queues = queues.clone();
udp_relay::start_udp_relay(
relay_wake.clone(),
Arc::new(move || shutdown_queues.is_shutting_down()),
)
};
let icmp_channels = {
let shutdown_queues = queues.clone();
icmp_relay::start_icmp_relay(
relay_wake.clone(),
Arc::new(move || shutdown_queues.is_shutting_down()),
)
};
let dns_channels = {
let shutdown_queues = queues.clone();
dns_relay::start_dns_relay(
relay_wake.clone(),
Arc::new(move || shutdown_queues.is_shutting_down()),
)
};
let mut dns_gateway = DnsGateway::new();
let poller = queues.guest_wake.poller().clone();
let mut events = polling::Events::new();
loop {
if queues.is_shutting_down() {
return;
}
let now = smoltcp_now(clock);
while let Some(frame) = device.stage_next_frame() {
match classify_guest_frame(frame, &gateway_addrs) {
FrameAction::TcpSyn {
source,
destination,
} => {
virtio_net_log!(
"virtio-net: guest TCP SYN source={} destination={}",
source,
destination
);
if !relays.has_socket_for(&source, &destination) {
relays.create_tcp_socket(source, destination, &mut sockets);
}
if matches!(
interface.poll_ingress_single(now, &mut device, &mut sockets),
PollIngressSingleResult::None
) {
device.drop_staged_frame();
}
}
FrameAction::DnsQuery | FrameAction::Passthrough => {
if matches!(
interface.poll_ingress_single(now, &mut device, &mut sockets),
PollIngressSingleResult::None
) {
device.drop_staged_frame();
}
}
FrameAction::UdpFlow { destination } => {
if udp_relay::should_relay_udp(destination, &egress)
&& udp_sockets.ensure_socket(destination, &mut sockets)
{
if matches!(
interface.poll_ingress_single(now, &mut device, &mut sockets),
PollIngressSingleResult::None
) {
device.drop_staged_frame();
}
} else {
device.drop_staged_frame();
}
}
}
}
relay_accepted_tcp_connection(
&mut tcp_receiver,
&mut relays,
&mut interface,
&mut sockets,
config.gateway_ipv4,
config.guest_ipv4,
);
flush_interface_egress(&mut interface, &mut device, &mut sockets, now);
interface.poll_maintenance(now);
wake_guest_if_needed(&queues, &device);
relays.relay_data(&mut sockets);
let mut woke_dns = false;
woke_dns |= dispatch_dns_udp(
dns_socket_handle,
&mut sockets,
&egress,
config.upstream_dns,
&mut dns_gateway,
&dns_channels.to_relay,
);
woke_dns |= process_dns_tcp(
&dns_tcp_handles,
&mut dns_tcp_conns,
&mut sockets,
&egress,
config.upstream_dns,
&mut dns_gateway,
&dns_channels.to_relay,
);
if woke_dns {
dns_channels.relay_thread_wake.wake();
}
deliver_dns_responses(
dns_socket_handle,
&dns_tcp_handles,
&mut dns_tcp_conns,
&mut sockets,
&egress,
&mut dns_gateway,
&dns_channels.from_relay,
);
if udp_sockets.drain_to_relay(&mut sockets, &udp_channels.to_relay) {
udp_channels.relay_thread_wake.wake();
}
udp_sockets.deliver_replies(&mut sockets, &udp_channels.from_relay);
udp_sockets.expire_idle(&mut sockets);
let mut woke_icmp = false;
woke_icmp |= drain_icmp_echo(
&mut sockets,
icmp4_handle,
false,
&egress,
&gateway_addrs,
&icmp_channels.to_relay,
);
woke_icmp |= drain_icmp_echo(
&mut sockets,
icmp6_handle,
true,
&egress,
&gateway_addrs,
&icmp_channels.to_relay,
);
if woke_icmp {
icmp_channels.relay_thread_wake.wake();
}
deliver_icmp_replies(
&mut sockets,
icmp4_handle,
icmp6_handle,
&icmp_channels.from_relay,
);
for connection in relays.take_new_connections(&mut sockets) {
spawn_tcp_relay(
connection.destination,
connection.relay_target,
connection.from_smoltcp,
connection.to_smoltcp,
relay_wake.clone(),
connection.exit_state,
);
}
relays.cleanup_closed(&mut sockets);
flush_interface_egress(&mut interface, &mut device, &mut sockets, now);
wake_guest_if_needed(&queues, &device);
let timeout = interface
.poll_delay(now, &sockets)
.map(|duration| Duration::from_millis(duration.total_millis().min(u32::MAX as u64)));
let timeout = match timeout {
Some(timeout) => Some(timeout),
None => Some(Duration::from_millis(DEFAULT_IDLE_TIMEOUT_MS as u64)),
};
events.clear();
let _ = poller.wait(&mut events, timeout);
}
}
fn create_interface(device: &mut VirtioNetworkDevice, config: &VirtioPollConfig) -> Interface {
let mut interface = Interface::new(
Config::new(HardwareAddress::Ethernet(EthernetAddress(
config.gateway_mac,
))),
device,
Instant::ZERO,
);
interface.update_ip_addrs(|addresses| {
addresses
.push(IpCidr::new(IpAddress::Ipv4(config.gateway_ipv4), 30))
.expect("failed to add gateway IPv4 address");
addresses
.push(IpCidr::new(
IpAddress::Ipv6(config.gateway_ipv6),
config.prefix_len6,
))
.expect("failed to add gateway IPv6 address");
addresses
.push(IpCidr::new(
IpAddress::Ipv6(link_local_from_mac(config.gateway_mac)),
64,
))
.expect("failed to add gateway IPv6 link-local address");
});
interface
.routes_mut()
.add_default_ipv4_route(config.gateway_ipv4)
.expect("failed to add default IPv4 route");
interface
.routes_mut()
.add_default_ipv6_route(config.gateway_ipv6)
.expect("failed to add default IPv6 route");
interface.set_any_ip(true);
interface
}
fn link_local_from_mac(mac: [u8; 6]) -> Ipv6Addr {
Ipv6Addr::new(
0xfe80,
0,
0,
0,
u16::from_be_bytes([mac[0] ^ 0x02, mac[1]]),
u16::from_be_bytes([mac[2], 0xff]),
u16::from_be_bytes([0xfe, mac[3]]),
u16::from_be_bytes([mac[4], mac[5]]),
)
}
fn add_dns_socket(sockets: &mut SocketSet<'_>) -> SocketHandle {
let rx_meta = vec![PacketMetadata::EMPTY; DNS_PACKET_SLOTS];
let tx_meta = vec![PacketMetadata::EMPTY; DNS_PACKET_SLOTS];
let rx_buffer = PacketBuffer::new(rx_meta, vec![0u8; DNS_BUFFER_BYTES]);
let tx_buffer = PacketBuffer::new(tx_meta, vec![0u8; DNS_BUFFER_BYTES]);
let mut socket = UdpSocket::new(rx_buffer, tx_buffer);
socket
.bind(smoltcp::wire::IpListenEndpoint {
addr: None,
port: DNS_SOCKET_PORT,
})
.expect("failed to bind gateway DNS socket");
sockets.add(socket)
}
fn add_icmp_raw_sockets(sockets: &mut SocketSet<'_>) -> (SocketHandle, SocketHandle) {
fn raw_socket(version: IpVersion, protocol: IpProtocol) -> RawSocket<'static> {
let rx = RawPacketBuffer::new(
vec![RawPacketMetadata::EMPTY; ICMP_PACKET_SLOTS],
vec![0u8; ICMP_BUFFER_BYTES],
);
let tx = RawPacketBuffer::new(
vec![RawPacketMetadata::EMPTY; ICMP_PACKET_SLOTS],
vec![0u8; ICMP_BUFFER_BYTES],
);
RawSocket::new(Some(version), Some(protocol), rx, tx)
}
let v4 = sockets.add(raw_socket(IpVersion::Ipv4, IpProtocol::Icmp));
let v6 = sockets.add(raw_socket(IpVersion::Ipv6, IpProtocol::Icmpv6));
(v4, v6)
}
fn drain_icmp_echo(
sockets: &mut SocketSet<'_>,
handle: SocketHandle,
is_ipv6: bool,
egress: &EgressPolicy,
gateway_addrs: &[IpAddr],
to_relay: &SyncSender<icmp_relay::IcmpEcho>,
) -> bool {
let mut echoes = Vec::new();
{
let socket = sockets.get_mut::<RawSocket>(handle);
while socket.can_recv() {
let Ok(packet) = socket.recv() else {
break;
};
let parsed = if is_ipv6 {
icmp_relay::parse_guest_echo_v6(packet)
} else {
icmp_relay::parse_guest_echo_v4(packet)
};
if let Some(echo) = parsed {
echoes.push(echo);
}
}
}
let mut woke = false;
let mut local_replies = Vec::new();
for echo in echoes {
if gateway_addrs.contains(&echo.destination) {
local_replies.push(echo);
} else if icmp_relay::should_relay_icmp(echo.destination, egress) {
match to_relay.try_send(echo) {
Ok(()) => woke = true,
Err(TrySendError::Full(_)) => {
virtio_net_log!("virtio-net: dropping guest ICMP echo (relay queue full)");
}
Err(TrySendError::Disconnected(_)) => return woke,
}
}
}
if !local_replies.is_empty() {
let socket = sockets.get_mut::<RawSocket>(handle);
for reply in local_replies {
let frame = if is_ipv6 {
icmp_relay::build_echo_reply_v6(&reply)
} else {
icmp_relay::build_echo_reply_v4(&reply)
};
if let Some(frame) = frame {
let _ = socket.send_slice(&frame);
}
}
}
woke
}
fn deliver_icmp_replies(
sockets: &mut SocketSet<'_>,
icmp4_handle: SocketHandle,
icmp6_handle: SocketHandle,
from_relay: &Receiver<icmp_relay::IcmpEcho>,
) {
while let Ok(reply) = from_relay.try_recv() {
let (handle, frame) = match reply.guest {
IpAddr::V4(_) => (icmp4_handle, icmp_relay::build_echo_reply_v4(&reply)),
IpAddr::V6(_) => (icmp6_handle, icmp_relay::build_echo_reply_v6(&reply)),
};
let Some(frame) = frame else {
continue;
};
let socket = sockets.get_mut::<RawSocket>(handle);
if socket.send_slice(&frame).is_err() {
virtio_net_log!(
"virtio-net: dropping ICMP reply to {} (raw socket buffer full)",
reply.guest
);
}
}
}
fn relay_accepted_tcp_connection(
tcp_receiver: &mut Option<Receiver<AcceptedTcpConnection>>,
relays: &mut TcpRelayTable,
interface: &mut Interface,
sockets: &mut SocketSet<'_>,
gateway_ipv4: Ipv4Addr,
guest_ipv4: Ipv4Addr,
) {
let mut disconnected = false;
if let Some(receiver) = tcp_receiver.as_mut() {
loop {
match receiver.try_recv() {
Ok(connection) => {
let guest_destination =
SocketAddr::new(std::net::IpAddr::V4(guest_ipv4), connection.guest_port);
virtio_net_log!(
"virtio-net: accepted published TCP connection peer={} host_port={} guest_destination={}",
connection.peer_addr,
connection.host_port,
guest_destination
);
if !relays.create_published_socket(
interface,
gateway_ipv4,
guest_destination,
connection.stream,
sockets,
) {
tracing::warn!(
host_port = connection.host_port,
guest_port = connection.guest_port,
peer_addr = %connection.peer_addr,
"dropping published TCP connection because the guest relay path could not be created"
);
}
}
Err(TryRecvError::Empty) => break,
Err(TryRecvError::Disconnected) => {
disconnected = true;
break;
}
}
}
}
if disconnected {
*tcp_receiver = None;
}
}
const MAX_PENDING_DNS: usize = 512;
struct DnsGateway {
next_id: u64,
pending_udp: HashMap<u64, PendingDnsUdp>,
}
struct PendingDnsUdp {
endpoint: smoltcp::wire::IpEndpoint,
local_address: Option<IpAddress>,
learn: bool,
}
impl DnsGateway {
fn new() -> Self {
Self {
next_id: 0,
pending_udp: HashMap::new(),
}
}
fn next_id(&mut self) -> u64 {
let id = self.next_id;
self.next_id = self.next_id.wrapping_add(1);
id
}
}
enum DnsDecision {
Immediate(Vec<u8>),
Forward { learn: bool },
}
fn classify_dns_query(query: &[u8], egress: &EgressPolicy) -> DnsDecision {
if !egress.dns_filter_active() {
return DnsDecision::Forward { learn: false };
}
match dns::question_name(query) {
Some(name) if egress.hostname_allowed(&name) => DnsDecision::Forward { learn: true },
Some(name) => {
virtio_net_log!(
"virtio-net: blocking DNS query by allow-host policy name={}",
name
);
DnsDecision::Immediate(dns::error_response(query, dns::DNS_RCODE_NXDOMAIN))
}
None => DnsDecision::Immediate(dns::error_response(query, dns::DNS_RCODE_SERVFAIL)),
}
}
fn dispatch_dns_udp(
dns_socket_handle: SocketHandle,
sockets: &mut SocketSet<'_>,
egress: &EgressPolicy,
upstream_dns: Ipv4Addr,
gateway: &mut DnsGateway,
to_relay: &SyncSender<DnsQuery>,
) -> bool {
let mut queued = false;
let socket = sockets.get_mut::<UdpSocket>(dns_socket_handle);
while socket.can_recv() {
let (query, endpoint, local_address) = match socket.recv() {
Ok((q, m)) => (q.to_vec(), m.endpoint, m.local_address),
Err(_) => break,
};
match classify_dns_query(&query, egress) {
DnsDecision::Immediate(response) => {
let response_meta = UdpMetadata {
endpoint,
local_address,
meta: Default::default(),
};
let _ = socket.send_slice(&response, response_meta);
}
DnsDecision::Forward { learn } => {
if gateway.pending_udp.len() >= MAX_PENDING_DNS {
virtio_net_log!("virtio-net: dropping guest DNS query (pending table full)");
continue;
}
let id = gateway.next_id();
gateway.pending_udp.insert(
id,
PendingDnsUdp {
endpoint,
local_address,
learn,
},
);
match to_relay.try_send(DnsQuery {
id,
transport: DnsTransport::Udp,
upstream: upstream_dns,
query,
}) {
Ok(()) => queued = true,
Err(TrySendError::Full(_)) => {
gateway.pending_udp.remove(&id);
virtio_net_log!("virtio-net: dropping guest DNS query (relay queue full)");
}
Err(TrySendError::Disconnected(_)) => {
gateway.pending_udp.remove(&id);
return queued;
}
}
}
}
}
queued
}
fn deliver_dns_responses(
dns_socket_handle: SocketHandle,
dns_tcp_handles: &[SocketHandle],
dns_tcp_conns: &mut [DnsTcpConn],
sockets: &mut SocketSet<'_>,
egress: &EgressPolicy,
gateway: &mut DnsGateway,
from_relay: &Receiver<DnsResponse>,
) {
while let Ok(response) = from_relay.try_recv() {
if let Some(pending) = gateway.pending_udp.remove(&response.id) {
if let Some(answer) = response.answer {
if pending.learn {
egress.learn_ip_records(&dns::answer_ip_records(&answer));
}
let socket = sockets.get_mut::<UdpSocket>(dns_socket_handle);
let response_meta = UdpMetadata {
endpoint: pending.endpoint,
local_address: pending.local_address,
meta: Default::default(),
};
let _ = socket.send_slice(&answer, response_meta);
}
continue;
}
for (handle, conn) in dns_tcp_handles.iter().zip(dns_tcp_conns.iter_mut()) {
match conn.awaiting {
Some(pending) if pending.id == response.id => {
conn.awaiting = None;
if let Some(answer) = response.answer {
if pending.learn {
egress.learn_ip_records(&dns::answer_ip_records(&answer));
}
frame_dns_tcp_response(conn, &answer);
}
conn.done = true;
let socket = sockets.get_mut::<tcp::Socket>(*handle);
drain_dns_tcp_tx(socket, conn);
break;
}
_ => {}
}
}
}
}
fn frame_dns_tcp_response(conn: &mut DnsTcpConn, response: &[u8]) {
if let Ok(resp_len) = u16::try_from(response.len()) {
conn.tx.extend_from_slice(&resp_len.to_be_bytes());
conn.tx.extend_from_slice(response);
}
}
fn add_dns_tcp_sockets(sockets: &mut SocketSet<'_>) -> Vec<SocketHandle> {
(0..DNS_TCP_LISTENERS)
.map(|_| {
let rx_buffer = tcp::SocketBuffer::new(vec![0u8; DNS_TCP_RX_BYTES]);
let tx_buffer = tcp::SocketBuffer::new(vec![0u8; DNS_TCP_TX_BYTES]);
let mut socket = tcp::Socket::new(rx_buffer, tx_buffer);
socket
.listen(IpListenEndpoint {
addr: None,
port: DNS_SOCKET_PORT,
})
.expect("failed to listen on gateway DNS TCP socket");
sockets.add(socket)
})
.collect()
}
#[derive(Default)]
struct DnsTcpConn {
rx: Vec<u8>,
tx: Vec<u8>,
tx_sent: usize,
done: bool,
awaiting: Option<DnsTcpPending>,
}
#[derive(Clone, Copy)]
struct DnsTcpPending {
id: u64,
learn: bool,
}
fn process_dns_tcp(
handles: &[SocketHandle],
conns: &mut [DnsTcpConn],
sockets: &mut SocketSet<'_>,
egress: &EgressPolicy,
upstream_dns: Ipv4Addr,
gateway: &mut DnsGateway,
to_relay: &SyncSender<DnsQuery>,
) -> bool {
let mut queued = false;
for (handle, conn) in handles.iter().zip(conns.iter_mut()) {
let socket = sockets.get_mut::<tcp::Socket>(*handle);
if !socket.is_open() {
if !conn.rx.is_empty() || !conn.tx.is_empty() || conn.done || conn.awaiting.is_some() {
*conn = DnsTcpConn::default();
}
let _ = socket.listen(IpListenEndpoint {
addr: None,
port: DNS_SOCKET_PORT,
});
continue;
}
if conn.awaiting.is_some() {
continue;
}
if conn.done {
drain_dns_tcp_tx(socket, conn);
continue;
}
while socket.can_recv() {
let appended = socket.recv(|data| (data.len(), data.to_vec()));
match appended {
Ok(bytes) if !bytes.is_empty() => conn.rx.extend_from_slice(&bytes),
_ => break,
}
}
if conn.rx.len() > DNS_TCP_MAX_MSG + 2 {
conn.done = true;
socket.close();
continue;
}
if conn.rx.len() >= 2 {
let msg_len = u16::from_be_bytes([conn.rx[0], conn.rx[1]]) as usize;
if msg_len == 0 || msg_len > DNS_TCP_MAX_MSG {
conn.done = true;
socket.close();
continue;
}
if conn.rx.len() >= 2 + msg_len {
let query = conn.rx[2..2 + msg_len].to_vec();
virtio_net_log!(
"virtio-net: DNS/TCP query query_len={} upstream_dns={}",
query.len(),
upstream_dns
);
match classify_dns_query(&query, egress) {
DnsDecision::Immediate(response) => {
frame_dns_tcp_response(conn, &response);
conn.done = true;
drain_dns_tcp_tx(socket, conn);
}
DnsDecision::Forward { learn } => {
let id = gateway.next_id();
match to_relay.try_send(DnsQuery {
id,
transport: DnsTransport::Tcp,
upstream: upstream_dns,
query,
}) {
Ok(()) => {
conn.awaiting = Some(DnsTcpPending { id, learn });
queued = true;
}
Err(_) => {
conn.done = true;
socket.close();
}
}
}
}
}
}
}
queued
}
fn drain_dns_tcp_tx(socket: &mut tcp::Socket<'_>, conn: &mut DnsTcpConn) {
while conn.tx_sent < conn.tx.len() && socket.can_send() {
match socket.send_slice(&conn.tx[conn.tx_sent..]) {
Ok(n) if n > 0 => conn.tx_sent += n,
_ => break,
}
}
if conn.tx_sent >= conn.tx.len() {
socket.close();
}
}
fn flush_interface_egress(
interface: &mut Interface,
device: &mut VirtioNetworkDevice,
sockets: &mut SocketSet<'_>,
now: Instant,
) {
loop {
let result = interface.poll_egress(now, device, sockets);
if matches!(result, PollResult::None) {
break;
}
}
}
fn wake_guest_if_needed(queues: &NetworkFrameQueues, device: &VirtioNetworkDevice) {
if device.frames_emitted.swap(false, Ordering::Relaxed) {
queues.host_wake.wake();
}
}
fn smoltcp_now(clock: StdInstant) -> Instant {
let elapsed = clock.elapsed();
Instant::from_millis(elapsed.as_millis() as i64)
}
fn classify_guest_frame(frame: &[u8], gateway_addrs: &[IpAddr]) -> FrameAction {
let ethernet = match EthernetFrame::new_checked(frame) {
Ok(frame) => frame,
Err(_) => return FrameAction::Passthrough,
};
let (src_ip, dst_ip, protocol, transport): (IpAddr, IpAddr, _, _) = match ethernet.ethertype() {
EthernetProtocol::Ipv4 => {
let ipv4 = match Ipv4Packet::new_checked(ethernet.payload()) {
Ok(packet) => packet,
Err(_) => return FrameAction::Passthrough,
};
(
IpAddr::V4(ipv4.src_addr()),
IpAddr::V4(ipv4.dst_addr()),
ipv4.next_header(),
ipv4.payload(),
)
}
EthernetProtocol::Ipv6 => {
let ipv6 = match Ipv6Packet::new_checked(ethernet.payload()) {
Ok(packet) => packet,
Err(_) => return FrameAction::Passthrough,
};
(
IpAddr::V6(ipv6.src_addr()),
IpAddr::V6(ipv6.dst_addr()),
ipv6.next_header(),
ipv6.payload(),
)
}
_ => return FrameAction::Passthrough,
};
match protocol {
smoltcp::wire::IpProtocol::Tcp => {
let tcp = match TcpPacket::new_checked(transport) {
Ok(packet) => packet,
Err(_) => return FrameAction::Passthrough,
};
if tcp.syn() && !tcp.ack() {
if tcp.dst_port() == DNS_SOCKET_PORT && gateway_addrs.contains(&dst_ip) {
FrameAction::Passthrough
} else {
FrameAction::TcpSyn {
source: SocketAddr::new(src_ip, tcp.src_port()),
destination: SocketAddr::new(dst_ip, tcp.dst_port()),
}
}
} else {
FrameAction::Passthrough
}
}
smoltcp::wire::IpProtocol::Udp => {
let udp = match UdpPacket::new_checked(transport) {
Ok(packet) => packet,
Err(_) => return FrameAction::Passthrough,
};
if udp.dst_port() == DNS_SOCKET_PORT {
FrameAction::DnsQuery
} else {
FrameAction::UdpFlow {
destination: SocketAddr::new(dst_ip, udp.dst_port()),
}
}
}
_ => FrameAction::Passthrough,
}
}
#[cfg(feature = "fuzzing")]
pub fn fuzz_classify_guest_frame(frame: &[u8]) {
let _ = classify_guest_frame(frame, &[]);
}
#[cfg(test)]
mod tests {
use super::*;
fn tcp_syn_frame(dst_ip: [u8; 4], dst_port: u16) -> Vec<u8> {
let mut f = Vec::new();
f.extend_from_slice(&[0xff; 6]);
f.extend_from_slice(&[0x02, 0, 0, 0, 0, 1]);
f.extend_from_slice(&[0x08, 0x00]);
f.extend_from_slice(&[0x45, 0x00, 0x00, 0x28, 0, 0, 0, 0, 0x40, 0x06, 0, 0]);
f.extend_from_slice(&[10, 0, 0, 2]); f.extend_from_slice(&dst_ip);
f.extend_from_slice(&54321u16.to_be_bytes());
f.extend_from_slice(&dst_port.to_be_bytes());
f.extend_from_slice(&[0, 0, 0, 0, 0, 0, 0, 0]); f.extend_from_slice(&[0x50, 0x02, 0xff, 0xff, 0, 0, 0, 0]); f
}
#[test]
fn dns_tcp_to_gateway_is_intercepted_not_relayed() {
let gw = IpAddr::V4(Ipv4Addr::new(100, 96, 0, 1));
assert_eq!(
classify_guest_frame(&tcp_syn_frame([100, 96, 0, 1], 53), &[gw]),
FrameAction::Passthrough
);
}
#[test]
fn dns_tcp_to_external_resolver_still_relayed() {
let gw = IpAddr::V4(Ipv4Addr::new(100, 96, 0, 1));
assert!(matches!(
classify_guest_frame(&tcp_syn_frame([1, 1, 1, 1], 53), &[gw]),
FrameAction::TcpSyn { .. }
));
}
#[test]
fn non_dns_tcp_to_gateway_still_relayed() {
let gw = IpAddr::V4(Ipv4Addr::new(100, 96, 0, 1));
assert!(matches!(
classify_guest_frame(&tcp_syn_frame([100, 96, 0, 1], 443), &[gw]),
FrameAction::TcpSyn { .. }
));
}
}