1use 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 #[inspect(display)]
103 host_addr: SocketAddr,
104 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 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(ð.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 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(ð.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 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 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 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 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 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 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 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 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 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 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 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 pub fn udp_connection_count(&self) -> usize {
727 self.inner.udp.connections.len()
728 }
729}
730
731fn 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 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 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 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 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 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 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 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 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, }
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 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 for conn in access.inner.udp.connections.values_mut() {
900 conn.last_activity = Instant::now() - Duration::from_millis(150);
901 }
902
903 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 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 let sender = UdpSocket::bind("127.0.0.1:0").unwrap();
944 sender.send_to(b"hello", host_addr).unwrap();
945
946 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 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 #[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 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 let sender = UdpSocket::bind("127.0.0.1:0").unwrap();
1089 sender.send_to(b"hello", host_addr).unwrap();
1090
1091 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 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 assert_eq!(dst_ip, client_ip);
1122 assert!(
1125 !src_ip.is_loopback(),
1126 "forwarded UDP source IP should not be loopback, got {src_ip}"
1127 );
1128 assert_ne!(
1130 src_ip, client_ip,
1131 "forwarded UDP source IP should not be the guest's own IP"
1132 );
1133 }
1134
1135 #[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 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 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 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 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 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 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 assert_eq!(udp.dst_port(), listener_guest_port);
1260 assert!(
1262 !src_ip.is_loopback(),
1263 "loopback UDP source IP should not be 127.x.x.x, got {src_ip}"
1264 );
1265 assert_ne!(
1267 src_ip, guest_ip,
1268 "loopback UDP source IP should not be the guest's own IP"
1269 );
1270 }
1271}