1use crate::device::VirtioNetworkDevice;
54use crate::dns;
55use crate::egress::EgressPolicy;
56use crate::icmp_relay;
57use crate::queues::NetworkFrameQueues;
58use crate::tcp_listeners::AcceptedTcpConnection;
59use crate::tcp_relay::{spawn_tcp_relay, TcpRelayTable};
60use crate::udp_relay;
61use crate::virtio_net_log;
62use smoltcp::iface::{
63 Config, Interface, PollIngressSingleResult, PollResult, SocketHandle, SocketSet,
64};
65use smoltcp::socket::raw::{
66 PacketBuffer as RawPacketBuffer, PacketMetadata as RawPacketMetadata, Socket as RawSocket,
67};
68use smoltcp::socket::tcp;
69use smoltcp::socket::udp::{PacketBuffer, PacketMetadata, Socket as UdpSocket, UdpMetadata};
70use smoltcp::time::Instant;
71use smoltcp::wire::{
72 EthernetAddress, EthernetFrame, EthernetProtocol, HardwareAddress, IpAddress, IpCidr,
73 IpListenEndpoint, IpProtocol, IpVersion, Ipv4Packet, Ipv6Packet, TcpPacket, UdpPacket,
74};
75use std::io::{Read, Write};
76use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, UdpSocket as HostUdpSocket};
77use std::sync::atomic::Ordering;
78use std::sync::mpsc::{Receiver, SyncSender, TryRecvError, TrySendError};
79use std::sync::Arc;
80use std::thread::{self, JoinHandle};
81use std::time::{Duration, Instant as StdInstant};
82
83const DNS_SOCKET_PORT: u16 = 53;
84const DNS_PACKET_SLOTS: usize = 8;
85const DNS_BUFFER_BYTES: usize = 2048;
86const DNS_TCP_LISTENERS: usize = 4;
92const DNS_TCP_RX_BYTES: usize = 4096;
93const DNS_TCP_TX_BYTES: usize = 8192;
94const DNS_TCP_MAX_MSG: usize = 4096;
98const DEFAULT_IDLE_TIMEOUT_MS: i32 = 100;
99const ICMP_PACKET_SLOTS: usize = 16;
101const ICMP_BUFFER_BYTES: usize = 32 * 1024;
103
104#[derive(Debug, Clone, Copy)]
110pub struct VirtioPollConfig {
111 pub gateway_mac: [u8; 6],
113 pub guest_mac: [u8; 6],
115 pub gateway_ipv4: Ipv4Addr,
117 pub guest_ipv4: Ipv4Addr,
119 pub gateway_ipv6: Ipv6Addr,
121 pub guest_ipv6: Ipv6Addr,
123 pub prefix_len6: u8,
125 pub upstream_dns: Ipv4Addr,
127 pub mtu: usize,
129}
130
131#[derive(Debug, Clone, Copy, PartialEq, Eq)]
132enum FrameAction {
133 TcpSyn {
134 source: SocketAddr,
135 destination: SocketAddr,
136 },
137 DnsQuery,
138 UdpFlow {
140 destination: SocketAddr,
141 },
142 Passthrough,
143}
144
145pub fn start_network_stack(
155 queues: Arc<NetworkFrameQueues>,
156 config: VirtioPollConfig,
157 tcp_receiver: Option<Receiver<AcceptedTcpConnection>>,
158 egress: EgressPolicy,
159) -> std::io::Result<JoinHandle<()>> {
160 virtio_net_log!(
161 "virtio-net: spawning poll thread guest_ip={} gateway_ip={} mtu={}",
162 config.guest_ipv4,
163 config.gateway_ipv4,
164 config.mtu
165 );
166 thread::Builder::new()
167 .name("smolvm-net-poll".into())
168 .spawn(move || run_network_stack(queues, config, tcp_receiver, egress))
169}
170
171fn run_network_stack(
172 queues: Arc<NetworkFrameQueues>,
173 config: VirtioPollConfig,
174 mut tcp_receiver: Option<Receiver<AcceptedTcpConnection>>,
175 egress: EgressPolicy,
176) {
177 virtio_net_log!(
191 "virtio-net: poll loop started guest_ip={} gateway_ip={}",
192 config.guest_ipv4,
193 config.gateway_ipv4
194 );
195 let clock = StdInstant::now();
196 let mut device = VirtioNetworkDevice::new(queues.clone(), config.mtu);
197 let mut interface = create_interface(&mut device, &config);
198 let mut sockets = SocketSet::new(vec![]);
199 let dns_socket_handle = add_dns_socket(&mut sockets);
200 let dns_tcp_handles = add_dns_tcp_sockets(&mut sockets);
201 let mut dns_tcp_conns: Vec<DnsTcpConn> = (0..dns_tcp_handles.len())
202 .map(|_| DnsTcpConn::default())
203 .collect();
204 let (icmp4_handle, icmp6_handle) = add_icmp_raw_sockets(&mut sockets);
205 let gateway_addrs = [
208 IpAddr::V4(config.gateway_ipv4),
209 IpAddr::V6(config.gateway_ipv6),
210 IpAddr::V6(link_local_from_mac(config.gateway_mac)),
211 ];
212 let relay_wake = Arc::new(queues.relay_wake.clone());
213 let mut relays = TcpRelayTable::new(None, egress.clone());
214 let mut udp_sockets = udp_relay::UdpSocketTable::new();
215 let udp_channels = {
216 let shutdown_queues = queues.clone();
217 udp_relay::start_udp_relay(
218 relay_wake.clone(),
219 Arc::new(move || shutdown_queues.is_shutting_down()),
220 )
221 };
222 let icmp_channels = {
223 let shutdown_queues = queues.clone();
224 icmp_relay::start_icmp_relay(
225 relay_wake.clone(),
226 Arc::new(move || shutdown_queues.is_shutting_down()),
227 )
228 };
229
230 let poller = queues.guest_wake.poller().clone();
239 let mut events = polling::Events::new();
240
241 loop {
242 if queues.is_shutting_down() {
243 return;
244 }
245 let now = smoltcp_now(clock);
246
247 while let Some(frame) = device.stage_next_frame() {
248 match classify_guest_frame(frame, &gateway_addrs) {
254 FrameAction::TcpSyn {
255 source,
256 destination,
257 } => {
258 virtio_net_log!(
259 "virtio-net: guest TCP SYN source={} destination={}",
260 source,
261 destination
262 );
263 if !relays.has_socket_for(&source, &destination) {
264 relays.create_tcp_socket(source, destination, &mut sockets);
265 }
266 if matches!(
267 interface.poll_ingress_single(now, &mut device, &mut sockets),
268 PollIngressSingleResult::None
269 ) {
270 device.drop_staged_frame();
271 }
272 }
273 FrameAction::DnsQuery | FrameAction::Passthrough => {
274 if matches!(
275 interface.poll_ingress_single(now, &mut device, &mut sockets),
276 PollIngressSingleResult::None
277 ) {
278 device.drop_staged_frame();
279 }
280 }
281 FrameAction::UdpFlow { destination } => {
282 if udp_relay::should_relay_udp(destination, &egress)
285 && udp_sockets.ensure_socket(destination, &mut sockets)
286 {
287 if matches!(
288 interface.poll_ingress_single(now, &mut device, &mut sockets),
289 PollIngressSingleResult::None
290 ) {
291 device.drop_staged_frame();
292 }
293 } else {
294 device.drop_staged_frame();
295 }
296 }
297 }
298 }
299
300 relay_accepted_tcp_connection(
301 &mut tcp_receiver,
302 &mut relays,
303 &mut interface,
304 &mut sockets,
305 config.gateway_ipv4,
306 config.guest_ipv4,
307 );
308
309 flush_interface_egress(&mut interface, &mut device, &mut sockets, now);
312 interface.poll_maintenance(now);
313 wake_guest_if_needed(&queues, &device);
314
315 relays.relay_data(&mut sockets);
318 process_dns_queries(
319 dns_socket_handle,
320 &mut sockets,
321 &egress,
322 config.upstream_dns,
323 );
324 process_dns_tcp(
325 &dns_tcp_handles,
326 &mut dns_tcp_conns,
327 &mut sockets,
328 &egress,
329 config.upstream_dns,
330 );
331
332 if udp_sockets.drain_to_relay(&mut sockets, &udp_channels.to_relay) {
335 udp_channels.relay_thread_wake.wake();
336 }
337 udp_sockets.deliver_replies(&mut sockets, &udp_channels.from_relay);
338 udp_sockets.expire_idle(&mut sockets);
339
340 let mut woke_icmp = false;
344 woke_icmp |= drain_icmp_echo(
345 &mut sockets,
346 icmp4_handle,
347 false,
348 &egress,
349 &gateway_addrs,
350 &icmp_channels.to_relay,
351 );
352 woke_icmp |= drain_icmp_echo(
353 &mut sockets,
354 icmp6_handle,
355 true,
356 &egress,
357 &gateway_addrs,
358 &icmp_channels.to_relay,
359 );
360 if woke_icmp {
361 icmp_channels.relay_thread_wake.wake();
362 }
363 deliver_icmp_replies(
364 &mut sockets,
365 icmp4_handle,
366 icmp6_handle,
367 &icmp_channels.from_relay,
368 );
369
370 for connection in relays.take_new_connections(&mut sockets) {
373 spawn_tcp_relay(
374 connection.destination,
375 connection.relay_target,
376 connection.from_smoltcp,
377 connection.to_smoltcp,
378 relay_wake.clone(),
379 connection.exit_state,
380 );
381 }
382
383 relays.cleanup_closed(&mut sockets);
384
385 flush_interface_egress(&mut interface, &mut device, &mut sockets, now);
388 wake_guest_if_needed(&queues, &device);
389
390 let timeout = interface
391 .poll_delay(now, &sockets)
392 .map(|duration| Duration::from_millis(duration.total_millis().min(u32::MAX as u64)));
393 let timeout = match timeout {
394 Some(timeout) => Some(timeout),
395 None => Some(Duration::from_millis(DEFAULT_IDLE_TIMEOUT_MS as u64)),
396 };
397
398 events.clear();
402 let _ = poller.wait(&mut events, timeout);
403 }
404}
405
406fn create_interface(device: &mut VirtioNetworkDevice, config: &VirtioPollConfig) -> Interface {
407 let mut interface = Interface::new(
417 Config::new(HardwareAddress::Ethernet(EthernetAddress(
418 config.gateway_mac,
419 ))),
420 device,
421 Instant::ZERO,
422 );
423 interface.update_ip_addrs(|addresses| {
424 addresses
425 .push(IpCidr::new(IpAddress::Ipv4(config.gateway_ipv4), 30))
426 .expect("failed to add gateway IPv4 address");
427 addresses
428 .push(IpCidr::new(
429 IpAddress::Ipv6(config.gateway_ipv6),
430 config.prefix_len6,
431 ))
432 .expect("failed to add gateway IPv6 address");
433 addresses
437 .push(IpCidr::new(
438 IpAddress::Ipv6(link_local_from_mac(config.gateway_mac)),
439 64,
440 ))
441 .expect("failed to add gateway IPv6 link-local address");
442 });
443 interface
447 .routes_mut()
448 .add_default_ipv4_route(config.gateway_ipv4)
449 .expect("failed to add default IPv4 route");
450 interface
451 .routes_mut()
452 .add_default_ipv6_route(config.gateway_ipv6)
453 .expect("failed to add default IPv6 route");
454 interface.set_any_ip(true);
455 interface
456}
457
458fn link_local_from_mac(mac: [u8; 6]) -> Ipv6Addr {
461 Ipv6Addr::new(
462 0xfe80,
463 0,
464 0,
465 0,
466 u16::from_be_bytes([mac[0] ^ 0x02, mac[1]]),
467 u16::from_be_bytes([mac[2], 0xff]),
468 u16::from_be_bytes([0xfe, mac[3]]),
469 u16::from_be_bytes([mac[4], mac[5]]),
470 )
471}
472
473fn add_dns_socket(sockets: &mut SocketSet<'_>) -> SocketHandle {
484 let rx_meta = vec![PacketMetadata::EMPTY; DNS_PACKET_SLOTS];
485 let tx_meta = vec![PacketMetadata::EMPTY; DNS_PACKET_SLOTS];
486 let rx_buffer = PacketBuffer::new(rx_meta, vec![0u8; DNS_BUFFER_BYTES]);
487 let tx_buffer = PacketBuffer::new(tx_meta, vec![0u8; DNS_BUFFER_BYTES]);
488 let mut socket = UdpSocket::new(rx_buffer, tx_buffer);
489 socket
490 .bind(smoltcp::wire::IpListenEndpoint {
491 addr: None,
492 port: DNS_SOCKET_PORT,
493 })
494 .expect("failed to bind gateway DNS socket");
495 sockets.add(socket)
496}
497
498fn add_icmp_raw_sockets(sockets: &mut SocketSet<'_>) -> (SocketHandle, SocketHandle) {
506 fn raw_socket(version: IpVersion, protocol: IpProtocol) -> RawSocket<'static> {
507 let rx = RawPacketBuffer::new(
508 vec![RawPacketMetadata::EMPTY; ICMP_PACKET_SLOTS],
509 vec![0u8; ICMP_BUFFER_BYTES],
510 );
511 let tx = RawPacketBuffer::new(
512 vec![RawPacketMetadata::EMPTY; ICMP_PACKET_SLOTS],
513 vec![0u8; ICMP_BUFFER_BYTES],
514 );
515 RawSocket::new(Some(version), Some(protocol), rx, tx)
516 }
517
518 let v4 = sockets.add(raw_socket(IpVersion::Ipv4, IpProtocol::Icmp));
519 let v6 = sockets.add(raw_socket(IpVersion::Ipv6, IpProtocol::Icmpv6));
520 (v4, v6)
521}
522
523fn drain_icmp_echo(
528 sockets: &mut SocketSet<'_>,
529 handle: SocketHandle,
530 is_ipv6: bool,
531 egress: &EgressPolicy,
532 gateway_addrs: &[IpAddr],
533 to_relay: &SyncSender<icmp_relay::IcmpEcho>,
534) -> bool {
535 let mut echoes = Vec::new();
538 {
539 let socket = sockets.get_mut::<RawSocket>(handle);
540 while socket.can_recv() {
541 let Ok(packet) = socket.recv() else {
542 break;
543 };
544 let parsed = if is_ipv6 {
545 icmp_relay::parse_guest_echo_v6(packet)
546 } else {
547 icmp_relay::parse_guest_echo_v4(packet)
548 };
549 if let Some(echo) = parsed {
550 echoes.push(echo);
551 }
552 }
553 }
554
555 let mut woke = false;
557 let mut local_replies = Vec::new();
558 for echo in echoes {
559 if gateway_addrs.contains(&echo.destination) {
560 local_replies.push(echo);
561 } else if icmp_relay::should_relay_icmp(echo.destination, egress) {
562 match to_relay.try_send(echo) {
563 Ok(()) => woke = true,
564 Err(TrySendError::Full(_)) => {
565 virtio_net_log!("virtio-net: dropping guest ICMP echo (relay queue full)");
566 }
567 Err(TrySendError::Disconnected(_)) => return woke,
568 }
569 }
570 }
572
573 if !local_replies.is_empty() {
575 let socket = sockets.get_mut::<RawSocket>(handle);
576 for reply in local_replies {
577 let frame = if is_ipv6 {
578 icmp_relay::build_echo_reply_v6(&reply)
579 } else {
580 icmp_relay::build_echo_reply_v4(&reply)
581 };
582 if let Some(frame) = frame {
583 let _ = socket.send_slice(&frame);
584 }
585 }
586 }
587 woke
588}
589
590fn deliver_icmp_replies(
594 sockets: &mut SocketSet<'_>,
595 icmp4_handle: SocketHandle,
596 icmp6_handle: SocketHandle,
597 from_relay: &Receiver<icmp_relay::IcmpEcho>,
598) {
599 while let Ok(reply) = from_relay.try_recv() {
600 let (handle, frame) = match reply.guest {
601 IpAddr::V4(_) => (icmp4_handle, icmp_relay::build_echo_reply_v4(&reply)),
602 IpAddr::V6(_) => (icmp6_handle, icmp_relay::build_echo_reply_v6(&reply)),
603 };
604 let Some(frame) = frame else {
605 continue;
606 };
607 let socket = sockets.get_mut::<RawSocket>(handle);
608 if socket.send_slice(&frame).is_err() {
609 virtio_net_log!(
610 "virtio-net: dropping ICMP reply to {} (raw socket buffer full)",
611 reply.guest
612 );
613 }
614 }
615}
616
617fn relay_accepted_tcp_connection(
620 tcp_receiver: &mut Option<Receiver<AcceptedTcpConnection>>,
621 relays: &mut TcpRelayTable,
622 interface: &mut Interface,
623 sockets: &mut SocketSet<'_>,
624 gateway_ipv4: Ipv4Addr,
625 guest_ipv4: Ipv4Addr,
626) {
627 let mut disconnected = false;
637
638 if let Some(receiver) = tcp_receiver.as_mut() {
639 loop {
640 match receiver.try_recv() {
641 Ok(connection) => {
642 let guest_destination =
643 SocketAddr::new(std::net::IpAddr::V4(guest_ipv4), connection.guest_port);
644 virtio_net_log!(
645 "virtio-net: accepted published TCP connection peer={} host_port={} guest_destination={}",
646 connection.peer_addr,
647 connection.host_port,
648 guest_destination
649 );
650 if !relays.create_published_socket(
651 interface,
652 gateway_ipv4,
653 guest_destination,
654 connection.stream,
655 sockets,
656 ) {
657 tracing::warn!(
658 host_port = connection.host_port,
659 guest_port = connection.guest_port,
660 peer_addr = %connection.peer_addr,
661 "dropping published TCP connection because the guest relay path could not be created"
662 );
663 }
664 }
665 Err(TryRecvError::Empty) => break,
666 Err(TryRecvError::Disconnected) => {
667 disconnected = true;
668 break;
669 }
670 }
671 }
672 }
673
674 if disconnected {
675 *tcp_receiver = None;
676 }
677}
678
679fn process_dns_queries(
680 dns_socket_handle: SocketHandle,
681 sockets: &mut SocketSet<'_>,
682 egress: &EgressPolicy,
683 upstream_dns: Ipv4Addr,
684) {
685 let socket = sockets.get_mut::<UdpSocket>(dns_socket_handle);
689 while socket.can_recv() {
690 let (query, metadata) = match socket.recv() {
691 Ok((q, m)) => (q.to_vec(), m),
692 Err(_) => break,
693 };
694 virtio_net_log!(
695 "virtio-net: forwarding guest DNS query guest={} local_address={:?} query_len={} upstream_dns={}",
696 metadata.endpoint,
697 metadata.local_address,
698 query.len(),
699 upstream_dns
700 );
701 let response =
704 match filtered_dns_response(&query, egress, |q| forward_dns_query(upstream_dns, q)) {
705 Some(response) => response,
706 None => continue,
707 };
708 virtio_net_log!(
709 "virtio-net: forwarded DNS response back to guest guest={} response_len={}",
710 metadata.endpoint,
711 response.len()
712 );
713
714 let response_meta = UdpMetadata {
715 endpoint: metadata.endpoint,
716 local_address: metadata.local_address,
717 meta: Default::default(),
718 };
719 let _ = socket.send_slice(&response, response_meta);
720 }
721}
722
723fn forward_dns_query(upstream_dns: Ipv4Addr, query: &[u8]) -> std::io::Result<Vec<u8>> {
724 let socket = HostUdpSocket::bind((Ipv4Addr::UNSPECIFIED, 0))?;
732 socket.set_read_timeout(Some(Duration::from_secs(2)))?;
733 let local_addr = socket.local_addr()?;
734 virtio_net_log!(
735 "virtio-net: sending DNS query to upstream resolver local_addr={} upstream_dns={} query_len={}",
736 local_addr,
737 upstream_dns,
738 query.len()
739 );
740 socket.send_to(query, (upstream_dns, DNS_SOCKET_PORT))?;
741
742 let mut buffer = vec![0u8; DNS_BUFFER_BYTES];
743 let (bytes_read, _) = socket.recv_from(&mut buffer)?;
744 buffer.truncate(bytes_read);
745 virtio_net_log!(
746 "virtio-net: received DNS response from upstream resolver upstream_dns={} response_len={}",
747 upstream_dns,
748 buffer.len()
749 );
750 Ok(buffer)
751}
752
753fn filtered_dns_response(
763 query: &[u8],
764 egress: &EgressPolicy,
765 forward: impl FnOnce(&[u8]) -> std::io::Result<Vec<u8>>,
766) -> Option<Vec<u8>> {
767 if !egress.dns_filter_active() {
768 return match forward(query) {
769 Ok(response) => Some(response),
770 Err(err) => {
771 virtio_net_log!("virtio-net: host DNS forwarding failed error={}", err);
772 None
773 }
774 };
775 }
776 match dns::question_name(query) {
777 Some(name) if egress.hostname_allowed(&name) => match forward(query) {
778 Ok(response) => {
779 egress.learn_ip_records(&dns::answer_ip_records(&response));
780 Some(response)
781 }
782 Err(err) => {
783 virtio_net_log!("virtio-net: host DNS forwarding failed error={}", err);
784 None
785 }
786 },
787 Some(name) => {
788 virtio_net_log!(
789 "virtio-net: blocking DNS query by allow-host policy name={}",
790 name
791 );
792 Some(dns::error_response(query, dns::DNS_RCODE_NXDOMAIN))
793 }
794 None => Some(dns::error_response(query, dns::DNS_RCODE_SERVFAIL)),
795 }
796}
797
798fn forward_dns_query_tcp(upstream_dns: Ipv4Addr, query: &[u8]) -> std::io::Result<Vec<u8>> {
802 use std::io::{Error, ErrorKind};
803 let len = u16::try_from(query.len())
804 .map_err(|_| Error::new(ErrorKind::InvalidInput, "DNS query too large for TCP"))?;
805 let mut stream = std::net::TcpStream::connect_timeout(
806 &SocketAddr::new(IpAddr::V4(upstream_dns), DNS_SOCKET_PORT),
807 Duration::from_secs(2),
808 )?;
809 stream.set_read_timeout(Some(Duration::from_secs(2)))?;
810 stream.set_write_timeout(Some(Duration::from_secs(2)))?;
811 stream.write_all(&len.to_be_bytes())?;
812 stream.write_all(query)?;
813 stream.flush()?;
814
815 let mut len_buf = [0u8; 2];
816 stream.read_exact(&mut len_buf)?;
817 let resp_len = u16::from_be_bytes(len_buf) as usize;
818 if resp_len == 0 || resp_len > DNS_TCP_MAX_MSG {
819 return Err(Error::new(
820 ErrorKind::InvalidData,
821 "upstream DNS/TCP response length out of range",
822 ));
823 }
824 let mut response = vec![0u8; resp_len];
825 stream.read_exact(&mut response)?;
826 Ok(response)
827}
828
829fn add_dns_tcp_sockets(sockets: &mut SocketSet<'_>) -> Vec<SocketHandle> {
833 (0..DNS_TCP_LISTENERS)
834 .map(|_| {
835 let rx_buffer = tcp::SocketBuffer::new(vec![0u8; DNS_TCP_RX_BYTES]);
836 let tx_buffer = tcp::SocketBuffer::new(vec![0u8; DNS_TCP_TX_BYTES]);
837 let mut socket = tcp::Socket::new(rx_buffer, tx_buffer);
838 socket
839 .listen(IpListenEndpoint {
840 addr: None,
841 port: DNS_SOCKET_PORT,
842 })
843 .expect("failed to listen on gateway DNS TCP socket");
844 sockets.add(socket)
845 })
846 .collect()
847}
848
849#[derive(Default)]
853struct DnsTcpConn {
854 rx: Vec<u8>,
856 tx: Vec<u8>,
858 tx_sent: usize,
860 done: bool,
862}
863
864fn process_dns_tcp(
869 handles: &[SocketHandle],
870 conns: &mut [DnsTcpConn],
871 sockets: &mut SocketSet<'_>,
872 egress: &EgressPolicy,
873 upstream_dns: Ipv4Addr,
874) {
875 for (handle, conn) in handles.iter().zip(conns.iter_mut()) {
876 let socket = sockets.get_mut::<tcp::Socket>(*handle);
877
878 if !socket.is_open() {
883 if !conn.rx.is_empty() || !conn.tx.is_empty() || conn.done {
884 *conn = DnsTcpConn::default();
885 }
886 let _ = socket.listen(IpListenEndpoint {
887 addr: None,
888 port: DNS_SOCKET_PORT,
889 });
890 continue;
891 }
892
893 if conn.done {
895 drain_dns_tcp_tx(socket, conn);
896 continue;
897 }
898
899 while socket.can_recv() {
901 let appended = socket.recv(|data| (data.len(), data.to_vec()));
902 match appended {
903 Ok(bytes) if !bytes.is_empty() => conn.rx.extend_from_slice(&bytes),
904 _ => break,
905 }
906 }
907
908 if conn.rx.len() > DNS_TCP_MAX_MSG + 2 {
910 conn.done = true;
911 socket.close();
912 continue;
913 }
914
915 if conn.rx.len() >= 2 {
916 let msg_len = u16::from_be_bytes([conn.rx[0], conn.rx[1]]) as usize;
917 if msg_len == 0 || msg_len > DNS_TCP_MAX_MSG {
918 conn.done = true;
919 socket.close();
920 continue;
921 }
922 if conn.rx.len() >= 2 + msg_len {
923 let query = conn.rx[2..2 + msg_len].to_vec();
924 virtio_net_log!(
925 "virtio-net: DNS/TCP query query_len={} upstream_dns={}",
926 query.len(),
927 upstream_dns
928 );
929 if let Some(response) = filtered_dns_response(&query, egress, |q| {
930 forward_dns_query_tcp(upstream_dns, q)
931 }) {
932 if let Ok(resp_len) = u16::try_from(response.len()) {
933 conn.tx.extend_from_slice(&resp_len.to_be_bytes());
934 conn.tx.extend_from_slice(&response);
935 }
936 }
937 conn.done = true;
938 drain_dns_tcp_tx(socket, conn);
939 }
940 }
941 }
942}
943
944fn drain_dns_tcp_tx(socket: &mut tcp::Socket<'_>, conn: &mut DnsTcpConn) {
947 while conn.tx_sent < conn.tx.len() && socket.can_send() {
948 match socket.send_slice(&conn.tx[conn.tx_sent..]) {
949 Ok(n) if n > 0 => conn.tx_sent += n,
950 _ => break,
951 }
952 }
953 if conn.tx_sent >= conn.tx.len() {
954 socket.close();
955 }
956}
957
958fn flush_interface_egress(
959 interface: &mut Interface,
960 device: &mut VirtioNetworkDevice,
961 sockets: &mut SocketSet<'_>,
962 now: Instant,
963) {
964 loop {
968 let result = interface.poll_egress(now, device, sockets);
969 if matches!(result, PollResult::None) {
970 break;
971 }
972 }
973}
974
975fn wake_guest_if_needed(queues: &NetworkFrameQueues, device: &VirtioNetworkDevice) {
976 if device.frames_emitted.swap(false, Ordering::Relaxed) {
980 queues.host_wake.wake();
981 }
982}
983
984fn smoltcp_now(clock: StdInstant) -> Instant {
985 let elapsed = clock.elapsed();
986 Instant::from_millis(elapsed.as_millis() as i64)
987}
988
989fn classify_guest_frame(frame: &[u8], gateway_addrs: &[IpAddr]) -> FrameAction {
990 let ethernet = match EthernetFrame::new_checked(frame) {
991 Ok(frame) => frame,
992 Err(_) => return FrameAction::Passthrough,
993 };
994
995 let (src_ip, dst_ip, protocol, transport): (IpAddr, IpAddr, _, _) = match ethernet.ethertype() {
1000 EthernetProtocol::Ipv4 => {
1001 let ipv4 = match Ipv4Packet::new_checked(ethernet.payload()) {
1002 Ok(packet) => packet,
1003 Err(_) => return FrameAction::Passthrough,
1004 };
1005 (
1006 IpAddr::V4(ipv4.src_addr()),
1007 IpAddr::V4(ipv4.dst_addr()),
1008 ipv4.next_header(),
1009 ipv4.payload(),
1010 )
1011 }
1012 EthernetProtocol::Ipv6 => {
1013 let ipv6 = match Ipv6Packet::new_checked(ethernet.payload()) {
1014 Ok(packet) => packet,
1015 Err(_) => return FrameAction::Passthrough,
1016 };
1017 (
1018 IpAddr::V6(ipv6.src_addr()),
1019 IpAddr::V6(ipv6.dst_addr()),
1020 ipv6.next_header(),
1021 ipv6.payload(),
1022 )
1023 }
1024 _ => return FrameAction::Passthrough,
1025 };
1026
1027 match protocol {
1028 smoltcp::wire::IpProtocol::Tcp => {
1029 let tcp = match TcpPacket::new_checked(transport) {
1030 Ok(packet) => packet,
1031 Err(_) => return FrameAction::Passthrough,
1032 };
1033
1034 if tcp.syn() && !tcp.ack() {
1035 if tcp.dst_port() == DNS_SOCKET_PORT && gateway_addrs.contains(&dst_ip) {
1040 FrameAction::Passthrough
1041 } else {
1042 FrameAction::TcpSyn {
1043 source: SocketAddr::new(src_ip, tcp.src_port()),
1044 destination: SocketAddr::new(dst_ip, tcp.dst_port()),
1045 }
1046 }
1047 } else {
1048 FrameAction::Passthrough
1049 }
1050 }
1051 smoltcp::wire::IpProtocol::Udp => {
1052 let udp = match UdpPacket::new_checked(transport) {
1053 Ok(packet) => packet,
1054 Err(_) => return FrameAction::Passthrough,
1055 };
1056
1057 if udp.dst_port() == DNS_SOCKET_PORT {
1058 FrameAction::DnsQuery
1059 } else {
1060 FrameAction::UdpFlow {
1061 destination: SocketAddr::new(dst_ip, udp.dst_port()),
1062 }
1063 }
1064 }
1065 _ => FrameAction::Passthrough,
1066 }
1067}
1068
1069#[cfg(feature = "fuzzing")]
1075pub fn fuzz_classify_guest_frame(frame: &[u8]) {
1076 let _ = classify_guest_frame(frame, &[]);
1077}
1078
1079#[cfg(test)]
1080mod tests {
1081 use super::*;
1082
1083 fn tcp_syn_frame(dst_ip: [u8; 4], dst_port: u16) -> Vec<u8> {
1087 let mut f = Vec::new();
1088 f.extend_from_slice(&[0xff; 6]);
1090 f.extend_from_slice(&[0x02, 0, 0, 0, 0, 1]);
1091 f.extend_from_slice(&[0x08, 0x00]);
1092 f.extend_from_slice(&[0x45, 0x00, 0x00, 0x28, 0, 0, 0, 0, 0x40, 0x06, 0, 0]);
1094 f.extend_from_slice(&[10, 0, 0, 2]); f.extend_from_slice(&dst_ip);
1096 f.extend_from_slice(&54321u16.to_be_bytes());
1098 f.extend_from_slice(&dst_port.to_be_bytes());
1099 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
1102 }
1103
1104 #[test]
1105 fn dns_tcp_to_gateway_is_intercepted_not_relayed() {
1106 let gw = IpAddr::V4(Ipv4Addr::new(100, 96, 0, 1));
1107 assert_eq!(
1109 classify_guest_frame(&tcp_syn_frame([100, 96, 0, 1], 53), &[gw]),
1110 FrameAction::Passthrough
1111 );
1112 }
1113
1114 #[test]
1115 fn dns_tcp_to_external_resolver_still_relayed() {
1116 let gw = IpAddr::V4(Ipv4Addr::new(100, 96, 0, 1));
1117 assert!(matches!(
1120 classify_guest_frame(&tcp_syn_frame([1, 1, 1, 1], 53), &[gw]),
1121 FrameAction::TcpSyn { .. }
1122 ));
1123 }
1124
1125 #[test]
1126 fn non_dns_tcp_to_gateway_still_relayed() {
1127 let gw = IpAddr::V4(Ipv4Addr::new(100, 96, 0, 1));
1128 assert!(matches!(
1130 classify_guest_frame(&tcp_syn_frame([100, 96, 0, 1], 443), &[gw]),
1131 FrameAction::TcpSyn { .. }
1132 ));
1133 }
1134}