1use alloc::{
34 boxed::Box,
35 collections::VecDeque,
36 string::{String, ToString},
37 sync::Arc,
38 vec,
39 vec::Vec,
40};
41use core::sync::atomic::{AtomicU64, Ordering};
42
43use ax_hal::time::{NANOS_PER_MICROS, monotonic_time_nanos};
44use ax_sync::SpinRwLock as RwLock;
45use smoltcp::{
46 iface::SocketSet,
47 phy::{DeviceCapabilities, Medium, PacketMeta},
48 storage::{PacketMetadata, RingBuffer},
49 time::Instant,
50 wire::{
51 IpAddress, IpCidr, IpProtocol, IpVersion, Ipv4Address, Ipv4Cidr, Ipv4Packet, Ipv6Packet,
52 TcpPacket,
53 },
54};
55
56use crate::{
57 LISTEN_TABLE,
58 config::{DeviceBinding, InterfaceId, RouteInfo},
59 consts::{SOCKET_BUFFER_SIZE, STANDARD_MTU},
60 device::{ArpEntry, Device, DeviceRxPacket, DeviceRxPoll, NetDeviceError},
61 ip_tos::apply_egress_ip_tos,
62 rx_meta::packet_meta_for_rx_packet,
63};
64
65const DEVICE_RX_WORKER_BATCH: usize = 16;
66
67#[derive(Debug, Clone)]
74pub struct NetDevStats {
75 pub interface_id: InterfaceId,
76 pub name: String,
77 pub rx_bytes: u64,
78 pub rx_packets: u64,
79 pub rx_errors: u64,
80 pub rx_dropped: u64,
81 pub tx_bytes: u64,
82 pub tx_packets: u64,
83 pub tx_errors: u64,
84 pub tx_dropped: u64,
85}
86
87#[derive(Debug)]
88pub struct Rule {
89 pub filter: IpCidr,
91 pub via: Option<IpAddress>,
93 pub dev: usize,
95 pub interface_id: InterfaceId,
97 pub src: IpAddress,
99 pub metric: u32,
101 pub order: u64,
103}
104
105impl Rule {
106 pub fn new(
108 filter: IpCidr,
109 via: Option<IpAddress>,
110 dev: usize,
111 interface_id: InterfaceId,
112 src: IpAddress,
113 metric: u32,
114 ) -> Self {
115 Self {
116 filter,
117 via,
118 dev,
119 interface_id,
120 src,
121 metric,
122 order: 0,
123 }
124 }
125
126 fn to_info(&self) -> RouteInfo {
127 RouteInfo {
128 filter: self.filter,
129 via: self.via,
130 interface_id: self.interface_id,
131 source: self.src,
132 metric: self.metric,
133 }
134 }
135}
136
137#[derive(Debug, Clone, Copy)]
138struct RxMetadata {
139 interface_id: InterfaceId,
140 packet_meta: PacketMeta,
141}
142
143type RouterPacketBuffer = smoltcp::storage::PacketBuffer<'static, RxMetadata>;
144type DevicePacketBuffer = smoltcp::storage::PacketBuffer<'static, InterfaceId>;
145
146#[derive(Clone)]
149struct TxPacket {
150 len: usize,
151 bytes: [u8; STANDARD_MTU],
152}
153
154impl TxPacket {
155 fn as_bytes(&self) -> &[u8] {
156 &self.bytes[..self.len]
157 }
158}
159
160struct OwnedRxPacket {
161 metadata: RxMetadata,
162 packet: DeviceRxPacket,
163}
164
165fn rx_metadata(interface_id: InterfaceId, packet: &[u8]) -> RxMetadata {
166 RxMetadata {
167 interface_id,
168 packet_meta: packet_meta_for_rx_packet(packet),
169 }
170}
171
172struct DeviceHandle {
174 interface_id: InterfaceId,
176 name: String,
178 inner: Box<dyn Device>,
180 rx_buffer: DevicePacketBuffer,
182 rx_bytes: AtomicU64,
186 rx_packets: AtomicU64,
187 rx_errors: AtomicU64,
188 rx_dropped: AtomicU64,
189 tx_bytes: AtomicU64,
190 tx_packets: AtomicU64,
191 tx_errors: AtomicU64,
192 tx_dropped: AtomicU64,
193}
194
195impl DeviceHandle {
196 fn new(interface_id: InterfaceId, device: Box<dyn Device>) -> Self {
197 let name = device.name().to_string();
198 Self {
199 interface_id,
200 name,
201 inner: device,
202 rx_buffer: DevicePacketBuffer::new(
203 vec![PacketMetadata::EMPTY; DEVICE_RX_WORKER_BATCH],
204 vec![0u8; STANDARD_MTU * DEVICE_RX_WORKER_BATCH],
205 ),
206 rx_bytes: AtomicU64::new(0),
207 rx_packets: AtomicU64::new(0),
208 rx_errors: AtomicU64::new(0),
209 rx_dropped: AtomicU64::new(0),
210 tx_bytes: AtomicU64::new(0),
211 tx_packets: AtomicU64::new(0),
212 tx_errors: AtomicU64::new(0),
213 tx_dropped: AtomicU64::new(0),
214 }
215 }
216
217 fn count_rx(&self, len: usize) {
223 self.rx_bytes.fetch_add(len as u64, Ordering::Relaxed);
229 self.rx_packets.fetch_add(1, Ordering::Relaxed);
230 }
231
232 fn count_tx(&self, len: usize) {
237 self.tx_bytes.fetch_add(len as u64, Ordering::Relaxed);
238 self.tx_packets.fetch_add(1, Ordering::Relaxed);
239 }
240
241 fn count_rx_errors(&self, n: u64) {
242 self.rx_errors.fetch_add(n, Ordering::Relaxed);
243 }
244
245 fn count_rx_dropped(&self, n: u64) {
246 self.rx_dropped.fetch_add(n, Ordering::Relaxed);
247 }
248
249 fn count_tx_errors(&self, n: u64) {
250 self.tx_errors.fetch_add(n, Ordering::Relaxed);
251 }
252
253 fn count_tx_dropped(&self, n: u64) {
254 self.tx_dropped.fetch_add(n, Ordering::Relaxed);
255 }
256
257 fn drain_device_counters(&mut self) {
258 for len in self.inner.drain_deferred_tx() {
259 self.count_tx(len);
260 }
261 for len in self.inner.drain_deferred_rx() {
262 self.count_rx(len);
263 }
264 let n = self.inner.drain_deferred_tx_errors();
265 if n > 0 {
266 self.count_tx_errors(n);
267 }
268 let n = self.inner.drain_deferred_tx_drops();
269 if n > 0 {
270 self.count_tx_dropped(n);
271 }
272 let n = self.inner.drain_deferred_rx_errors();
273 if n > 0 {
274 self.count_rx_errors(n);
275 }
276 let n = self.inner.drain_deferred_rx_drops();
277 if n > 0 {
278 self.count_rx_dropped(n);
279 }
280 }
281
282 fn stats(&self) -> NetDevStats {
283 NetDevStats {
284 interface_id: self.interface_id,
285 name: self.name.clone(),
286 rx_bytes: self.rx_bytes.load(Ordering::Relaxed),
287 rx_packets: self.rx_packets.load(Ordering::Relaxed),
288 rx_errors: self.rx_errors.load(Ordering::Relaxed),
289 rx_dropped: self.rx_dropped.load(Ordering::Relaxed),
290 tx_bytes: self.tx_bytes.load(Ordering::Relaxed),
291 tx_packets: self.tx_packets.load(Ordering::Relaxed),
292 tx_errors: self.tx_errors.load(Ordering::Relaxed),
293 tx_dropped: self.tx_dropped.load(Ordering::Relaxed),
294 }
295 }
296
297 fn send(&mut self, next_hop: IpAddress, packet: &[u8], timestamp: Instant) -> bool {
298 match self.try_send(next_hop, packet, timestamp) {
299 Ok(consumed) => consumed,
300 Err(NetDeviceError::Again) => false,
301 Err(error) => {
302 warn!("{}: transmit failed: {error:?}", self.name);
303 self.count_tx_errors(1);
304 self.drain_device_counters();
305 false
306 }
307 }
308 }
309
310 fn try_send(
311 &mut self,
312 next_hop: IpAddress,
313 packet: &[u8],
314 timestamp: Instant,
315 ) -> Result<bool, NetDeviceError> {
316 if packet.len() > STANDARD_MTU {
317 warn!(
318 "{}: packet to {} exceeds MTU ({} bytes), dropping",
319 self.name,
320 next_hop,
321 packet.len()
322 );
323 self.count_tx_dropped(1);
324 return Ok(false);
325 }
326 let frame_len = self.inner.try_send(next_hop, packet, timestamp)?;
327 if frame_len > 0 {
328 self.count_tx(frame_len);
329 }
330 self.drain_device_counters();
331 Ok(true)
332 }
333}
334
335fn now() -> Instant {
336 Instant::from_micros_const((monotonic_time_nanos() / NANOS_PER_MICROS) as i64)
337}
338
339#[derive(Debug, Clone, Copy)]
340pub struct RouteDecision {
341 pub dev: usize,
343 pub interface_id: InterfaceId,
345 pub source: IpAddress,
347 pub next_hop: IpAddress,
349 pub metric: u32,
351}
352
353pub struct RouteTable {
355 rules: Vec<Rule>,
356 next_order: u64,
357}
358impl RouteTable {
359 pub fn new() -> Self {
361 Self {
362 rules: Vec::new(),
363 next_order: 0,
364 }
365 }
366
367 pub fn add_rule(&mut self, mut rule: Rule) {
369 rule.order = self.next_order;
370 self.next_order = self.next_order.saturating_add(1);
371 self.rules.push(rule);
372 self.sort_rules();
373 }
374
375 fn sort_rules(&mut self) {
376 self.rules.sort_by(|a, b| {
377 b.filter
378 .prefix_len()
379 .cmp(&a.filter.prefix_len())
380 .then_with(|| a.metric.cmp(&b.metric))
381 .then_with(|| a.order.cmp(&b.order))
382 });
383 }
384
385 pub fn select_route_if(
387 &self,
388 dst: &IpAddress,
389 mut is_usable: impl FnMut(InterfaceId) -> bool,
390 ) -> Option<RouteDecision> {
391 self.rules
392 .iter()
393 .find(|rule| rule.filter.contains_addr(dst) && is_usable(rule.interface_id))
394 .map(|rule| RouteDecision {
395 dev: rule.dev,
396 interface_id: rule.interface_id,
397 source: rule.src,
398 next_hop: rule.via.unwrap_or(*dst),
399 metric: rule.metric,
400 })
401 }
402
403 pub fn select_route_for_source(
405 &self,
406 dst: &IpAddress,
407 source: &IpAddress,
408 ) -> Option<RouteDecision> {
409 self.rules
410 .iter()
411 .find(|rule| rule.filter.contains_addr(dst) && &rule.src == source)
412 .map(|rule| RouteDecision {
413 dev: rule.dev,
414 interface_id: rule.interface_id,
415 source: rule.src,
416 next_hop: rule.via.unwrap_or(*dst),
417 metric: rule.metric,
418 })
419 }
420
421 pub fn default_routes(&self) -> Vec<RouteInfo> {
423 self.rules
424 .iter()
425 .filter(|rule| match rule.filter {
426 IpCidr::Ipv4(cidr) => {
427 cidr.address() == Ipv4Address::UNSPECIFIED && cidr.prefix_len() == 0
428 }
429 _ => false,
430 })
431 .map(Rule::to_info)
432 .collect()
433 }
434
435 pub fn remove_ipv4_rules_for_interface(&mut self, interface_id: InterfaceId) {
437 self.rules.retain(|rule| {
438 !matches!(
439 rule.filter,
440 IpCidr::Ipv4(_) if rule.interface_id == interface_id
441 )
442 });
443 }
444
445 pub fn replace_ipv4_rules_for_interface(
447 &mut self,
448 interface_id: InterfaceId,
449 mut new_rules: Vec<Rule>,
450 ) {
451 self.remove_ipv4_rules_for_interface(interface_id);
452 for rule in &mut new_rules {
453 rule.order = self.next_order;
454 self.next_order = self.next_order.saturating_add(1);
455 }
456 self.rules.extend(new_rules);
457 self.sort_rules();
458 }
459}
460
461pub(crate) type SharedRouteTable = Arc<RwLock<RouteTable>>;
462
463pub struct Router {
465 rx_buffer: RouterPacketBuffer,
466 tx_buffer: RingBuffer<'static, TxPacket>,
467 pending_fanout: Vec<usize>,
470 ready_rx: VecDeque<OwnedRxPacket>,
472 devices: Vec<DeviceHandle>,
473 table: SharedRouteTable,
474}
475impl Router {
476 pub fn new(table: SharedRouteTable) -> Self {
478 let rx_buffer = RouterPacketBuffer::new(
479 vec![PacketMetadata::EMPTY; SOCKET_BUFFER_SIZE],
480 vec![0u8; STANDARD_MTU * SOCKET_BUFFER_SIZE],
481 );
482 let tx_buffer = RingBuffer::new(vec![
483 TxPacket {
484 len: 0,
485 bytes: [0; STANDARD_MTU],
486 };
487 SOCKET_BUFFER_SIZE
488 ]);
489 Self {
490 rx_buffer,
491 tx_buffer,
492 pending_fanout: Vec::new(),
493 ready_rx: VecDeque::with_capacity(SOCKET_BUFFER_SIZE),
494 devices: Vec::new(),
495 table,
496 }
497 }
498
499 pub fn add_rule(&mut self, rule: Rule) {
501 self.table.write().add_rule(rule);
502 }
503
504 pub fn add_device(&mut self, interface_id: InterfaceId, device: Box<dyn Device>) -> usize {
506 self.devices.push(DeviceHandle::new(interface_id, device));
507 self.devices.len() - 1
508 }
509
510 pub fn interface_id_for_dev(&self, dev: usize) -> Option<InterfaceId> {
512 self.devices.get(dev).map(|device| device.interface_id)
513 }
514
515 pub fn device_index_for_interface_id(&self, interface_id: InterfaceId) -> Option<usize> {
517 self.devices
518 .iter()
519 .position(|device| device.interface_id == interface_id)
520 }
521
522 pub fn device_names(&self) -> Vec<String> {
524 self.devices
525 .iter()
526 .map(|device| device.name.clone())
527 .collect()
528 }
529
530 pub fn set_ipv4_config(
532 &mut self,
533 dev: usize,
534 interface_id: InterfaceId,
535 metric: u32,
536 address: Option<Ipv4Cidr>,
537 gateway: Option<IpAddress>,
538 ) {
539 let new_rules = self.ipv4_rules(dev, interface_id, metric, address, gateway);
540 self.table
541 .write()
542 .replace_ipv4_rules_for_interface(interface_id, new_rules);
543 }
544
545 pub(crate) fn ipv4_rules(
547 &mut self,
548 dev: usize,
549 interface_id: InterfaceId,
550 metric: u32,
551 address: Option<Ipv4Cidr>,
552 gateway: Option<IpAddress>,
553 ) -> Vec<Rule> {
554 self.devices[dev].inner.set_ipv4_addr(address);
555
556 let mut rules = Vec::new();
557 if let Some(address) = address {
558 rules.push(Rule::new(
559 address.into(),
560 None,
561 dev,
562 interface_id,
563 address.address().into(),
564 metric,
565 ));
566 if let Some(gateway) = gateway {
567 rules.push(Rule::new(
568 Ipv4Cidr::new(Ipv4Address::UNSPECIFIED, 0).into(),
569 Some(gateway),
570 dev,
571 interface_id,
572 address.address().into(),
573 metric,
574 ));
575 }
576 }
577 rules
578 }
579
580 pub fn poll(
582 &mut self,
583 _timestamp: Instant,
584 sockets: &mut SocketSet<'_>,
585 mut snoop: impl FnMut(InterfaceId, &[u8]),
586 ) -> bool {
587 let mut moved_rx = false;
588 let Router {
589 rx_buffer,
590 ready_rx,
591 devices,
592 ..
593 } = self;
594 for device in devices {
595 if device.interface_id == InterfaceId::LOOPBACK {
596 continue;
597 }
598 let mut budget = DEVICE_RX_WORKER_BATCH;
599 while budget > 0 && ready_rx.len() < SOCKET_BUFFER_SIZE {
600 let interface_id = device.interface_id;
601 match device.inner.poll_owned_rx(now()) {
602 DeviceRxPoll::Packet(packet) => {
603 let metadata = packet.read_with(|bytes| {
604 snoop_tcp_packet(bytes, sockets);
605 snoop(interface_id, bytes);
606 rx_metadata(interface_id, bytes)
607 });
608 let frame_len = packet.frame_len();
609 ready_rx.push_back(OwnedRxPacket { metadata, packet });
610 device.count_rx(frame_len);
611 moved_rx = true;
612 budget -= 1;
613 continue;
614 }
615 DeviceRxPoll::Idle => break,
616 DeviceRxPoll::Unsupported => {}
617 }
618 if rx_buffer.is_full() || device.rx_buffer.is_full() {
619 break;
620 }
621 let mut frame_snoop = |_packet: &[u8]| {};
622 let direct = device.inner.recv_direct(
623 now(),
624 &mut |packet| {
625 snoop_tcp_packet(packet, sockets);
626 snoop(interface_id, packet);
627 let Ok(dst) =
628 rx_buffer.enqueue(packet.len(), rx_metadata(interface_id, packet))
629 else {
630 return false;
631 };
632 dst.copy_from_slice(packet);
633 true
634 },
635 &mut frame_snoop,
636 );
637 if let Some(frame_len) = direct {
638 if frame_len == 0 {
639 break;
640 }
641 device.count_rx(frame_len);
642 moved_rx = true;
643 budget -= 1;
644 continue;
645 }
646 let frame_len = device.inner.recv(
647 device.interface_id,
648 &mut device.rx_buffer,
649 now(),
650 &mut frame_snoop,
651 );
652 if frame_len == 0 {
653 break;
654 }
655 let Ok((interface_id, packet)) = device.rx_buffer.dequeue() else {
656 device.count_rx_errors(1);
657 break;
658 };
659 snoop_tcp_packet(packet, sockets);
660 snoop(interface_id, packet);
661 let Ok(dst) = rx_buffer.enqueue(packet.len(), rx_metadata(interface_id, packet))
662 else {
663 device.count_rx_dropped(1);
664 break;
665 };
666 dst.copy_from_slice(packet);
667 device.count_rx(frame_len);
668 moved_rx = true;
669 budget -= 1;
670 }
671 device.drain_device_counters();
672 }
673 moved_rx
674 }
675
676 pub fn send_on_device(
678 &mut self,
679 dev: usize,
680 next_hop: IpAddress,
681 packet: &[u8],
682 _timestamp: Instant,
683 ) -> bool {
684 let Router {
685 rx_buffer, devices, ..
686 } = self;
687 let device = &mut devices[dev];
688 if device.interface_id == InterfaceId::LOOPBACK {
689 let ok =
699 inject_loopback_rx_direct(rx_buffer, next_hop, packet, &mut SocketSet::new(vec![]));
700 if ok {
701 device.count_tx(packet.len());
702 device.count_rx(packet.len());
703 } else {
704 device.count_rx_dropped(1);
705 }
706 return ok;
707 }
708 device.send(next_hop, packet, now())
709 }
710
711 pub fn arp_entries(&self, timestamp: Instant) -> Vec<ArpEntry> {
713 let mut entries = Vec::new();
714 for device in &self.devices {
715 entries.extend(device.inner.arp_entries(timestamp));
716 }
717 entries
718 }
719
720 pub fn net_dev_stats(&self) -> Vec<NetDevStats> {
722 self.devices.iter().map(|device| device.stats()).collect()
723 }
724
725 pub fn register_waker(&self, _binding: DeviceBinding, _waker: &core::task::Waker) {
728 crate::request_poll();
729 }
730
731 pub fn dispatch(&mut self, _timestamp: Instant, sockets: &mut SocketSet<'_>) -> bool {
733 let mut poll_next = false;
734 let Router {
735 rx_buffer,
736 tx_buffer,
737 pending_fanout,
738 devices,
739 table,
740 ..
741 } = self;
742 while let Some(packet) = tx_buffer.get_allocated(0, 1).first() {
743 let packet = packet.as_bytes();
744 let outcome = match IpVersion::of_packet(packet).expect("got invalid IP packet") {
745 IpVersion::Ipv4 => {
746 let packet = smoltcp::wire::Ipv4Packet::new_checked(packet)
747 .expect("got invalid IPv4 packet");
748 let src_addr = IpAddress::Ipv4(packet.src_addr());
749 let dst_addr = IpAddress::Ipv4(packet.dst_addr());
750 if packet.dst_addr().is_broadcast() {
751 dispatch_link_local_fanout(
752 devices,
753 pending_fanout,
754 dst_addr,
755 packet.into_inner(),
756 )
757 } else {
758 dispatch_unicast_packet(
759 rx_buffer,
760 devices,
761 table,
762 src_addr,
763 dst_addr,
764 packet.into_inner(),
765 sockets,
766 )
767 }
768 }
769 IpVersion::Ipv6 => {
770 let packet = smoltcp::wire::Ipv6Packet::new_checked(packet)
771 .expect("got invalid IPv6 packet");
772 let src_addr = IpAddress::Ipv6(packet.src_addr());
773 let dst_addr = IpAddress::Ipv6(packet.dst_addr());
774 if packet.dst_addr().is_multicast() {
775 dispatch_link_local_fanout(
776 devices,
777 pending_fanout,
778 dst_addr,
779 packet.into_inner(),
780 )
781 } else {
782 dispatch_unicast_packet(
783 rx_buffer,
784 devices,
785 table,
786 src_addr,
787 dst_addr,
788 packet.into_inner(),
789 sockets,
790 )
791 }
792 }
793 };
794 match outcome {
795 DispatchOutcome::Consumed(next) => {
796 poll_next |= next;
797 tx_buffer
798 .dequeue_one()
799 .expect("the packet was only peeked while dispatching");
800 }
801 DispatchOutcome::Retry(next) => {
802 poll_next |= next;
803 break;
804 }
805 }
806 }
807 if tx_buffer.is_empty() {
808 tx_buffer.clear();
811 }
812 poll_next
813 }
814}
815
816fn dispatch_link_local_fanout(
817 devices: &mut [DeviceHandle],
818 pending: &mut Vec<usize>,
819 dst_addr: IpAddress,
820 packet: &[u8],
821) -> DispatchOutcome {
822 if pending.is_empty() {
823 pending.extend(devices.iter().enumerate().filter_map(|(index, dev)| {
826 (dev.interface_id != InterfaceId::LOOPBACK).then_some(index)
827 }));
828 }
829 let mut poll_next = false;
830 pending.retain(|&index| {
831 let dev = &mut devices[index];
832 match dev.try_send(dst_addr, packet, now()) {
833 Ok(consumed) => {
834 poll_next |= consumed;
835 false
836 }
837 Err(NetDeviceError::Again) => true,
838 Err(error) => {
839 warn!("{}: transmit failed: {error:?}", dev.name);
840 dev.count_tx_errors(1);
841 dev.drain_device_counters();
842 false
843 }
844 }
845 });
846 if pending.is_empty() {
847 DispatchOutcome::Consumed(poll_next)
848 } else {
849 DispatchOutcome::Retry(poll_next)
850 }
851}
852
853#[derive(Clone, Copy, Debug, Eq, PartialEq)]
854enum DispatchOutcome {
855 Consumed(bool),
856 Retry(bool),
857}
858
859fn dispatch_unicast_packet(
860 rx_buffer: &mut RouterPacketBuffer,
861 devices: &mut [DeviceHandle],
862 table: &SharedRouteTable,
863 src_addr: IpAddress,
864 dst_addr: IpAddress,
865 packet: &[u8],
866 sockets: &mut SocketSet<'_>,
867) -> DispatchOutcome {
868 let route = {
869 let routes = table.read();
870 let Some(route) = routes.select_route_for_source(&dst_addr, &src_addr) else {
871 debug!(
872 "No route found for source {} destination {}",
873 src_addr, dst_addr
874 );
875 return DispatchOutcome::Consumed(false);
881 };
882 route
883 };
884
885 let dev = &mut devices[route.dev];
886 if dev.interface_id == InterfaceId::LOOPBACK {
887 let ok = inject_loopback_rx_direct(rx_buffer, dst_addr, packet, sockets);
893 if ok {
894 dev.count_tx(packet.len());
895 dev.count_rx(packet.len());
896 } else {
897 dev.count_rx_dropped(1);
902 }
903 DispatchOutcome::Consumed(ok)
904 } else {
905 match dev.try_send(route.next_hop, packet, now()) {
906 Ok(consumed) => DispatchOutcome::Consumed(consumed),
907 Err(NetDeviceError::Again) => DispatchOutcome::Retry(false),
908 Err(error) => {
909 warn!("{}: transmit failed: {error:?}", dev.name);
910 dev.count_tx_errors(1);
911 dev.drain_device_counters();
912 DispatchOutcome::Consumed(false)
913 }
914 }
915 }
916}
917
918fn inject_loopback_rx_direct(
920 rx_buffer: &mut RouterPacketBuffer,
921 dst_addr: IpAddress,
922 packet: &[u8],
923 sockets: &mut SocketSet<'_>,
924) -> bool {
925 snoop_tcp_packet(packet, sockets);
926 let Ok(dst) = rx_buffer.enqueue(packet.len(), rx_metadata(InterfaceId::LOOPBACK, packet))
927 else {
928 warn!("Loopback: RX buffer full, dropping packet to {}", dst_addr);
929 return false;
930 };
931 dst.copy_from_slice(packet);
932 true
933}
934
935pub struct TxToken<'a>(&'a mut RingBuffer<'static, TxPacket>);
937
938impl smoltcp::phy::TxToken for TxToken<'_> {
939 fn consume<R, F>(self, len: usize, f: F) -> R
940 where
941 F: FnOnce(&mut [u8]) -> R,
942 {
943 let slot = self
946 .0
947 .enqueue_one()
948 .expect("This was checked before creating the TxToken");
949 slot.len = len;
950 let packet = &mut slot.bytes[..len];
951 let result = f(packet);
952 apply_egress_ip_tos(packet);
953 result
954 }
955}
956
957fn snoop_tcp_packet(buf: &[u8], sockets: &mut SocketSet<'_>) {
959 if buf.is_empty() {
960 return;
961 }
962 let (src_addr, dst_addr, payload) = match IpVersion::of_packet(buf) {
963 Ok(IpVersion::Ipv4) => {
964 let Ok(packet) = Ipv4Packet::new_checked(buf) else {
965 return;
966 };
967 if packet.next_header() != IpProtocol::Tcp {
968 return;
969 }
970 (
971 IpAddress::Ipv4(packet.src_addr()),
972 IpAddress::Ipv4(packet.dst_addr()),
973 packet.payload(),
974 )
975 }
976 Ok(IpVersion::Ipv6) => {
977 let Ok(packet) = Ipv6Packet::new_checked(buf) else {
978 return;
979 };
980 if packet.next_header() != IpProtocol::Tcp {
981 return;
982 }
983 (
984 IpAddress::Ipv6(packet.src_addr()),
985 IpAddress::Ipv6(packet.dst_addr()),
986 packet.payload(),
987 )
988 }
989 Err(_) => return,
990 };
991 let Ok(tcp_packet) = TcpPacket::new_checked(payload) else {
992 return;
993 };
994 let src_addr = (src_addr, tcp_packet.src_port()).into();
995 let dst_addr = (dst_addr, tcp_packet.dst_port()).into();
996 let is_first = tcp_packet.syn() && !tcp_packet.ack();
997 if is_first {
998 LISTEN_TABLE.incoming_tcp_packet(src_addr, dst_addr, sockets);
999 }
1000}
1001
1002enum RxTokenPacket<'a> {
1003 Borrowed(&'a [u8]),
1004 Owned(DeviceRxPacket),
1005}
1006
1007pub struct RxToken<'a> {
1009 interface_id: InterfaceId,
1010 packet_meta: PacketMeta,
1011 packet: RxTokenPacket<'a>,
1012}
1013
1014impl<'a> smoltcp::phy::RxToken for RxToken<'a> {
1015 fn consume<R, F>(self, f: F) -> R
1016 where
1017 F: FnOnce(&[u8]) -> R,
1018 {
1019 let _ingress_if = self.interface_id;
1020 match self.packet {
1021 RxTokenPacket::Borrowed(packet) => f(packet),
1022 RxTokenPacket::Owned(packet) => packet.consume(f),
1023 }
1024 }
1025
1026 fn meta(&self) -> PacketMeta {
1027 self.packet_meta
1028 }
1029}
1030
1031impl smoltcp::phy::Device for Router {
1032 type RxToken<'a> = RxToken<'a>;
1033 type TxToken<'a> = TxToken<'a>;
1034
1035 fn receive(&mut self, _timestamp: Instant) -> Option<(Self::RxToken<'_>, Self::TxToken<'_>)> {
1036 if self.tx_buffer.is_full() {
1037 return None;
1038 }
1039 let Self {
1040 rx_buffer,
1041 ready_rx,
1042 tx_buffer,
1043 ..
1044 } = self;
1045 let rx_token = if !rx_buffer.is_empty() {
1046 let (metadata, packet) = rx_buffer.dequeue().unwrap();
1047 RxToken {
1048 interface_id: metadata.interface_id,
1049 packet_meta: metadata.packet_meta,
1050 packet: RxTokenPacket::Borrowed(packet),
1051 }
1052 } else {
1053 let packet = ready_rx.pop_front()?;
1054 RxToken {
1055 interface_id: packet.metadata.interface_id,
1056 packet_meta: packet.metadata.packet_meta,
1057 packet: RxTokenPacket::Owned(packet.packet),
1058 }
1059 };
1060 Some((rx_token, TxToken(tx_buffer)))
1061 }
1062
1063 fn transmit(&mut self, _timestamp: Instant) -> Option<Self::TxToken<'_>> {
1064 if self.tx_buffer.is_full() {
1065 None
1066 } else {
1067 Some(TxToken(&mut self.tx_buffer))
1068 }
1069 }
1070
1071 fn capabilities(&self) -> DeviceCapabilities {
1072 let mut caps = DeviceCapabilities::default();
1073 caps.medium = Medium::Ip;
1074 caps.max_transmission_unit = STANDARD_MTU;
1075 caps.max_burst_size = Some(SOCKET_BUFFER_SIZE);
1076 caps
1081 }
1082}
1083
1084#[cfg(test)]
1085mod tests {
1086 use smoltcp::{
1087 phy::{Device as _, TxToken as _},
1088 storage::PacketBuffer,
1089 };
1090
1091 use super::*;
1092 use crate::device::TxChecksumCapabilities;
1093
1094 #[test]
1095 fn stack_tcp_and_udp_emit_complete_software_checksums() {
1096 use smoltcp::{
1097 iface::{Config, Interface},
1098 socket::{tcp, udp},
1099 wire::{HardwareAddress, UdpPacket},
1100 };
1101
1102 let mut router = Router::new(Arc::new(RwLock::new(RouteTable::new())));
1103 router.add_device(
1104 IF0,
1105 Box::new(crate::device::EthernetDevice::new(
1106 "checksum".into(),
1107 Box::new(ChecksumPort),
1108 None,
1109 )),
1110 );
1111 let now = Instant::from_millis(0);
1112 let mut interface = Interface::new(Config::new(HardwareAddress::Ip), &mut router, now);
1113 interface.update_ip_addrs(|addrs| {
1114 addrs
1115 .push(ipv4_cidr(Ipv4Address::new(10, 0, 0, 2), 24))
1116 .unwrap()
1117 });
1118 let destination = IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1));
1119 let mut sockets = SocketSet::new(vec![]);
1120 let mut tcp = tcp::Socket::new(
1121 tcp::SocketBuffer::new(vec![0; 1024]),
1122 tcp::SocketBuffer::new(vec![0; 1024]),
1123 );
1124 tcp.connect(interface.context(), (destination, 4321), 1234)
1125 .unwrap();
1126 sockets.add(tcp);
1127 interface.poll_egress(now, &mut router, &mut sockets);
1128 let packet = router
1129 .tx_buffer
1130 .dequeue_one()
1131 .expect("TCP SYN must be emitted");
1132 let ip = Ipv4Packet::new_checked(packet.as_bytes()).unwrap();
1133 let tcp = TcpPacket::new_checked(ip.payload()).unwrap();
1134 assert!(tcp.syn());
1135 assert!(tcp.verify_checksum(&ip.src_addr().into(), &ip.dst_addr().into()));
1136
1137 let mut udp = udp::Socket::new(
1138 udp::PacketBuffer::new(vec![udp::PacketMetadata::EMPTY; 1], vec![0; 64]),
1139 udp::PacketBuffer::new(vec![udp::PacketMetadata::EMPTY; 1], vec![0; 64]),
1140 );
1141 udp.bind(1235).unwrap();
1142 udp.send_slice(b"checksum", (destination, 4322)).unwrap();
1143 sockets.add(udp);
1144 interface.poll_egress(now, &mut router, &mut sockets);
1145 let packet = router
1146 .tx_buffer
1147 .dequeue_one()
1148 .expect("UDP packet must be emitted");
1149 let ip = Ipv4Packet::new_checked(packet.as_bytes()).unwrap();
1150 let udp = UdpPacket::new_checked(ip.payload()).unwrap();
1151 assert_ne!(udp.checksum(), 0, "ordinary UDP must generate a checksum");
1152 assert!(udp.verify_checksum(&ip.src_addr().into(), &ip.dst_addr().into()));
1153 assert_eq!(udp.payload(), b"checksum");
1154 }
1155
1156 #[test]
1157 fn loopback_preserves_raw_udp_checksum() {
1158 let table = Arc::new(RwLock::new(RouteTable::new()));
1159 let mut router = Router::new(table);
1160 let mut sockets = SocketSet::new(vec![]);
1161 for checksum in [0u16, 0x1234] {
1162 let mut packet = [0u8; 32];
1163 packet[0] = 0x45;
1164 packet[2..4].copy_from_slice(&32u16.to_be_bytes());
1165 packet[8] = 64;
1166 packet[9] = 17;
1167 packet[12..16].copy_from_slice(&[127, 0, 0, 1]);
1168 packet[16..20].copy_from_slice(&[127, 0, 0, 1]);
1169 packet[20..22].copy_from_slice(&1234u16.to_be_bytes());
1170 packet[22..24].copy_from_slice(&4321u16.to_be_bytes());
1171 packet[24..26].copy_from_slice(&12u16.to_be_bytes());
1172 packet[26..28].copy_from_slice(&checksum.to_be_bytes());
1173 assert!(inject_loopback_rx_direct(
1174 &mut router.rx_buffer,
1175 IpAddress::Ipv4(Ipv4Address::LOCALHOST),
1176 &packet,
1177 &mut sockets
1178 ));
1179 let (_, received) = router.rx_buffer.dequeue().unwrap();
1180 assert_eq!(
1181 received, &packet,
1182 "loopback rewrote raw UDP transport bytes"
1183 );
1184 }
1185 }
1186
1187 const IF0: InterfaceId = InterfaceId::new(2);
1188 const IF1: InterfaceId = InterfaceId::new(3);
1189 const SRC0: IpAddress = IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 2));
1190 const SRC1: IpAddress = IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 2));
1191
1192 struct EmptyDevice;
1193
1194 impl Device for EmptyDevice {
1195 fn name(&self) -> &str {
1196 "empty"
1197 }
1198
1199 fn recv(
1200 &mut self,
1201 _interface_id: InterfaceId,
1202 _buffer: &mut PacketBuffer<InterfaceId>,
1203 _timestamp: Instant,
1204 _snoop: &mut dyn FnMut(&[u8]),
1205 ) -> usize {
1206 0
1207 }
1208
1209 fn send(&mut self, _next_hop: IpAddress, _packet: &[u8], _timestamp: Instant) -> usize {
1210 0
1211 }
1212 }
1213
1214 struct RetryDevice;
1215
1216 impl Device for RetryDevice {
1217 fn name(&self) -> &str {
1218 "retry"
1219 }
1220
1221 fn recv(
1222 &mut self,
1223 _interface_id: InterfaceId,
1224 _buffer: &mut PacketBuffer<InterfaceId>,
1225 _timestamp: Instant,
1226 _snoop: &mut dyn FnMut(&[u8]),
1227 ) -> usize {
1228 0
1229 }
1230
1231 fn send(&mut self, _next_hop: IpAddress, _packet: &[u8], _timestamp: Instant) -> usize {
1232 0
1233 }
1234
1235 fn try_send(
1236 &mut self,
1237 _next_hop: IpAddress,
1238 _packet: &[u8],
1239 _timestamp: Instant,
1240 ) -> crate::device::NetDeviceResult<usize> {
1241 Err(NetDeviceError::Again)
1242 }
1243 }
1244
1245 struct ChecksumPort;
1246
1247 impl crate::device::EthernetFramePort for ChecksumPort {
1248 fn device_name(&self) -> &str {
1249 "checksum"
1250 }
1251 fn mac_address(&self) -> [u8; 6] {
1252 [2, 0, 0, 0, 0, 1]
1253 }
1254 fn checksum_capabilities(&self) -> TxChecksumCapabilities {
1255 TxChecksumCapabilities::TCP_UDP
1256 }
1257 fn transmit(
1258 &mut self,
1259 _: &crate::device::ProtocolEthernetFrame,
1260 ) -> crate::device::NetDeviceResult {
1261 Err(NetDeviceError::Again)
1262 }
1263 fn receive(
1264 &mut self,
1265 ) -> crate::device::NetDeviceResult<crate::device::ProtocolEthernetFrame> {
1266 Err(NetDeviceError::Again)
1267 }
1268 }
1269
1270 fn test_device_handle(device: Box<dyn Device>) -> DeviceHandle {
1271 DeviceHandle::new(IF0, device)
1272 }
1273
1274 fn ipv4_cidr(addr: Ipv4Address, prefix_len: u8) -> IpCidr {
1275 Ipv4Cidr::new(addr, prefix_len).into()
1276 }
1277
1278 #[test]
1279 fn route_lookup_uses_longest_prefix() {
1280 let mut table = RouteTable::new();
1281 table.add_rule(Rule::new(
1282 ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
1283 Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
1284 0,
1285 IF0,
1286 SRC0,
1287 100,
1288 ));
1289 table.add_rule(Rule::new(
1290 ipv4_cidr(Ipv4Address::new(10, 0, 1, 0), 24),
1291 None,
1292 1,
1293 IF1,
1294 SRC1,
1295 200,
1296 ));
1297
1298 let route = table
1299 .select_route_if(&IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 99)), |_| true)
1300 .unwrap();
1301 assert_eq!(route.dev, 1);
1302 assert_eq!(route.interface_id, IF1);
1303 assert_eq!(route.source, SRC1);
1304 assert_eq!(
1305 route.next_hop,
1306 IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 99))
1307 );
1308 }
1309
1310 #[test]
1311 fn transient_tx_backpressure_keeps_the_router_packet_queued() {
1312 let table = Arc::new(RwLock::new(RouteTable::new()));
1313 let mut router = Router::new(Arc::clone(&table));
1314 router.add_device(IF0, Box::new(RetryDevice));
1315 router.add_rule(Rule::new(
1316 ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
1317 Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
1318 0,
1319 IF0,
1320 SRC0,
1321 100,
1322 ));
1323 router
1324 .transmit(Instant::from_millis(0))
1325 .expect("the empty router TX queue has capacity")
1326 .consume(20, |packet| {
1327 packet[0] = 0x45;
1328 packet[2..4].copy_from_slice(&20u16.to_be_bytes());
1329 packet[12..16].copy_from_slice(&[10, 0, 0, 2]);
1330 packet[16..20].copy_from_slice(&[198, 51, 100, 1]);
1331 });
1332
1333 let mut sockets = SocketSet::new(vec![]);
1334 assert!(!router.dispatch(Instant::from_millis(0), &mut sockets));
1335 assert_eq!(router.tx_buffer.len(), 1);
1336 assert_eq!(router.tx_buffer.get_allocated(0, 1)[0].as_bytes().len(), 20);
1337 assert_eq!(router.devices[0].stats().tx_packets, 0);
1338 assert_eq!(router.devices[0].stats().tx_dropped, 0);
1339 }
1340
1341 #[test]
1342 fn drained_tx_queue_reuses_packet_storage() {
1343 use smoltcp::phy::RxToken as _;
1344
1345 let mut router = Router::new(Arc::new(RwLock::new(RouteTable::new())));
1346 router.add_device(InterfaceId::LOOPBACK, Box::new(EmptyDevice));
1347 router.add_rule(Rule::new(
1348 ipv4_cidr(Ipv4Address::LOCALHOST, 8),
1349 None,
1350 0,
1351 InterfaceId::LOOPBACK,
1352 IpAddress::Ipv4(Ipv4Address::LOCALHOST),
1353 0,
1354 ));
1355 let now = Instant::from_millis(0);
1356 let mut sockets = SocketSet::new(vec![]);
1357 let mut first_slot = None;
1358
1359 for len in [64, STANDARD_MTU, 64] {
1362 let mut packet = vec![0; len];
1363 packet[0] = 0x45;
1364 packet[2..4].copy_from_slice(&(len as u16).to_be_bytes());
1365 packet[12..16].copy_from_slice(&[127, 0, 0, 1]);
1366 packet[16..20].copy_from_slice(&[127, 0, 0, 1]);
1367 let address = router.transmit(now).unwrap().consume(len, |dst| {
1368 dst.copy_from_slice(&packet);
1369 dst.as_ptr() as usize
1370 });
1371 assert_eq!(
1372 address,
1373 *first_slot.get_or_insert(address),
1374 "a drained TX queue must reuse the first packet's storage"
1375 );
1376 assert!(router.dispatch(now, &mut sockets));
1377 let (rx, _tx) = router.receive(now).unwrap();
1378 rx.consume(|received| assert_eq!(received, packet));
1379 }
1380 assert!(router.receive(now).is_none());
1381 }
1382
1383 #[test]
1384 fn transmit_token_survives_backpressure_and_payload_wrap() {
1385 check_tx_token_after_payload_wrap(false);
1386 }
1387
1388 #[test]
1389 fn receive_token_survives_backpressure_and_payload_wrap() {
1390 check_tx_token_after_payload_wrap(true);
1391 }
1392
1393 fn check_tx_token_after_payload_wrap(reply_to_rx: bool) {
1394 use ax_sync::SpinLock;
1395 use smoltcp::phy::RxToken as _;
1396
1397 #[derive(Default)]
1398 struct TxProbe {
1399 allowance: usize,
1400 packets: Vec<Vec<u8>>,
1401 }
1402
1403 struct BackpressureDevice(Arc<SpinLock<TxProbe>>);
1404
1405 impl Device for BackpressureDevice {
1406 fn name(&self) -> &str {
1407 "backpressure"
1408 }
1409
1410 fn recv(
1411 &mut self,
1412 _: InterfaceId,
1413 _: &mut PacketBuffer<InterfaceId>,
1414 _: Instant,
1415 _: &mut dyn FnMut(&[u8]),
1416 ) -> usize {
1417 0
1418 }
1419
1420 fn send(&mut self, _: IpAddress, _: &[u8], _: Instant) -> usize {
1421 panic!("dispatch must use the fallible TX contract")
1422 }
1423
1424 fn try_send(
1425 &mut self,
1426 _: IpAddress,
1427 packet: &[u8],
1428 _: Instant,
1429 ) -> crate::device::NetDeviceResult<usize> {
1430 let mut probe = self.0.lock_irqsave();
1431 if probe.allowance == 0 {
1432 return Err(NetDeviceError::Again);
1433 }
1434 probe.allowance -= 1;
1435 probe.packets.push(packet.to_vec());
1436 Ok(packet.len())
1437 }
1438 }
1439
1440 fn packet(len: usize, id: u8) -> Vec<u8> {
1441 let mut packet = vec![id; len];
1442 packet[0] = 0x45;
1443 packet[2..4].copy_from_slice(&(len as u16).to_be_bytes());
1444 packet[12..16].copy_from_slice(&[10, 0, 0, 2]);
1445 packet[16..20].copy_from_slice(&[198, 51, 100, 1]);
1446 packet
1447 }
1448
1449 let mut router = Router::new(Arc::new(RwLock::new(RouteTable::new())));
1450 let probe = Arc::new(SpinLock::new(TxProbe::default()));
1451 router.add_device(IF0, Box::new(BackpressureDevice(Arc::clone(&probe))));
1452 router.add_rule(Rule::new(
1453 ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
1454 Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
1455 0,
1456 IF0,
1457 SRC0,
1458 100,
1459 ));
1460 let now = Instant::from_millis(0);
1461 let mut sockets = SocketSet::new(vec![]);
1462 let mut expected = vec![];
1463
1464 for id in 0..2 {
1466 let packet = packet(20, id);
1467 router.transmit(now).unwrap().consume(packet.len(), |dst| {
1468 dst.copy_from_slice(&packet);
1469 });
1470 expected.push(packet);
1471 }
1472 probe.lock_irqsave().allowance = 1;
1473 assert!(router.dispatch(now, &mut sockets));
1474 assert_eq!(probe.lock_irqsave().packets, expected[..1]);
1475
1476 for id in 0..SOCKET_BUFFER_SIZE - 1 {
1477 let packet = packet(STANDARD_MTU, id as u8);
1478 router.transmit(now).unwrap().consume(packet.len(), |dst| {
1479 dst.copy_from_slice(&packet);
1480 });
1481 expected.push(packet);
1482 assert!(!router.dispatch(now, &mut sockets));
1483 }
1484 assert!(router.transmit(now).is_none());
1485
1486 let incoming = packet(20, 0);
1487 router
1488 .rx_buffer
1489 .enqueue(incoming.len(), rx_metadata(IF0, &incoming))
1490 .unwrap()
1491 .copy_from_slice(&incoming);
1492 assert!(router.receive(now).is_none());
1493
1494 probe.lock_irqsave().allowance = 1;
1497 assert!(router.dispatch(now, &mut sockets));
1498 assert_eq!(probe.lock_irqsave().packets, expected[..2]);
1499 let token = if reply_to_rx {
1500 let (rx, tx) = router.receive(now).unwrap();
1501 rx.consume(|packet| assert_eq!(packet, incoming));
1502 tx
1503 } else {
1504 router.transmit(now).unwrap()
1505 };
1506 let final_packet = packet(STANDARD_MTU, 0xfe);
1507 assert_eq!(
1508 token.consume(final_packet.len(), |dst| {
1509 dst.copy_from_slice(&final_packet);
1510 42
1511 }),
1512 42
1513 );
1514 expected.push(final_packet);
1515
1516 probe.lock_irqsave().allowance = expected.len();
1517 assert!(router.dispatch(now, &mut sockets));
1518 assert_eq!(probe.lock_irqsave().packets, expected);
1519 assert!(router.transmit(now).is_some());
1520 assert!(!router.dispatch(now, &mut sockets));
1521 assert_eq!(router.devices[0].stats().tx_packets, expected.len() as u64);
1522 assert_eq!(router.devices[0].stats().tx_errors, 0);
1523 assert_eq!(router.devices[0].stats().tx_dropped, 0);
1524 }
1525
1526 #[test]
1527 fn fanout_retries_only_blocked_ports_without_repeating_accepted_packets() {
1528 use ax_sync::SpinLock;
1529
1530 #[derive(Default)]
1531 struct TxProbe {
1532 failures: VecDeque<NetDeviceError>,
1533 attempts: usize,
1534 packets: Vec<Vec<u8>>,
1535 }
1536
1537 struct FanoutDevice(Arc<SpinLock<TxProbe>>);
1538
1539 impl Device for FanoutDevice {
1540 fn name(&self) -> &str {
1541 "fanout"
1542 }
1543
1544 fn recv(
1545 &mut self,
1546 _: InterfaceId,
1547 _: &mut PacketBuffer<InterfaceId>,
1548 _: Instant,
1549 _: &mut dyn FnMut(&[u8]),
1550 ) -> usize {
1551 0
1552 }
1553
1554 fn send(&mut self, _: IpAddress, _: &[u8], _: Instant) -> usize {
1555 panic!("fanout must preserve the fallible TX contract")
1556 }
1557
1558 fn try_send(
1559 &mut self,
1560 _: IpAddress,
1561 packet: &[u8],
1562 _: Instant,
1563 ) -> crate::device::NetDeviceResult<usize> {
1564 let mut probe = self.0.lock_irqsave();
1565 probe.attempts += 1;
1566 if let Some(error) = probe.failures.pop_front() {
1567 return Err(error);
1568 }
1569 probe.packets.push(packet.to_vec());
1570 Ok(packet.len())
1571 }
1572 }
1573
1574 for ipv6 in [false, true] {
1575 let mut packet = if ipv6 { vec![0u8; 40] } else { vec![0u8; 20] };
1576 if ipv6 {
1577 packet[0] = 0x60;
1578 packet[24] = 0xff;
1579 packet[25] = 2;
1580 packet[39] = 1;
1581 } else {
1582 packet[0] = 0x45;
1583 packet[2..4].copy_from_slice(&20u16.to_be_bytes());
1584 packet[16..20].fill(0xff);
1585 }
1586 let mut router = Router::new(Arc::new(RwLock::new(RouteTable::new())));
1587 let probes: Vec<_> = [
1588 vec![],
1589 vec![NetDeviceError::Again, NetDeviceError::Again],
1590 vec![NetDeviceError::Again],
1591 vec![NetDeviceError::Io],
1592 vec![],
1593 ]
1594 .into_iter()
1595 .enumerate()
1596 .map(|(index, failures)| {
1597 let probe = Arc::new(SpinLock::new(TxProbe {
1598 failures: failures.into(),
1599 ..TxProbe::default()
1600 }));
1601 let id = if index == 4 {
1602 InterfaceId::LOOPBACK
1603 } else {
1604 InterfaceId::new(index as u32 + 2)
1605 };
1606 router.add_device(id, Box::new(FanoutDevice(Arc::clone(&probe))));
1607 probe
1608 })
1609 .collect();
1610 let mut next_packet = packet.clone();
1611 next_packet[1] = 1;
1612 for queued in [&packet, &next_packet] {
1613 router
1614 .transmit(Instant::from_millis(0))
1615 .unwrap()
1616 .consume(queued.len(), |dst| dst.copy_from_slice(queued));
1617 }
1618 let mut sockets = SocketSet::new(vec![]);
1619 for completed_port in [0, 2] {
1620 assert!(router.dispatch(Instant::from_millis(0), &mut sockets));
1621 assert_eq!(router.tx_buffer.get_allocated(0, 1)[0].as_bytes(), packet);
1622 assert_eq!(router.tx_buffer.len(), 2);
1623 assert_eq!(
1624 probes[completed_port].lock_irqsave().packets,
1625 vec![packet.clone()]
1626 );
1627 assert_eq!(probes[0].lock_irqsave().attempts, 1);
1628 assert_eq!(probes[3].lock_irqsave().attempts, 1);
1629 assert_eq!(router.devices[3].stats().tx_errors, 1);
1630 }
1631 probes[1]
1632 .lock_irqsave()
1633 .failures
1634 .push_back(NetDeviceError::Again);
1635 assert!(!router.dispatch(Instant::from_millis(0), &mut sockets));
1636 assert_eq!(router.tx_buffer.get_allocated(0, 1)[0].as_bytes(), packet);
1637 assert_eq!(probes[0].lock_irqsave().attempts, 1);
1638 assert_eq!(probes[2].lock_irqsave().attempts, 2);
1639 assert!(router.dispatch(Instant::from_millis(0), &mut sockets));
1640 assert!(router.tx_buffer.is_empty());
1641 for (index, attempts) in [2, 5, 3, 2].into_iter().enumerate() {
1642 let probe = probes[index].lock_irqsave();
1643 assert_eq!(probe.attempts, attempts);
1644 let expected = if index == 3 {
1645 vec![next_packet.clone()]
1646 } else {
1647 vec![packet.clone(), next_packet.clone()]
1648 };
1649 assert_eq!(probe.packets, expected);
1650 assert_eq!(
1651 router.devices[index].stats().tx_packets,
1652 expected.len() as u64
1653 );
1654 assert_eq!(router.devices[index].stats().tx_dropped, 0);
1655 }
1656 assert_eq!(probes[4].lock_irqsave().attempts, 0);
1657 }
1658 }
1659
1660 #[test]
1661 fn router_keeps_software_checksums_even_with_offload_capable_devices() {
1662 let table = Arc::new(RwLock::new(RouteTable::new()));
1663 let mut router = Router::new(table);
1664 router.add_device(
1665 IF0,
1666 Box::new(crate::device::EthernetDevice::new(
1667 "checksum".into(),
1668 Box::new(ChecksumPort),
1669 None,
1670 )),
1671 );
1672
1673 let caps = smoltcp::phy::Device::capabilities(&router);
1674 assert!(caps.checksum.tcp.rx());
1676 assert!(caps.checksum.tcp.tx());
1677 assert!(caps.checksum.udp.rx());
1678 assert!(caps.checksum.udp.tx());
1679
1680 router.add_device(IF1, Box::new(EmptyDevice));
1681 let caps = smoltcp::phy::Device::capabilities(&router);
1682 assert!(caps.checksum.tcp.tx());
1683 assert!(caps.checksum.udp.tx());
1684 }
1685
1686 #[test]
1687 fn route_lookup_uses_metric_for_same_prefix() {
1688 let mut table = RouteTable::new();
1689 let dst = IpAddress::Ipv4(Ipv4Address::new(203, 0, 113, 10));
1690 table.add_rule(Rule::new(
1691 ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
1692 Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
1693 0,
1694 IF0,
1695 SRC0,
1696 200,
1697 ));
1698 table.add_rule(Rule::new(
1699 ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
1700 Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 1))),
1701 1,
1702 IF1,
1703 SRC1,
1704 100,
1705 ));
1706
1707 let route = table.select_route_if(&dst, |_| true).unwrap();
1708 assert_eq!(route.interface_id, IF1);
1709 assert_eq!(route.metric, 100);
1710 assert_eq!(
1711 route.next_hop,
1712 IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 1))
1713 );
1714 }
1715
1716 #[test]
1717 fn route_lookup_keeps_stable_order_for_equal_metric() {
1718 let mut table = RouteTable::new();
1719 let dst = IpAddress::Ipv4(Ipv4Address::new(203, 0, 113, 10));
1720 table.add_rule(Rule::new(
1721 ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
1722 Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
1723 0,
1724 IF0,
1725 SRC0,
1726 100,
1727 ));
1728 table.add_rule(Rule::new(
1729 ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
1730 Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 1))),
1731 1,
1732 IF1,
1733 SRC1,
1734 100,
1735 ));
1736
1737 let route = table.select_route_if(&dst, |_| true).unwrap();
1738 assert_eq!(route.interface_id, IF0);
1739 assert_eq!(
1740 route.next_hop,
1741 IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))
1742 );
1743 }
1744
1745 #[test]
1746 fn route_lookup_skips_unusable_interface() {
1747 let mut table = RouteTable::new();
1748 let dst = IpAddress::Ipv4(Ipv4Address::new(203, 0, 113, 10));
1749 table.add_rule(Rule::new(
1750 ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
1751 Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
1752 0,
1753 IF0,
1754 SRC0,
1755 100,
1756 ));
1757 table.add_rule(Rule::new(
1758 ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
1759 Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 1))),
1760 1,
1761 IF1,
1762 SRC1,
1763 200,
1764 ));
1765
1766 let route = table
1767 .select_route_if(&dst, |interface_id| interface_id != IF0)
1768 .unwrap();
1769 assert_eq!(route.interface_id, IF1);
1770 }
1771
1772 #[test]
1773 fn snoop_tcp_packet_drops_truncated_ip_and_tcp_headers() {
1774 const IPV4_HEADER_LEN: usize = 20;
1775 const IPV6_HEADER_LEN: usize = 40;
1776 const TCP_HEADER_LEN: usize = 20;
1777
1778 let mut sockets = SocketSet::new(vec![]);
1779
1780 let mut ipv4_tcp = [0u8; IPV4_HEADER_LEN + TCP_HEADER_LEN];
1781 ipv4_tcp[0] = 0x45;
1782 let ipv4_tcp_len = ipv4_tcp.len() as u16;
1783 ipv4_tcp[2..4].copy_from_slice(&ipv4_tcp_len.to_be_bytes());
1784 ipv4_tcp[9] = IpProtocol::Tcp.into();
1785 for len in 0..ipv4_tcp.len() {
1786 snoop_tcp_packet(&ipv4_tcp[..len], &mut sockets);
1787 }
1788
1789 let mut ipv6_tcp = [0u8; IPV6_HEADER_LEN + TCP_HEADER_LEN];
1790 ipv6_tcp[0] = 0x60;
1791 ipv6_tcp[4..6].copy_from_slice(&20u16.to_be_bytes());
1792 ipv6_tcp[6] = IpProtocol::Tcp.into();
1793 for len in 0..ipv6_tcp.len() {
1794 snoop_tcp_packet(&ipv6_tcp[..len], &mut sockets);
1795 }
1796
1797 for tcp_len in 0..TCP_HEADER_LEN {
1801 let mut ipv4_tcp = vec![0u8; IPV4_HEADER_LEN + tcp_len];
1802 ipv4_tcp[0] = 0x45;
1803 let ipv4_len = ipv4_tcp.len() as u16;
1804 ipv4_tcp[2..4].copy_from_slice(&ipv4_len.to_be_bytes());
1805 ipv4_tcp[9] = IpProtocol::Tcp.into();
1806 snoop_tcp_packet(&ipv4_tcp, &mut sockets);
1807
1808 let mut ipv6_tcp = vec![0u8; IPV6_HEADER_LEN + tcp_len];
1809 ipv6_tcp[0] = 0x60;
1810 ipv6_tcp[4..6].copy_from_slice(&(tcp_len as u16).to_be_bytes());
1811 ipv6_tcp[6] = IpProtocol::Tcp.into();
1812 snoop_tcp_packet(&ipv6_tcp, &mut sockets);
1813 }
1814 }
1815
1816 #[test]
1817 fn default_routes_only_reports_zero_prefix_ipv4_rules() {
1818 let mut table = RouteTable::new();
1819 table.add_rule(Rule::new(
1820 ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
1821 Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
1822 0,
1823 IF0,
1824 SRC0,
1825 100,
1826 ));
1827 table.add_rule(Rule::new(
1828 ipv4_cidr(Ipv4Address::new(10, 0, 1, 0), 24),
1829 None,
1830 1,
1831 IF1,
1832 SRC1,
1833 100,
1834 ));
1835
1836 let routes = table.default_routes();
1837 assert_eq!(routes.len(), 1);
1838 assert_eq!(routes[0].interface_id, IF0);
1839 }
1840
1841 #[test]
1852 fn no_route_does_not_count_interface_tx_dropped() {
1853 use smoltcp::{iface::SocketSet, storage::PacketMetadata};
1854
1855 let dev0 = test_device_handle(Box::new(EmptyDevice));
1857 let dev1 = DeviceHandle::new(IF1, Box::new(EmptyDevice));
1858 let mut devices = vec![dev0, dev1];
1859
1860 let mut route_table = RouteTable::new();
1863 route_table.add_rule(Rule::new(
1864 ipv4_cidr(Ipv4Address::new(10, 0, 0, 0), 24),
1865 Some(SRC0),
1866 0,
1867 IF0, SRC0,
1869 100,
1870 ));
1871 let shared_table: SharedRouteTable = Arc::new(RwLock::new(route_table));
1872
1873 let mut rx_buffer: RouterPacketBuffer = PacketBuffer::new(
1874 vec![PacketMetadata::EMPTY; 1],
1875 vec![0u8; super::STANDARD_MTU],
1876 );
1877 let mut sockets = SocketSet::new(vec![]);
1878
1879 let src_addr = SRC0;
1880 let dst_addr = IpAddress::Ipv4(Ipv4Address::new(203, 0, 113, 10));
1881 let packet = [0u8; 64];
1882
1883 let before: Vec<_> = devices.iter().map(|d| d.stats()).collect();
1884
1885 let outcome = dispatch_unicast_packet(
1886 &mut rx_buffer,
1887 &mut devices,
1888 &shared_table,
1889 src_addr,
1890 dst_addr,
1891 &packet,
1892 &mut sockets,
1893 );
1894
1895 assert_eq!(
1896 outcome,
1897 DispatchOutcome::Consumed(false),
1898 "no-route dispatch must consume the packet without scheduling work"
1899 );
1900
1901 for (i, dev) in devices.iter().enumerate() {
1902 let snap = dev.stats();
1903 assert_eq!(
1904 snap.tx_dropped, before[i].tx_dropped,
1905 "device {i} tx_dropped changed from {} to {} after no-route dispatch",
1906 before[i].tx_dropped, snap.tx_dropped,
1907 );
1908 }
1909 }
1910}
1911
1912#[cfg(test)]
1913mod l2_counter_tests {
1914 use smoltcp::{
1915 storage::{PacketBuffer, PacketMetadata},
1916 time::Instant,
1917 wire::{IpAddress, Ipv4Address},
1918 };
1919
1920 use super::*;
1921
1922 const IF0: InterfaceId = InterfaceId::new(2);
1923
1924 struct CountingMockDevice {
1926 name: &'static str,
1927 send_returns: usize,
1928 recv_returns: usize,
1929 deferred_tx_lens: Vec<usize>,
1931 deferred_rx_lens: Vec<usize>,
1933 }
1934
1935 impl Device for CountingMockDevice {
1936 fn name(&self) -> &str {
1937 self.name
1938 }
1939
1940 fn recv(
1941 &mut self,
1942 _interface_id: InterfaceId,
1943 _buffer: &mut PacketBuffer<InterfaceId>,
1944 _timestamp: Instant,
1945 _snoop: &mut dyn FnMut(&[u8]),
1946 ) -> usize {
1947 self.recv_returns
1948 }
1949
1950 fn send(&mut self, _next_hop: IpAddress, _packet: &[u8], _timestamp: Instant) -> usize {
1951 self.send_returns
1952 }
1953
1954 fn drain_deferred_tx(&mut self) -> Vec<usize> {
1955 core::mem::take(&mut self.deferred_tx_lens)
1956 }
1957
1958 fn drain_deferred_rx(&mut self) -> Vec<usize> {
1959 core::mem::take(&mut self.deferred_rx_lens)
1960 }
1961 }
1962
1963 fn test_device_handle(device: Box<dyn Device>) -> DeviceHandle {
1964 DeviceHandle::new(IF0, device)
1965 }
1966
1967 fn test_ip() -> IpAddress {
1968 IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))
1969 }
1970
1971 fn test_packet_buffer() -> PacketBuffer<'static, InterfaceId> {
1972 PacketBuffer::new(
1973 vec![PacketMetadata::EMPTY; 1],
1974 vec![0u8; super::STANDARD_MTU],
1975 )
1976 }
1977
1978 #[test]
1981 fn count_rx_accumulates_bytes_and_packets() {
1982 let device = test_device_handle(Box::new(CountingMockDevice {
1983 name: "mock",
1984 send_returns: 0,
1985 deferred_tx_lens: vec![],
1986 deferred_rx_lens: vec![],
1987 recv_returns: 0,
1988 }));
1989
1990 device.count_rx(100);
1991 assert_eq!(device.stats().rx_bytes, 100);
1992 assert_eq!(device.stats().rx_packets, 1);
1993
1994 device.count_rx(200);
1995 assert_eq!(device.stats().rx_bytes, 300);
1996 assert_eq!(device.stats().rx_packets, 2);
1997 }
1998
1999 #[test]
2000 fn count_tx_accumulates_bytes_and_packets() {
2001 let device = test_device_handle(Box::new(CountingMockDevice {
2002 name: "mock",
2003 send_returns: 0,
2004 deferred_tx_lens: vec![],
2005 deferred_rx_lens: vec![],
2006 recv_returns: 0,
2007 }));
2008
2009 device.count_tx(64);
2010 assert_eq!(device.stats().tx_bytes, 64);
2011 assert_eq!(device.stats().tx_packets, 1);
2012
2013 device.count_tx(1500);
2014 assert_eq!(device.stats().tx_bytes, 1564);
2015 assert_eq!(device.stats().tx_packets, 2);
2016 }
2017
2018 #[test]
2021 fn stats_reflects_current_counters_after_counting() {
2022 let device = test_device_handle(Box::new(CountingMockDevice {
2023 name: "mock",
2024 send_returns: 0,
2025 deferred_tx_lens: vec![],
2026 deferred_rx_lens: vec![],
2027 recv_returns: 0,
2028 }));
2029
2030 device.count_rx(100);
2031 device.count_tx(64);
2032
2033 let snap = device.stats();
2034 assert_eq!(snap.rx_bytes, 100);
2035 assert_eq!(snap.rx_packets, 1);
2036 assert_eq!(snap.tx_bytes, 64);
2037 assert_eq!(snap.tx_packets, 1);
2038 }
2039
2040 #[test]
2043 fn send_returns_frame_len_tx_counts_l2_not_ip_payload() {
2044 let mut device = test_device_handle(Box::new(CountingMockDevice {
2045 name: "mock",
2046 send_returns: 1514, deferred_tx_lens: vec![],
2048 deferred_rx_lens: vec![],
2049 recv_returns: 0,
2050 }));
2051
2052 let frame_len = device
2054 .inner
2055 .send(test_ip(), &[0u8; 100], Instant::from_millis(0));
2056 assert_eq!(frame_len, 1514);
2057 if frame_len > 0 {
2058 device.count_tx(frame_len);
2059 }
2060
2061 let snap = device.stats();
2062 assert_eq!(snap.tx_bytes, 1514);
2064 assert_eq!(snap.tx_packets, 1);
2065 }
2066
2067 #[test]
2068 fn send_returns_zero_no_tx_counted() {
2069 let mut device = test_device_handle(Box::new(CountingMockDevice {
2070 name: "mock",
2071 send_returns: 0, deferred_tx_lens: vec![],
2073 deferred_rx_lens: vec![],
2074 recv_returns: 0,
2075 }));
2076
2077 let frame_len = device
2078 .inner
2079 .send(test_ip(), &[0u8; 100], Instant::from_millis(0));
2080 assert_eq!(frame_len, 0);
2081 if frame_len > 0 {
2083 device.count_tx(frame_len);
2084 }
2085
2086 let snap = device.stats();
2087 assert_eq!(snap.tx_bytes, 0);
2088 assert_eq!(snap.tx_packets, 0);
2089 }
2090
2091 #[test]
2094 fn recv_returns_frame_len_rx_counts_it() {
2095 let mut device = test_device_handle(Box::new(CountingMockDevice {
2096 name: "mock",
2097 send_returns: 0,
2098 deferred_tx_lens: vec![],
2099 deferred_rx_lens: vec![],
2100 recv_returns: 1514,
2101 }));
2102
2103 let frame_len = device.inner.recv(
2104 IF0,
2105 &mut test_packet_buffer(),
2106 Instant::from_millis(0),
2107 &mut |_| {},
2108 );
2109 assert_eq!(frame_len, 1514);
2110 if frame_len > 0 {
2111 device.count_rx(frame_len);
2112 }
2113
2114 let snap = device.stats();
2115 assert_eq!(snap.rx_bytes, 1514);
2116 assert_eq!(snap.rx_packets, 1);
2117 }
2118
2119 #[test]
2120 fn recv_returns_zero_no_rx_counted() {
2121 let mut device = test_device_handle(Box::new(CountingMockDevice {
2122 name: "mock",
2123 send_returns: 0,
2124 deferred_tx_lens: vec![],
2125 deferred_rx_lens: vec![],
2126 recv_returns: 0, }));
2128
2129 let frame_len = device.inner.recv(
2130 IF0,
2131 &mut test_packet_buffer(),
2132 Instant::from_millis(0),
2133 &mut |_| {},
2134 );
2135 assert_eq!(frame_len, 0);
2136 if frame_len > 0 {
2137 device.count_rx(frame_len);
2138 }
2139
2140 let snap = device.stats();
2141 assert_eq!(snap.rx_bytes, 0);
2142 assert_eq!(snap.rx_packets, 0);
2143 }
2144
2145 #[test]
2151 fn protocol_executor_three_path_combined_drain() {
2152 let mut device = test_device_handle(Box::new(CountingMockDevice {
2153 name: "mock",
2154 send_returns: 0,
2155 deferred_tx_lens: vec![60, 60], deferred_rx_lens: vec![42], recv_returns: 1514, }));
2159
2160 let frame_len = device.inner.recv(
2165 IF0,
2166 &mut test_packet_buffer(),
2167 Instant::from_millis(0),
2168 &mut |_| {},
2169 );
2170 if frame_len > 0 {
2171 device.count_rx(frame_len);
2172 }
2173 for len in device.inner.drain_deferred_tx() {
2174 device.count_tx(len);
2175 }
2176 for len in device.inner.drain_deferred_rx() {
2177 device.count_rx(len);
2178 }
2179
2180 let snap = device.stats();
2181 assert_eq!(snap.rx_packets, 2);
2183 assert_eq!(snap.rx_bytes, 1556);
2184 assert_eq!(snap.tx_packets, 2);
2186 assert_eq!(snap.tx_bytes, 120);
2187 }
2188}