Skip to main content

consomme/
udp.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4use super::Access;
5use super::BindError;
6use super::Client;
7use super::DropReason;
8use super::dhcp::DHCP_SERVER;
9use super::dhcpv6::DHCPV6_ALL_AGENTS_MULTICAST;
10use super::dhcpv6::DHCPV6_SERVER;
11use crate::ChecksumState;
12use crate::ConsommeState;
13use crate::FourTuple;
14use crate::IpAddresses;
15use crate::IpVersion;
16use crate::Ipv4Addresses;
17use crate::Ipv6Addresses;
18use crate::PortForwardKey;
19use crate::dns_resolver::DnsFlow;
20use crate::dns_resolver::DnsRequest;
21use crate::dns_resolver::DnsResponse;
22use inspect::Inspect;
23use inspect::InspectMut;
24use inspect_counters::Counter;
25use pal_async::interest::InterestSlot;
26use pal_async::interest::PollEvents;
27use pal_async::socket::PolledSocket;
28use smoltcp::phy::ChecksumCapabilities;
29use smoltcp::wire::ETHERNET_HEADER_LEN;
30use smoltcp::wire::EthernetAddress;
31use smoltcp::wire::EthernetFrame;
32use smoltcp::wire::EthernetProtocol;
33use smoltcp::wire::EthernetRepr;
34use smoltcp::wire::IPV4_HEADER_LEN;
35use smoltcp::wire::IPV6_HEADER_LEN;
36use smoltcp::wire::IpAddress;
37use smoltcp::wire::IpProtocol;
38use smoltcp::wire::Ipv4Packet;
39use smoltcp::wire::Ipv4Repr;
40use smoltcp::wire::Ipv6Packet;
41use smoltcp::wire::Ipv6Repr;
42use smoltcp::wire::UDP_HEADER_LEN;
43use smoltcp::wire::UdpPacket;
44use smoltcp::wire::UdpRepr;
45use socket2::Socket;
46use std::collections::HashMap;
47use std::collections::hash_map;
48use std::io::ErrorKind;
49use std::net::IpAddr;
50use std::net::Ipv4Addr;
51use std::net::Ipv6Addr;
52use std::net::SocketAddr;
53use std::net::SocketAddrV4;
54use std::net::SocketAddrV6;
55use std::net::UdpSocket;
56use std::task::Context;
57use std::task::Poll;
58use std::time::Duration;
59use std::time::Instant;
60
61use crate::DNS_PORT;
62
63#[cfg(unix)]
64use crate::unix as platform;
65#[cfg(windows)]
66use crate::windows as platform;
67
68pub(crate) struct Udp {
69    connections: HashMap<SocketAddr, UdpConnection>,
70    listeners: HashMap<PortForwardKey, UdpListener>,
71    timeout: Duration,
72}
73
74impl Udp {
75    pub fn new(timeout: Duration) -> Self {
76        Self {
77            connections: HashMap::new(),
78            listeners: HashMap::new(),
79            timeout,
80        }
81    }
82}
83
84impl InspectMut for Udp {
85    fn inspect_mut(&mut self, req: inspect::Request<'_>) {
86        let mut resp = req.respond();
87        for (addr, conn) in &mut self.connections {
88            let key = addr.to_string();
89            resp.field_mut(&key, conn);
90        }
91        for (key, listener) in &mut self.listeners {
92            resp.field_mut(&format!("listener:{key}"), listener);
93        }
94    }
95}
96
97#[derive(InspectMut)]
98struct UdpListener {
99    #[inspect(skip)]
100    socket: Option<PolledSocket<UdpSocket>>,
101    // The host address listening for UDP packets.
102    #[inspect(display)]
103    host_addr: SocketAddr,
104    /// The guest port to forward received packets to.
105    guest_port: u16,
106    stats: Stats,
107}
108
109#[derive(InspectMut)]
110struct UdpConnection {
111    #[inspect(skip)]
112    socket: Option<PolledSocket<UdpSocket>>,
113    host_port: u16,
114    #[inspect(display)]
115    guest_mac: EthernetAddress,
116    stats: Stats,
117    #[inspect(mut)]
118    recycle: bool,
119    #[inspect(debug)]
120    last_activity: Instant,
121    gso_size: Option<u16>,
122}
123
124#[derive(Inspect, Default)]
125struct Stats {
126    tx_packets: Counter,
127    tx_dropped: Counter,
128    tx_errors: Counter,
129    rx_packets: Counter,
130}
131
132impl UdpConnection {
133    fn poll_conn(
134        &mut self,
135        cx: &mut Context<'_>,
136        dst_addr: &SocketAddr,
137        state: &mut ConsommeState,
138        client: &mut impl Client,
139    ) -> bool {
140        if self.recycle {
141            return false;
142        }
143
144        let mut eth = EthernetFrame::new_unchecked(&mut state.buffer);
145        loop {
146            // Receive UDP packets while there are receive buffers available. This
147            // means we won't drop UDP packets at this level--instead, we only drop
148            // UDP packets if the kernel socket's receive buffer fills up. If this
149            // results in latency problems, then we could try sizing this buffer
150            // more carefully.
151            if client.rx_mtu() == 0 {
152                break true;
153            }
154
155            let header_offset = match dst_addr {
156                SocketAddr::V4(_) => IPV4_HEADER_LEN + UDP_HEADER_LEN,
157                SocketAddr::V6(_) => IPV6_HEADER_LEN + UDP_HEADER_LEN,
158            };
159
160            match self.socket.as_mut().unwrap().poll_io(
161                cx,
162                InterestSlot::Read,
163                PollEvents::IN,
164                |socket| {
165                    socket
166                        .get()
167                        .recv_from(&mut eth.payload_mut()[header_offset..])
168                },
169            ) {
170                Poll::Ready(Ok((n, src_addr))) => {
171                    let ft = match ConsommeState::translate_remote_address(
172                        &state.params,
173                        &mut state.local_addr_map,
174                        &src_addr,
175                        dst_addr.port(),
176                    ) {
177                        Some(ft) => ft,
178                        None => FourTuple {
179                            src: src_addr,
180                            dst: *dst_addr,
181                        },
182                    };
183                    let packet_len = build_udp_packet(
184                        &mut eth,
185                        ft.src.ip().into(),
186                        ft.dst.ip().into(),
187                        ft.src.port(),
188                        ft.dst.port(),
189                        n,
190                        state.params.gateway_mac,
191                        self.guest_mac,
192                    );
193                    let checksum_state = match dst_addr {
194                        SocketAddr::V4(_) => ChecksumState::UDP4,
195                        SocketAddr::V6(_) => ChecksumState::NONE,
196                    };
197                    client.recv(&eth.as_ref()[..packet_len], &checksum_state);
198                    self.stats.rx_packets.increment();
199                    self.last_activity = Instant::now();
200                }
201                Poll::Ready(Err(err)) => {
202                    tracelimit::error_ratelimited!(
203                        error = &err as &dyn std::error::Error,
204                        guest = %dst_addr,
205                        "udp recv error"
206                    );
207                    break false;
208                }
209                Poll::Pending => break true,
210            }
211        }
212    }
213}
214
215impl UdpListener {
216    fn poll_listener(
217        &mut self,
218        cx: &mut Context<'_>,
219        state: &mut ConsommeState,
220        client: &mut impl Client,
221        connections: &HashMap<SocketAddr, UdpConnection>,
222    ) {
223        let Some(socket) = self.socket.as_mut() else {
224            return;
225        };
226        let mut eth = EthernetFrame::new_unchecked(&mut state.buffer);
227        loop {
228            if client.rx_mtu() == 0 {
229                break;
230            }
231
232            let header_offset = match self.host_addr.ip() {
233                IpAddr::V4(_) => IPV4_HEADER_LEN + UDP_HEADER_LEN,
234                IpAddr::V6(_) => IPV6_HEADER_LEN + UDP_HEADER_LEN,
235            };
236            match socket.poll_io(cx, InterestSlot::Read, PollEvents::IN, |socket| {
237                socket
238                    .get()
239                    .recv_from(&mut eth.payload_mut()[header_offset..])
240            }) {
241                Poll::Ready(Ok((n, mut other_addr))) => {
242                    // Check if this connection originated from the same guest in order to adjust
243                    // the port in the crafted packet to match the guest value.
244                    if state.params.is_local_address(&other_addr) {
245                        for (guest_addr, connection) in connections.iter() {
246                            if other_addr.port() == connection.host_port {
247                                other_addr.set_port(guest_addr.port());
248                                break;
249                            }
250                        }
251                    }
252                    let Some(ft) = ConsommeState::translate_remote_address(
253                        &state.params,
254                        &mut state.local_addr_map,
255                        &other_addr,
256                        self.guest_port,
257                    ) else {
258                        continue;
259                    };
260                    tracing::trace!(
261                        ?other_addr,
262                        guest_port = self.guest_port,
263                        "Received UDP packet on listener"
264                    );
265                    let packet_len = build_udp_packet(
266                        &mut eth,
267                        ft.src.ip().into(),
268                        ft.dst.ip().into(),
269                        ft.src.port(),
270                        ft.dst.port(),
271                        n,
272                        state.params.gateway_mac,
273                        state.params.client_mac,
274                    );
275                    let checksum_state = match other_addr {
276                        SocketAddr::V4(_) => ChecksumState::UDP4,
277                        SocketAddr::V6(_) => ChecksumState::NONE,
278                    };
279                    client.recv(&eth.as_ref()[..packet_len], &checksum_state);
280                    self.stats.rx_packets.increment();
281                }
282                Poll::Ready(Err(err)) => {
283                    tracelimit::error_ratelimited!(
284                        error = &err as &dyn std::error::Error,
285                        host_addr = %self.host_addr,
286                        guest_port = self.guest_port,
287                        "udp listener recv error"
288                    );
289                    break;
290                }
291                Poll::Pending => break,
292            }
293        }
294    }
295}
296
297impl<T: Client> Access<'_, T> {
298    pub(crate) fn poll_udp(&mut self, cx: &mut Context<'_>) {
299        let timeout = self.inner.udp.timeout;
300        let now = Instant::now();
301
302        self.inner.udp.connections.retain(|dst_addr, conn| {
303            // Check if connection has timed out
304            if now.duration_since(conn.last_activity) > timeout {
305                tracing::debug!(
306                    guest = %dst_addr,
307                    "UDP connection timed out"
308                );
309                return false;
310            }
311
312            conn.poll_conn(cx, dst_addr, &mut self.inner.state, self.client)
313        });
314
315        for listener in self.inner.udp.listeners.values_mut() {
316            listener.poll_listener(
317                cx,
318                &mut self.inner.state,
319                self.client,
320                &self.inner.udp.connections,
321            );
322        }
323
324        while let Poll::Ready(Some(response)) = self.inner.dns.poll_udp_response(cx) {
325            if let Err(e) = self.send_dns_response(&response) {
326                tracelimit::error_ratelimited!(error = ?e, "Failed to send DNS response");
327            }
328        }
329    }
330
331    pub(crate) fn refresh_udp_driver(&mut self) {
332        self.inner.udp.connections.retain(|dst_addr, conn| {
333            let socket = conn.socket.take().unwrap().into_inner();
334            match PolledSocket::new(self.client.driver(), socket) {
335                Ok(socket) => {
336                    conn.socket = Some(socket);
337                    true
338                }
339                Err(err) => {
340                    tracing::warn!(
341                        error = &err as &dyn std::error::Error,
342                        guest = %dst_addr,
343                        "failed to update driver for udp connection"
344                    );
345                    false
346                }
347            }
348        });
349        self.inner.udp.listeners.retain(|key, listener| {
350            let socket = listener.socket.take().unwrap().into_inner();
351            match PolledSocket::new(self.client.driver(), socket) {
352                Ok(socket) => {
353                    listener.socket = Some(socket);
354                    true
355                }
356                Err(err) => {
357                    tracing::warn!(
358                        guest_port = key.guest_port,
359                        family = %key.family,
360                        error = &err as &dyn std::error::Error,
361                        "failed to update driver for udp listener"
362                    );
363                    false
364                }
365            }
366        });
367    }
368
369    pub(crate) fn handle_udp(
370        &mut self,
371        frame: &EthernetRepr,
372        addresses: &IpAddresses,
373        payload: &[u8],
374        checksum: &ChecksumState,
375    ) -> Result<(), DropReason> {
376        let udp_packet = UdpPacket::new_checked(payload)?;
377
378        // Parse UDP header and check gateway handling
379        let (guest_addr, dst_sock_addr) = match addresses {
380            IpAddresses::V4(addrs) => {
381                let udp = UdpRepr::parse(
382                    &udp_packet,
383                    &addrs.src_addr.into(),
384                    &addrs.dst_addr.into(),
385                    &checksum.caps(),
386                )?;
387
388                if udp.dst_port == DNS_PORT
389                    && self.inner.dns.should_intercept_static_queries()
390                    && self.handle_dns(
391                        frame,
392                        addrs.src_addr.into(),
393                        addrs.dst_addr.into(),
394                        &udp_packet,
395                        true,
396                    )?
397                {
398                    return Ok(());
399                }
400
401                // Check for gateway-destined packets
402                if addrs.dst_addr == self.inner.state.params.gateway_ip
403                    || addrs.dst_addr.is_broadcast()
404                {
405                    if self.handle_gateway_udp(frame, addrs, &udp_packet)? {
406                        return Ok(());
407                    }
408                }
409
410                let guest_addr = SocketAddr::V4(SocketAddrV4::new(addrs.src_addr, udp.src_port));
411
412                let dst_sock_addr = SocketAddr::V4(SocketAddrV4::new(addrs.dst_addr, udp.dst_port));
413
414                (guest_addr, dst_sock_addr)
415            }
416            IpAddresses::V6(addrs) => {
417                let udp = UdpRepr::parse(
418                    &udp_packet,
419                    &addrs.src_addr.into(),
420                    &addrs.dst_addr.into(),
421                    &checksum.caps(),
422                )?;
423
424                if udp.dst_port == DNS_PORT
425                    && self.inner.dns.should_intercept_static_queries()
426                    && self.handle_dns(
427                        frame,
428                        addrs.src_addr.into(),
429                        addrs.dst_addr.into(),
430                        &udp_packet,
431                        true,
432                    )?
433                {
434                    return Ok(());
435                }
436
437                // Check for gateway-destined packets (IPv6 uses multicast instead of broadcast)
438                if addrs.dst_addr == self.inner.state.params.gateway_link_local_ipv6
439                    || addrs.dst_addr == DHCPV6_ALL_AGENTS_MULTICAST
440                {
441                    if self.handle_gateway_udp_v6(frame, addrs, &udp_packet)? {
442                        return Ok(());
443                    }
444                }
445
446                let guest_addr =
447                    SocketAddr::V6(SocketAddrV6::new(addrs.src_addr, udp.src_port, 0, 0));
448
449                let dst_sock_addr =
450                    SocketAddr::V6(SocketAddrV6::new(addrs.dst_addr, udp.dst_port, 0, 0));
451
452                (guest_addr, dst_sock_addr)
453            }
454        };
455
456        // Resolve virtual mapped addresses back to the real host address.
457        let mut dst_sock_addr = self.inner.state.resolve_destination(&dst_sock_addr);
458        if self.inner.state.params.is_local_address(&dst_sock_addr) {
459            // This packet is destined for a local address. If the port matches a listener,
460            // translate it so that the connection loops back to the expected destination.
461            let key = PortForwardKey::from_socket_addr(dst_sock_addr, dst_sock_addr.port());
462            if let Some(listener) = self.inner.udp.listeners.get(&key) {
463                dst_sock_addr.set_port(listener.host_addr.port());
464            }
465        }
466
467        let conn = self.get_or_insert(guest_addr, Some(frame.src_addr))?;
468        let socket = conn.socket.as_ref().unwrap().get();
469        if conn.gso_size != checksum.gso {
470            platform::set_udp_gso_size(socket, checksum.gso.unwrap_or(0))
471                .map_err(DropReason::Io)?;
472            conn.gso_size = checksum.gso;
473        }
474        let result = platform::send_to(socket, udp_packet.payload(), &dst_sock_addr, checksum.gso);
475        match result {
476            Ok(_) => {
477                conn.stats.tx_packets.increment();
478                conn.last_activity = Instant::now();
479                Ok(())
480            }
481            Err(err) if err.kind() == ErrorKind::WouldBlock => {
482                conn.stats.tx_dropped.increment();
483                Err(DropReason::SendBufferFull)
484            }
485            Err(err) => {
486                conn.stats.tx_errors.increment();
487                Err(DropReason::Io(err))
488            }
489        }
490    }
491
492    fn get_or_insert(
493        &mut self,
494        guest_addr: SocketAddr,
495        guest_mac: Option<EthernetAddress>,
496    ) -> Result<&mut UdpConnection, DropReason> {
497        let entry = self.inner.udp.connections.entry(guest_addr);
498        match entry {
499            hash_map::Entry::Occupied(conn) => Ok(conn.into_mut()),
500            hash_map::Entry::Vacant(e) => {
501                let bind_addr: SocketAddr = match guest_addr {
502                    SocketAddr::V4(_) => {
503                        SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0))
504                    }
505                    SocketAddr::V6(_) => {
506                        SocketAddr::V6(SocketAddrV6::new(Ipv6Addr::UNSPECIFIED, 0, 0, 0))
507                    }
508                };
509
510                let socket = UdpSocket::bind(bind_addr).map_err(DropReason::Io)?;
511                let socket =
512                    PolledSocket::new(self.client.driver(), socket).map_err(DropReason::Io)?;
513                let host_port = socket.get().local_addr().map_err(DropReason::Io)?.port();
514                let conn = UdpConnection {
515                    socket: Some(socket),
516                    host_port,
517                    guest_mac: guest_mac.unwrap_or(self.inner.state.params.client_mac),
518                    stats: Default::default(),
519                    recycle: false,
520                    last_activity: Instant::now(),
521                    gso_size: None,
522                };
523                Ok(e.insert(conn))
524            }
525        }
526    }
527
528    fn handle_gateway_udp(
529        &mut self,
530        frame: &EthernetRepr,
531        addresses: &Ipv4Addresses,
532        udp: &UdpPacket<&[u8]>,
533    ) -> Result<bool, DropReason> {
534        match udp.dst_port() {
535            DHCP_SERVER => {
536                self.handle_dhcp(udp.payload())?;
537                Ok(true)
538            }
539            DNS_PORT => self.handle_dns(
540                frame,
541                addresses.src_addr.into(),
542                addresses.dst_addr.into(),
543                udp,
544                false,
545            ),
546            _ => Ok(false),
547        }
548    }
549
550    fn handle_gateway_udp_v6(
551        &mut self,
552        frame: &EthernetRepr,
553        addresses: &Ipv6Addresses,
554        udp: &UdpPacket<&[u8]>,
555    ) -> Result<bool, DropReason> {
556        let payload = udp.payload();
557        match udp.dst_port() {
558            DHCPV6_SERVER => {
559                self.handle_dhcpv6(payload, Some(addresses.src_addr))?;
560                Ok(true)
561            }
562            DNS_PORT => self.handle_dns(
563                frame,
564                addresses.src_addr.into(),
565                addresses.dst_addr.into(),
566                udp,
567                false,
568            ),
569            _ => Ok(false),
570        }
571    }
572
573    /// Binds to the specified host IP and port for forwarding inbound UDP
574    /// packets to the guest.
575    pub fn bind_udp_port(&mut self, socket: Socket, guest_port: u16) -> Result<(), BindError> {
576        let socket: UdpSocket = socket.into();
577        let socket = PolledSocket::new(self.client.driver(), socket).map_err(BindError::Io)?;
578        let host_addr = socket.get().local_addr().map_err(BindError::Io)?;
579        let key = PortForwardKey::from_socket_addr(host_addr, guest_port);
580        if self.inner.udp.listeners.contains_key(&key) {
581            return Err(BindError::PortAlreadyBound(guest_port));
582        }
583        self.inner.udp.listeners.insert(
584            key,
585            UdpListener {
586                socket: Some(socket),
587                host_addr,
588                guest_port,
589                stats: Default::default(),
590            },
591        );
592        Ok(())
593    }
594
595    /// Unbinds from the specified guest port and IP family.
596    pub fn unbind_udp_port(&mut self, family: IpVersion, port: u16) -> Result<(), BindError> {
597        if self
598            .inner
599            .udp
600            .listeners
601            .remove(&PortForwardKey::new(family, port))
602            .is_some()
603        {
604            Ok(())
605        } else {
606            Err(BindError::PortNotBound)
607        }
608    }
609
610    fn handle_dns(
611        &mut self,
612        frame: &EthernetRepr,
613        src_addr: IpAddress,
614        dst_addr: IpAddress,
615        udp: &UdpPacket<&[u8]>,
616        forward_static_misses: bool,
617    ) -> Result<bool, DropReason> {
618        let flow = DnsFlow {
619            src: SocketAddr::new(src_addr.into(), udp.src_port()),
620            dst: SocketAddr::new(dst_addr.into(), udp.dst_port()),
621            gateway_mac: self.inner.state.params.gateway_mac,
622            client_mac: frame.src_addr,
623            transport: crate::dns_resolver::DnsTransport::Udp,
624        };
625
626        // Limit static DNS response sizes to the MTU, and to the 512-byte
627        // maximum for DNS over UDP (no EDNS0 negotiation is performed).
628        let ip_header_len = if matches!(dst_addr, IpAddress::Ipv4(_)) {
629            IPV4_HEADER_LEN
630        } else {
631            IPV6_HEADER_LEN
632        };
633
634        let max_response_len = self
635            .client
636            .rx_mtu()
637            .saturating_sub(ETHERNET_HEADER_LEN + ip_header_len + UDP_HEADER_LEN)
638            .min(crate::dns_resolver::MAX_DNS_UDP_RESPONSE_LEN);
639
640        if let Some(response_data) = self
641            .inner
642            .dns
643            .build_static_response(udp.payload(), max_response_len)
644        {
645            let response = DnsResponse {
646                flow,
647                response_data,
648            };
649            if let Err(e) = self.send_dns_response(&response) {
650                tracelimit::error_ratelimited!(error = ?e, "Failed to send static DNS response");
651            }
652            return Ok(true);
653        }
654
655        if forward_static_misses {
656            return Ok(false);
657        }
658
659        let request = DnsRequest {
660            flow,
661            dns_query: udp.payload(),
662        };
663
664        // Submit the DNS query with addressing information.
665        // The response will be queued and sent later in poll_udp, unless the
666        // resolver cannot accept the query, in which case it returns a
667        // SERVFAIL to emit immediately.
668        let immediate_response = self.inner.dns.submit_udp_query(&request).map_err(|e| {
669            tracelimit::error_ratelimited!(error = ?e, "Failed to start DNS query");
670            DropReason::Packet(smoltcp::wire::Error)
671        })?;
672
673        if let Some(response) = immediate_response {
674            if let Err(e) = self.send_dns_response(&response) {
675                tracelimit::error_ratelimited!(error = ?e, "Failed to send DNS SERVFAIL response");
676            }
677        }
678
679        Ok(true)
680    }
681
682    fn send_dns_response(&mut self, response: &DnsResponse) -> Result<(), DropReason> {
683        tracing::debug!(
684            response_len = response.response_data.len(),
685            src = %response.flow.src,
686            dst = %response.flow.dst,
687            "Sending UDP DNS response"
688        );
689
690        let buffer = &mut self.inner.state.buffer;
691
692        // Determine header length based on IP version
693        let (ip_header_len, checksum_state) = match response.flow.src.ip() {
694            IpAddr::V4(_) => (IPV4_HEADER_LEN, ChecksumState::UDP4),
695            IpAddr::V6(_) => (IPV6_HEADER_LEN, ChecksumState::NONE),
696        };
697
698        let payload_offset = ETHERNET_HEADER_LEN + ip_header_len + UDP_HEADER_LEN;
699        let required_size = payload_offset + response.response_data.len();
700
701        if required_size > buffer.len() {
702            return Err(DropReason::SendBufferFull);
703        }
704
705        buffer[payload_offset..required_size].copy_from_slice(&response.response_data);
706
707        let mut eth_frame = EthernetFrame::new_unchecked(&mut buffer[..]);
708        let frame_len = build_udp_packet(
709            &mut eth_frame,
710            response.flow.dst.ip().into(),
711            response.flow.src.ip().into(),
712            response.flow.dst.port(),
713            response.flow.src.port(),
714            response.response_data.len(),
715            response.flow.gateway_mac,
716            response.flow.client_mac,
717        );
718
719        self.client.recv(&buffer[..frame_len], &checksum_state);
720
721        Ok(())
722    }
723
724    #[cfg(test)]
725    /// Returns the current number of active UDP connections.
726    pub fn udp_connection_count(&self) -> usize {
727        self.inner.udp.connections.len()
728    }
729}
730
731/// Helper function to build a complete UDP packet in an Ethernet frame.
732///
733/// This function constructs the Ethernet, IP (v4 or v6), and UDP headers, and assumes
734/// the UDP payload is already present in the buffer at the correct offset.
735///
736/// Returns the total length of the constructed frame.
737fn build_udp_packet<T: AsRef<[u8]> + AsMut<[u8]> + ?Sized>(
738    eth_frame: &mut EthernetFrame<&mut T>,
739    src_ip: IpAddress,
740    dst_ip: IpAddress,
741    src_port: u16,
742    dst_port: u16,
743    payload_len: usize,
744    src_mac: EthernetAddress,
745    dst_mac: EthernetAddress,
746) -> usize {
747    // Build Ethernet header
748    eth_frame.set_src_addr(src_mac);
749    eth_frame.set_dst_addr(dst_mac);
750
751    match (src_ip, dst_ip) {
752        (IpAddress::Ipv4(src_ip), IpAddress::Ipv4(dst_ip)) => {
753            eth_frame.set_ethertype(EthernetProtocol::Ipv4);
754
755            // Build IPv4 header
756            let mut ipv4_packet = Ipv4Packet::new_unchecked(eth_frame.payload_mut());
757            let ipv4_repr = Ipv4Repr {
758                src_addr: src_ip,
759                dst_addr: dst_ip,
760                next_header: IpProtocol::Udp,
761                payload_len: UDP_HEADER_LEN + payload_len,
762                hop_limit: 64,
763            };
764            ipv4_repr.emit(&mut ipv4_packet, &ChecksumCapabilities::default());
765
766            // Build UDP header (payload is already in place)
767            let mut udp_packet = UdpPacket::new_unchecked(ipv4_packet.payload_mut());
768            udp_packet.set_src_port(src_port);
769            udp_packet.set_dst_port(dst_port);
770            udp_packet.set_len((UDP_HEADER_LEN + payload_len) as u16);
771            udp_packet.fill_checksum(&src_ip.into(), &dst_ip.into());
772
773            // Return total frame length
774            ETHERNET_HEADER_LEN + ipv4_packet.total_len() as usize
775        }
776        (IpAddress::Ipv6(src_ip), IpAddress::Ipv6(dst_ip)) => {
777            eth_frame.set_ethertype(EthernetProtocol::Ipv6);
778
779            // Build IPv6 header
780            let mut ipv6_packet = Ipv6Packet::new_unchecked(eth_frame.payload_mut());
781            let ipv6_repr = Ipv6Repr {
782                src_addr: src_ip,
783                dst_addr: dst_ip,
784                next_header: IpProtocol::Udp,
785                payload_len: UDP_HEADER_LEN + payload_len,
786                hop_limit: 64,
787            };
788            ipv6_repr.emit(&mut ipv6_packet);
789
790            // Build UDP header (payload is already in place)
791            let mut udp_packet = UdpPacket::new_unchecked(ipv6_packet.payload_mut());
792            udp_packet.set_src_port(src_port);
793            udp_packet.set_dst_port(dst_port);
794            udp_packet.set_len((UDP_HEADER_LEN + payload_len) as u16);
795            udp_packet.fill_checksum(&src_ip.into(), &dst_ip.into());
796
797            // Return total frame length
798            ETHERNET_HEADER_LEN + ipv6_packet.total_len()
799        }
800        _ => panic!("mismatched IP address families"),
801    }
802}
803
804#[cfg(all(unix, test))]
805mod tests {
806    use super::*;
807    use crate::Consomme;
808    use crate::ConsommeParams;
809    use crate::IpVersion;
810    use crate::PortForwardKey;
811    use pal_async::DefaultDriver;
812    use parking_lot::Mutex;
813    use smoltcp::wire::Ipv4Address;
814    use std::io::ErrorKind;
815    use std::net::Ipv6Addr;
816    use std::net::SocketAddrV6;
817    use std::sync::Arc;
818
819    /// Mock test client that captures received packets
820    struct TestClient {
821        driver: Arc<DefaultDriver>,
822        received_packets: Arc<Mutex<Vec<Vec<u8>>>>,
823        rx_mtu: usize,
824    }
825
826    impl TestClient {
827        fn new(driver: Arc<DefaultDriver>) -> Self {
828            Self {
829                driver,
830                received_packets: Arc::new(Mutex::new(Vec::new())),
831                rx_mtu: 1514, // Standard Ethernet MTU
832            }
833        }
834    }
835
836    impl Client for TestClient {
837        fn driver(&self) -> &dyn pal_async::driver::Driver {
838            &*self.driver
839        }
840
841        fn recv(&mut self, data: &[u8], _checksum: &ChecksumState) {
842            self.received_packets.lock().push(data.to_vec());
843        }
844
845        fn rx_mtu(&mut self) -> usize {
846            self.rx_mtu
847        }
848    }
849
850    fn create_consomme_with_timeout(timeout: Duration) -> Consomme {
851        let mut params = ConsommeParams::new().expect("Failed to create params");
852        params.udp_timeout = timeout;
853        params.allow_host_local_access = true;
854        Consomme::new(params)
855    }
856
857    #[pal_async::async_test]
858    async fn test_udp_connection_timeout(driver: DefaultDriver) {
859        let driver = Arc::new(driver);
860        let mut consomme = create_consomme_with_timeout(Duration::from_millis(100));
861        let mut client = TestClient::new(driver);
862
863        let guest_mac = consomme.params_mut().client_mac;
864        let gateway_mac = consomme.params_mut().gateway_mac;
865        let guest_ip: Ipv4Address = consomme.params_mut().client_ip;
866        let target_ip: Ipv4Address = Ipv4Addr::LOCALHOST;
867
868        // Create a buffer and place the payload at the correct offset
869        let payload = b"test";
870        let mut buffer =
871            vec![0u8; ETHERNET_HEADER_LEN + IPV4_HEADER_LEN + UDP_HEADER_LEN + payload.len()];
872        buffer[ETHERNET_HEADER_LEN + IPV4_HEADER_LEN + UDP_HEADER_LEN..].copy_from_slice(payload);
873
874        let mut eth_frame = EthernetFrame::new_unchecked(&mut buffer[..]);
875        let packet_len = build_udp_packet(
876            &mut eth_frame,
877            IpAddress::Ipv4(guest_ip),
878            IpAddress::Ipv4(target_ip),
879            12345,
880            54321,
881            payload.len(),
882            guest_mac,
883            gateway_mac,
884        );
885
886        let mut access = consomme.access(&mut client);
887        let _ = access.send(&buffer[..packet_len], &ChecksumState::NONE);
888
889        let mut cx = Context::from_waker(std::task::Waker::noop());
890        access.poll(&mut cx);
891
892        assert_eq!(
893            access.udp_connection_count(),
894            1,
895            "Connection should be created"
896        );
897
898        // Manually update the last_activity to simulate timeout
899        for conn in access.inner.udp.connections.values_mut() {
900            conn.last_activity = Instant::now() - Duration::from_millis(150);
901        }
902
903        // Poll should remove timed out connections
904        access.poll(&mut cx);
905
906        assert_eq!(
907            access.udp_connection_count(),
908            0,
909            "Connection should be removed after timeout"
910        );
911    }
912
913    #[pal_async::async_test]
914    async fn test_udp_bind_port_forward(driver: DefaultDriver) {
915        let driver = Arc::new(driver);
916        let mut consomme = create_consomme_with_timeout(Duration::from_secs(30));
917        let mut client = TestClient::new(driver.clone());
918
919        // Bind a UDP listener socket on an ephemeral port.
920        let socket = Socket::new(socket2::Domain::IPV4, socket2::Type::DGRAM, None).unwrap();
921        socket
922            .bind(&SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0).into())
923            .unwrap();
924        let host_addr: SocketAddr = socket.local_addr().unwrap().as_socket().unwrap();
925
926        let guest_port = 5555;
927        let packets = client.received_packets.clone();
928        let mut access = consomme.access(&mut client);
929        access
930            .bind_udp_port(socket, guest_port)
931            .expect("bind should succeed");
932
933        assert!(
934            access
935                .inner
936                .udp
937                .listeners
938                .contains_key(&PortForwardKey::new(IpVersion::Ipv4, guest_port)),
939            "listener should be registered"
940        );
941
942        // Send a UDP packet to the listener from another socket.
943        let sender = UdpSocket::bind("127.0.0.1:0").unwrap();
944        sender.send_to(b"hello", host_addr).unwrap();
945
946        // Poll until the forwarded packet arrives (the first poll registers
947        // interest, subsequent polls receive the data).
948        let deadline = Instant::now() + Duration::from_secs(5);
949        loop {
950            std::future::poll_fn(|cx| {
951                access.poll(cx);
952                Poll::Ready(())
953            })
954            .await;
955
956            if !packets.lock().is_empty() {
957                break;
958            }
959            assert!(
960                Instant::now() < deadline,
961                "timed out waiting for forwarded UDP packet"
962            );
963            pal_async::timer::PolledTimer::new(&*driver)
964                .sleep(Duration::from_millis(10))
965                .await;
966        }
967
968        // Verify the packet targets the guest IP and the correct guest port.
969        let packets = packets.lock();
970        let pkt = &packets[0];
971        let eth = EthernetFrame::new_unchecked(pkt.as_slice());
972        let ipv4 = Ipv4Packet::new_unchecked(eth.payload());
973        let udp = UdpPacket::new_unchecked(ipv4.payload());
974        assert_eq!(
975            udp.dst_port(),
976            guest_port,
977            "forwarded packet should target the guest port"
978        );
979    }
980
981    #[pal_async::async_test]
982    async fn test_udp_bind_duplicate_port(driver: DefaultDriver) {
983        let driver = Arc::new(driver);
984        let mut consomme = create_consomme_with_timeout(Duration::from_secs(30));
985        let mut client = TestClient::new(driver.clone());
986
987        let guest_port = 6666;
988
989        let socket1 = Socket::new(socket2::Domain::IPV4, socket2::Type::DGRAM, None).unwrap();
990        socket1
991            .bind(&SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0).into())
992            .unwrap();
993
994        let socket2_inst = Socket::new(socket2::Domain::IPV4, socket2::Type::DGRAM, None).unwrap();
995        socket2_inst
996            .bind(&SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0).into())
997            .unwrap();
998
999        let mut access = consomme.access(&mut client);
1000        access
1001            .bind_udp_port(socket1, guest_port)
1002            .expect("first bind should succeed");
1003
1004        let err = access
1005            .bind_udp_port(socket2_inst, guest_port)
1006            .expect_err("duplicate bind should fail");
1007        assert!(
1008            matches!(err, BindError::PortAlreadyBound(_)),
1009            "error should be PortAlreadyBound"
1010        );
1011    }
1012
1013    #[pal_async::async_test]
1014    async fn test_udp_bind_same_port_different_families(driver: DefaultDriver) {
1015        let driver = Arc::new(driver);
1016        let mut consomme = create_consomme_with_timeout(Duration::from_secs(30));
1017        let mut client = TestClient::new(driver);
1018
1019        let guest_port = 6667;
1020
1021        let socket_v4 = Socket::new(socket2::Domain::IPV4, socket2::Type::DGRAM, None).unwrap();
1022        socket_v4
1023            .bind(&SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0).into())
1024            .unwrap();
1025
1026        let socket_v6 = Socket::new(socket2::Domain::IPV6, socket2::Type::DGRAM, None).unwrap();
1027        socket_v6.set_only_v6(true).unwrap();
1028        match socket_v6.bind(&SocketAddrV6::new(Ipv6Addr::LOCALHOST, 0, 0, 0).into()) {
1029            Ok(()) => {}
1030            Err(err)
1031                if matches!(
1032                    err.kind(),
1033                    ErrorKind::AddrNotAvailable | ErrorKind::Unsupported
1034                ) =>
1035            {
1036                return;
1037            }
1038            Err(err) => panic!("IPv6 bind failed: {err}"),
1039        }
1040
1041        let mut access = consomme.access(&mut client);
1042        access
1043            .bind_udp_port(socket_v4, guest_port)
1044            .expect("IPv4 bind should succeed");
1045        access
1046            .bind_udp_port(socket_v6, guest_port)
1047            .expect("IPv6 bind should succeed");
1048
1049        access
1050            .unbind_udp_port(IpVersion::Ipv4, guest_port)
1051            .expect("IPv4 unbind should succeed");
1052        assert!(
1053            access
1054                .inner
1055                .udp
1056                .listeners
1057                .contains_key(&PortForwardKey::new(IpVersion::Ipv6, guest_port)),
1058            "IPv6 listener should remain registered"
1059        );
1060    }
1061
1062    /// Test that when a UDP packet arrives from a loopback sender, the source
1063    /// IP forwarded to the guest is rewritten to a virtual address (not
1064    /// 127.0.0.1), so the guest routes its reply through the virtual adapter.
1065    #[pal_async::async_test]
1066    async fn test_udp_port_forward_loopback_src_rewritten(driver: DefaultDriver) {
1067        let driver = Arc::new(driver);
1068        let mut consomme = create_consomme_with_timeout(Duration::from_secs(30));
1069        let mut client = TestClient::new(driver.clone());
1070
1071        let client_ip: Ipv4Address = consomme.params_mut().client_ip;
1072
1073        // Bind a UDP listener socket on an ephemeral port.
1074        let socket = Socket::new(socket2::Domain::IPV4, socket2::Type::DGRAM, None).unwrap();
1075        socket
1076            .bind(&SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0).into())
1077            .unwrap();
1078        let host_addr: SocketAddr = socket.local_addr().unwrap().as_socket().unwrap();
1079
1080        let guest_port = 5556;
1081        let packets = client.received_packets.clone();
1082        let mut access = consomme.access(&mut client);
1083        access
1084            .bind_udp_port(socket, guest_port)
1085            .expect("bind should succeed");
1086
1087        // Send a UDP packet from loopback to the listener.
1088        let sender = UdpSocket::bind("127.0.0.1:0").unwrap();
1089        sender.send_to(b"hello", host_addr).unwrap();
1090
1091        // Poll until the forwarded packet arrives.
1092        let deadline = Instant::now() + Duration::from_secs(5);
1093        loop {
1094            std::future::poll_fn(|cx| {
1095                access.poll(cx);
1096                Poll::Ready(())
1097            })
1098            .await;
1099
1100            if !packets.lock().is_empty() {
1101                break;
1102            }
1103            assert!(
1104                Instant::now() < deadline,
1105                "timed out waiting for forwarded UDP packet"
1106            );
1107            pal_async::timer::PolledTimer::new(&*driver)
1108                .sleep(Duration::from_millis(10))
1109                .await;
1110        }
1111
1112        // Verify the source IP is NOT loopback and NOT the guest's own IP.
1113        let packets = packets.lock();
1114        let pkt = &packets[0];
1115        let eth = EthernetFrame::new_unchecked(pkt.as_slice());
1116        let ipv4 = Ipv4Packet::new_unchecked(eth.payload());
1117        let src_ip = ipv4.src_addr();
1118        let dst_ip = ipv4.dst_addr();
1119
1120        // Destination should be the guest.
1121        assert_eq!(dst_ip, client_ip);
1122        // Source must not be loopback — the guest would route replies to its
1123        // own loopback interface instead of back through the virtual NIC.
1124        assert!(
1125            !src_ip.is_loopback(),
1126            "forwarded UDP source IP should not be loopback, got {src_ip}"
1127        );
1128        // Source must not be the guest's own IP either.
1129        assert_ne!(
1130            src_ip, client_ip,
1131            "forwarded UDP source IP should not be the guest's own IP"
1132        );
1133    }
1134
1135    /// Test that the UDP loopback port remapping works end-to-end:
1136    /// when the guest sends a UDP packet to localhost on a listener port,
1137    /// consomme routes it through the host listener, and the returned packet
1138    /// has the guest's original source port (not the proxy ephemeral port).
1139    #[pal_async::async_test]
1140    async fn test_udp_loopback_port_remap(driver: DefaultDriver) {
1141        let driver = Arc::new(driver);
1142        let mut consomme = create_consomme_with_timeout(Duration::from_secs(30));
1143        let mut client = TestClient::new(driver.clone());
1144
1145        let guest_mac = consomme.params_mut().client_mac;
1146        let gateway_mac = consomme.params_mut().gateway_mac;
1147        let guest_ip: Ipv4Address = consomme.params_mut().client_ip;
1148        let dst_ip: IpAddress = IpAddress::Ipv4(Ipv4Addr::LOCALHOST);
1149        let listener_guest_port = 7070u16;
1150        let guest_src_port = 44444u16;
1151
1152        // Bind a UDP listener on an ephemeral host port, mapped to guest port 7070.
1153        let socket = Socket::new(socket2::Domain::IPV4, socket2::Type::DGRAM, None).unwrap();
1154        socket
1155            .bind(&SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0).into())
1156            .unwrap();
1157
1158        let packets = client.received_packets.clone();
1159        let mut access = consomme.access(&mut client);
1160        access
1161            .bind_udp_port(socket, listener_guest_port)
1162            .expect("bind should succeed");
1163
1164        // Guest sends a UDP packet to 127.0.0.1 on the listener port.
1165        // This simulates the guest trying to send to a host service that is
1166        // also forwarded back to the guest (loopback through consomme).
1167        let payload = b"loopback_test";
1168        let total_len = ETHERNET_HEADER_LEN + IPV4_HEADER_LEN + UDP_HEADER_LEN + payload.len();
1169        let mut buffer = vec![0u8; total_len];
1170        // Place the payload at the correct offset.
1171        buffer[ETHERNET_HEADER_LEN + IPV4_HEADER_LEN + UDP_HEADER_LEN..].copy_from_slice(payload);
1172
1173        let mut eth_frame = EthernetFrame::new_unchecked(&mut buffer[..]);
1174        let packet_len = build_udp_packet(
1175            &mut eth_frame,
1176            IpAddress::Ipv4(guest_ip),
1177            dst_ip,
1178            guest_src_port,
1179            listener_guest_port,
1180            payload.len(),
1181            guest_mac,
1182            gateway_mac,
1183        );
1184
1185        let _ = access.send(&buffer[..packet_len], &ChecksumState::NONE);
1186
1187        // Poll until the loopback packet arrives back at the guest.
1188        let deadline = Instant::now() + Duration::from_secs(5);
1189        loop {
1190            std::future::poll_fn(|cx| {
1191                access.poll(cx);
1192                Poll::Ready(())
1193            })
1194            .await;
1195
1196            let has_packet = packets.lock().iter().any(|p| {
1197                if p.len() < ETHERNET_HEADER_LEN + IPV4_HEADER_LEN + UDP_HEADER_LEN {
1198                    return false;
1199                }
1200                let eth = EthernetFrame::new_unchecked(p.as_slice());
1201                if eth.ethertype() != EthernetProtocol::Ipv4 {
1202                    return false;
1203                }
1204                let ipv4 = Ipv4Packet::new_unchecked(eth.payload());
1205                if ipv4.next_header() != IpProtocol::Udp {
1206                    return false;
1207                }
1208                let udp = UdpPacket::new_unchecked(ipv4.payload());
1209                udp.dst_port() == listener_guest_port
1210            });
1211            if has_packet {
1212                break;
1213            }
1214            assert!(
1215                Instant::now() < deadline,
1216                "timed out waiting for loopback UDP packet to be forwarded back to guest"
1217            );
1218            pal_async::timer::PolledTimer::new(&*driver)
1219                .sleep(Duration::from_millis(10))
1220                .await;
1221        }
1222
1223        // Find the packet forwarded to the guest on the listener port.
1224        let packets = packets.lock();
1225        let loopback_pkt = packets
1226            .iter()
1227            .find(|p| {
1228                if p.len() < ETHERNET_HEADER_LEN + IPV4_HEADER_LEN + UDP_HEADER_LEN {
1229                    return false;
1230                }
1231                let eth = EthernetFrame::new_unchecked(p.as_slice());
1232                if eth.ethertype() != EthernetProtocol::Ipv4 {
1233                    return false;
1234                }
1235                let ipv4 = Ipv4Packet::new_unchecked(eth.payload());
1236                if ipv4.next_header() != IpProtocol::Udp {
1237                    return false;
1238                }
1239                let udp = UdpPacket::new_unchecked(ipv4.payload());
1240                udp.dst_port() == listener_guest_port
1241            })
1242            .expect("should have received a loopback UDP packet");
1243
1244        let eth = EthernetFrame::new_unchecked(loopback_pkt.as_slice());
1245        let ipv4 = Ipv4Packet::new_unchecked(eth.payload());
1246        let udp = UdpPacket::new_unchecked(ipv4.payload());
1247        let src_ip = ipv4.src_addr();
1248
1249        // The source port should be the guest's original source port (remapped
1250        // from the proxy's ephemeral host port back to the guest port).
1251        assert_eq!(
1252            udp.src_port(),
1253            guest_src_port,
1254            "loopback UDP source port should be the guest's original source port \
1255             ({guest_src_port}), not a proxy ephemeral port; got {}",
1256            udp.src_port()
1257        );
1258        // Destination port should be the listener's guest port.
1259        assert_eq!(udp.dst_port(), listener_guest_port);
1260        // Source IP should not be loopback.
1261        assert!(
1262            !src_ip.is_loopback(),
1263            "loopback UDP source IP should not be 127.x.x.x, got {src_ip}"
1264        );
1265        // Source IP should not be the guest's own IP.
1266        assert_ne!(
1267            src_ip, guest_ip,
1268            "loopback UDP source IP should not be the guest's own IP"
1269        );
1270    }
1271}