use alloc::{
boxed::Box,
collections::VecDeque,
string::{String, ToString},
sync::Arc,
vec,
vec::Vec,
};
use core::sync::atomic::{AtomicU64, Ordering};
use ax_hal::time::{NANOS_PER_MICROS, monotonic_time_nanos};
use ax_sync::SpinRwLock as RwLock;
use smoltcp::{
iface::SocketSet,
phy::{DeviceCapabilities, Medium, PacketMeta},
storage::{PacketMetadata, RingBuffer},
time::Instant,
wire::{
IpAddress, IpCidr, IpProtocol, IpVersion, Ipv4Address, Ipv4Cidr, Ipv4Packet, Ipv6Packet,
TcpPacket,
},
};
use crate::{
LISTEN_TABLE,
config::{DeviceBinding, InterfaceId, RouteInfo},
consts::{SOCKET_BUFFER_SIZE, STANDARD_MTU},
device::{ArpEntry, Device, DeviceRxPacket, DeviceRxPoll, NetDeviceError},
ip_tos::apply_egress_ip_tos,
rx_meta::packet_meta_for_rx_packet,
};
const DEVICE_RX_WORKER_BATCH: usize = 16;
#[derive(Debug, Clone)]
pub struct NetDevStats {
pub interface_id: InterfaceId,
pub name: String,
pub rx_bytes: u64,
pub rx_packets: u64,
pub rx_errors: u64,
pub rx_dropped: u64,
pub tx_bytes: u64,
pub tx_packets: u64,
pub tx_errors: u64,
pub tx_dropped: u64,
}
#[derive(Debug)]
pub struct Rule {
pub filter: IpCidr,
pub via: Option<IpAddress>,
pub dev: usize,
pub interface_id: InterfaceId,
pub src: IpAddress,
pub metric: u32,
pub order: u64,
}
impl Rule {
pub fn new(
filter: IpCidr,
via: Option<IpAddress>,
dev: usize,
interface_id: InterfaceId,
src: IpAddress,
metric: u32,
) -> Self {
Self {
filter,
via,
dev,
interface_id,
src,
metric,
order: 0,
}
}
fn to_info(&self) -> RouteInfo {
RouteInfo {
filter: self.filter,
via: self.via,
interface_id: self.interface_id,
source: self.src,
metric: self.metric,
}
}
}
#[derive(Debug, Clone, Copy)]
struct RxMetadata {
interface_id: InterfaceId,
packet_meta: PacketMeta,
}
type RouterPacketBuffer = smoltcp::storage::PacketBuffer<'static, RxMetadata>;
type DevicePacketBuffer = smoltcp::storage::PacketBuffer<'static, InterfaceId>;
#[derive(Clone)]
struct TxPacket {
len: usize,
bytes: [u8; STANDARD_MTU],
}
impl TxPacket {
fn as_bytes(&self) -> &[u8] {
&self.bytes[..self.len]
}
}
struct OwnedRxPacket {
metadata: RxMetadata,
packet: DeviceRxPacket,
}
fn rx_metadata(interface_id: InterfaceId, packet: &[u8]) -> RxMetadata {
RxMetadata {
interface_id,
packet_meta: packet_meta_for_rx_packet(packet),
}
}
struct DeviceHandle {
interface_id: InterfaceId,
name: String,
inner: Box<dyn Device>,
rx_buffer: DevicePacketBuffer,
rx_bytes: AtomicU64,
rx_packets: AtomicU64,
rx_errors: AtomicU64,
rx_dropped: AtomicU64,
tx_bytes: AtomicU64,
tx_packets: AtomicU64,
tx_errors: AtomicU64,
tx_dropped: AtomicU64,
}
impl DeviceHandle {
fn new(interface_id: InterfaceId, device: Box<dyn Device>) -> Self {
let name = device.name().to_string();
Self {
interface_id,
name,
inner: device,
rx_buffer: DevicePacketBuffer::new(
vec![PacketMetadata::EMPTY; DEVICE_RX_WORKER_BATCH],
vec![0u8; STANDARD_MTU * DEVICE_RX_WORKER_BATCH],
),
rx_bytes: AtomicU64::new(0),
rx_packets: AtomicU64::new(0),
rx_errors: AtomicU64::new(0),
rx_dropped: AtomicU64::new(0),
tx_bytes: AtomicU64::new(0),
tx_packets: AtomicU64::new(0),
tx_errors: AtomicU64::new(0),
tx_dropped: AtomicU64::new(0),
}
}
fn count_rx(&self, len: usize) {
self.rx_bytes.fetch_add(len as u64, Ordering::Relaxed);
self.rx_packets.fetch_add(1, Ordering::Relaxed);
}
fn count_tx(&self, len: usize) {
self.tx_bytes.fetch_add(len as u64, Ordering::Relaxed);
self.tx_packets.fetch_add(1, Ordering::Relaxed);
}
fn count_rx_errors(&self, n: u64) {
self.rx_errors.fetch_add(n, Ordering::Relaxed);
}
fn count_rx_dropped(&self, n: u64) {
self.rx_dropped.fetch_add(n, Ordering::Relaxed);
}
fn count_tx_errors(&self, n: u64) {
self.tx_errors.fetch_add(n, Ordering::Relaxed);
}
fn count_tx_dropped(&self, n: u64) {
self.tx_dropped.fetch_add(n, Ordering::Relaxed);
}
fn drain_device_counters(&mut self) {
for len in self.inner.drain_deferred_tx() {
self.count_tx(len);
}
for len in self.inner.drain_deferred_rx() {
self.count_rx(len);
}
let n = self.inner.drain_deferred_tx_errors();
if n > 0 {
self.count_tx_errors(n);
}
let n = self.inner.drain_deferred_tx_drops();
if n > 0 {
self.count_tx_dropped(n);
}
let n = self.inner.drain_deferred_rx_errors();
if n > 0 {
self.count_rx_errors(n);
}
let n = self.inner.drain_deferred_rx_drops();
if n > 0 {
self.count_rx_dropped(n);
}
}
fn stats(&self) -> NetDevStats {
NetDevStats {
interface_id: self.interface_id,
name: self.name.clone(),
rx_bytes: self.rx_bytes.load(Ordering::Relaxed),
rx_packets: self.rx_packets.load(Ordering::Relaxed),
rx_errors: self.rx_errors.load(Ordering::Relaxed),
rx_dropped: self.rx_dropped.load(Ordering::Relaxed),
tx_bytes: self.tx_bytes.load(Ordering::Relaxed),
tx_packets: self.tx_packets.load(Ordering::Relaxed),
tx_errors: self.tx_errors.load(Ordering::Relaxed),
tx_dropped: self.tx_dropped.load(Ordering::Relaxed),
}
}
fn send(&mut self, next_hop: IpAddress, packet: &[u8], timestamp: Instant) -> bool {
match self.try_send(next_hop, packet, timestamp) {
Ok(consumed) => consumed,
Err(NetDeviceError::Again) => false,
Err(error) => {
warn!("{}: transmit failed: {error:?}", self.name);
self.count_tx_errors(1);
self.drain_device_counters();
false
}
}
}
fn try_send(
&mut self,
next_hop: IpAddress,
packet: &[u8],
timestamp: Instant,
) -> Result<bool, NetDeviceError> {
if packet.len() > STANDARD_MTU {
warn!(
"{}: packet to {} exceeds MTU ({} bytes), dropping",
self.name,
next_hop,
packet.len()
);
self.count_tx_dropped(1);
return Ok(false);
}
let frame_len = self.inner.try_send(next_hop, packet, timestamp)?;
if frame_len > 0 {
self.count_tx(frame_len);
}
self.drain_device_counters();
Ok(true)
}
}
fn now() -> Instant {
Instant::from_micros_const((monotonic_time_nanos() / NANOS_PER_MICROS) as i64)
}
#[derive(Debug, Clone, Copy)]
pub struct RouteDecision {
pub dev: usize,
pub interface_id: InterfaceId,
pub source: IpAddress,
pub next_hop: IpAddress,
pub metric: u32,
}
pub struct RouteTable {
rules: Vec<Rule>,
next_order: u64,
}
impl RouteTable {
pub fn new() -> Self {
Self {
rules: Vec::new(),
next_order: 0,
}
}
pub fn add_rule(&mut self, mut rule: Rule) {
rule.order = self.next_order;
self.next_order = self.next_order.saturating_add(1);
self.rules.push(rule);
self.sort_rules();
}
fn sort_rules(&mut self) {
self.rules.sort_by(|a, b| {
b.filter
.prefix_len()
.cmp(&a.filter.prefix_len())
.then_with(|| a.metric.cmp(&b.metric))
.then_with(|| a.order.cmp(&b.order))
});
}
pub fn select_route_if(
&self,
dst: &IpAddress,
mut is_usable: impl FnMut(InterfaceId) -> bool,
) -> Option<RouteDecision> {
self.rules
.iter()
.find(|rule| rule.filter.contains_addr(dst) && is_usable(rule.interface_id))
.map(|rule| RouteDecision {
dev: rule.dev,
interface_id: rule.interface_id,
source: rule.src,
next_hop: rule.via.unwrap_or(*dst),
metric: rule.metric,
})
}
pub fn select_route_for_source(
&self,
dst: &IpAddress,
source: &IpAddress,
) -> Option<RouteDecision> {
self.rules
.iter()
.find(|rule| rule.filter.contains_addr(dst) && &rule.src == source)
.map(|rule| RouteDecision {
dev: rule.dev,
interface_id: rule.interface_id,
source: rule.src,
next_hop: rule.via.unwrap_or(*dst),
metric: rule.metric,
})
}
pub fn default_routes(&self) -> Vec<RouteInfo> {
self.rules
.iter()
.filter(|rule| match rule.filter {
IpCidr::Ipv4(cidr) => {
cidr.address() == Ipv4Address::UNSPECIFIED && cidr.prefix_len() == 0
}
_ => false,
})
.map(Rule::to_info)
.collect()
}
pub fn remove_ipv4_rules_for_interface(&mut self, interface_id: InterfaceId) {
self.rules.retain(|rule| {
!matches!(
rule.filter,
IpCidr::Ipv4(_) if rule.interface_id == interface_id
)
});
}
pub fn replace_ipv4_rules_for_interface(
&mut self,
interface_id: InterfaceId,
mut new_rules: Vec<Rule>,
) {
self.remove_ipv4_rules_for_interface(interface_id);
for rule in &mut new_rules {
rule.order = self.next_order;
self.next_order = self.next_order.saturating_add(1);
}
self.rules.extend(new_rules);
self.sort_rules();
}
}
pub(crate) type SharedRouteTable = Arc<RwLock<RouteTable>>;
pub struct Router {
rx_buffer: RouterPacketBuffer,
tx_buffer: RingBuffer<'static, TxPacket>,
pending_fanout: Vec<usize>,
ready_rx: VecDeque<OwnedRxPacket>,
devices: Vec<DeviceHandle>,
table: SharedRouteTable,
}
impl Router {
pub fn new(table: SharedRouteTable) -> Self {
let rx_buffer = RouterPacketBuffer::new(
vec![PacketMetadata::EMPTY; SOCKET_BUFFER_SIZE],
vec![0u8; STANDARD_MTU * SOCKET_BUFFER_SIZE],
);
let tx_buffer = RingBuffer::new(vec![
TxPacket {
len: 0,
bytes: [0; STANDARD_MTU],
};
SOCKET_BUFFER_SIZE
]);
Self {
rx_buffer,
tx_buffer,
pending_fanout: Vec::new(),
ready_rx: VecDeque::with_capacity(SOCKET_BUFFER_SIZE),
devices: Vec::new(),
table,
}
}
pub fn add_rule(&mut self, rule: Rule) {
self.table.write().add_rule(rule);
}
pub fn add_device(&mut self, interface_id: InterfaceId, device: Box<dyn Device>) -> usize {
self.devices.push(DeviceHandle::new(interface_id, device));
self.devices.len() - 1
}
pub fn interface_id_for_dev(&self, dev: usize) -> Option<InterfaceId> {
self.devices.get(dev).map(|device| device.interface_id)
}
pub fn device_index_for_interface_id(&self, interface_id: InterfaceId) -> Option<usize> {
self.devices
.iter()
.position(|device| device.interface_id == interface_id)
}
pub fn device_names(&self) -> Vec<String> {
self.devices
.iter()
.map(|device| device.name.clone())
.collect()
}
pub fn set_ipv4_config(
&mut self,
dev: usize,
interface_id: InterfaceId,
metric: u32,
address: Option<Ipv4Cidr>,
gateway: Option<IpAddress>,
) {
let new_rules = self.ipv4_rules(dev, interface_id, metric, address, gateway);
self.table
.write()
.replace_ipv4_rules_for_interface(interface_id, new_rules);
}
pub(crate) fn ipv4_rules(
&mut self,
dev: usize,
interface_id: InterfaceId,
metric: u32,
address: Option<Ipv4Cidr>,
gateway: Option<IpAddress>,
) -> Vec<Rule> {
self.devices[dev].inner.set_ipv4_addr(address);
let mut rules = Vec::new();
if let Some(address) = address {
rules.push(Rule::new(
address.into(),
None,
dev,
interface_id,
address.address().into(),
metric,
));
if let Some(gateway) = gateway {
rules.push(Rule::new(
Ipv4Cidr::new(Ipv4Address::UNSPECIFIED, 0).into(),
Some(gateway),
dev,
interface_id,
address.address().into(),
metric,
));
}
}
rules
}
pub fn poll(
&mut self,
_timestamp: Instant,
sockets: &mut SocketSet<'_>,
mut snoop: impl FnMut(InterfaceId, &[u8]),
) -> bool {
let mut moved_rx = false;
let Router {
rx_buffer,
ready_rx,
devices,
..
} = self;
for device in devices {
if device.interface_id == InterfaceId::LOOPBACK {
continue;
}
let mut budget = DEVICE_RX_WORKER_BATCH;
while budget > 0 && ready_rx.len() < SOCKET_BUFFER_SIZE {
let interface_id = device.interface_id;
match device.inner.poll_owned_rx(now()) {
DeviceRxPoll::Packet(packet) => {
let metadata = packet.read_with(|bytes| {
snoop_tcp_packet(bytes, sockets);
snoop(interface_id, bytes);
rx_metadata(interface_id, bytes)
});
let frame_len = packet.frame_len();
ready_rx.push_back(OwnedRxPacket { metadata, packet });
device.count_rx(frame_len);
moved_rx = true;
budget -= 1;
continue;
}
DeviceRxPoll::Idle => break,
DeviceRxPoll::Unsupported => {}
}
if rx_buffer.is_full() || device.rx_buffer.is_full() {
break;
}
let mut frame_snoop = |_packet: &[u8]| {};
let direct = device.inner.recv_direct(
now(),
&mut |packet| {
snoop_tcp_packet(packet, sockets);
snoop(interface_id, packet);
let Ok(dst) =
rx_buffer.enqueue(packet.len(), rx_metadata(interface_id, packet))
else {
return false;
};
dst.copy_from_slice(packet);
true
},
&mut frame_snoop,
);
if let Some(frame_len) = direct {
if frame_len == 0 {
break;
}
device.count_rx(frame_len);
moved_rx = true;
budget -= 1;
continue;
}
let frame_len = device.inner.recv(
device.interface_id,
&mut device.rx_buffer,
now(),
&mut frame_snoop,
);
if frame_len == 0 {
break;
}
let Ok((interface_id, packet)) = device.rx_buffer.dequeue() else {
device.count_rx_errors(1);
break;
};
snoop_tcp_packet(packet, sockets);
snoop(interface_id, packet);
let Ok(dst) = rx_buffer.enqueue(packet.len(), rx_metadata(interface_id, packet))
else {
device.count_rx_dropped(1);
break;
};
dst.copy_from_slice(packet);
device.count_rx(frame_len);
moved_rx = true;
budget -= 1;
}
device.drain_device_counters();
}
moved_rx
}
pub fn send_on_device(
&mut self,
dev: usize,
next_hop: IpAddress,
packet: &[u8],
_timestamp: Instant,
) -> bool {
let Router {
rx_buffer, devices, ..
} = self;
let device = &mut devices[dev];
if device.interface_id == InterfaceId::LOOPBACK {
let ok =
inject_loopback_rx_direct(rx_buffer, next_hop, packet, &mut SocketSet::new(vec![]));
if ok {
device.count_tx(packet.len());
device.count_rx(packet.len());
} else {
device.count_rx_dropped(1);
}
return ok;
}
device.send(next_hop, packet, now())
}
pub fn arp_entries(&self, timestamp: Instant) -> Vec<ArpEntry> {
let mut entries = Vec::new();
for device in &self.devices {
entries.extend(device.inner.arp_entries(timestamp));
}
entries
}
pub fn net_dev_stats(&self) -> Vec<NetDevStats> {
self.devices.iter().map(|device| device.stats()).collect()
}
pub fn register_waker(&self, _binding: DeviceBinding, _waker: &core::task::Waker) {
crate::request_poll();
}
pub fn dispatch(&mut self, _timestamp: Instant, sockets: &mut SocketSet<'_>) -> bool {
let mut poll_next = false;
let Router {
rx_buffer,
tx_buffer,
pending_fanout,
devices,
table,
..
} = self;
while let Some(packet) = tx_buffer.get_allocated(0, 1).first() {
let packet = packet.as_bytes();
let outcome = match IpVersion::of_packet(packet).expect("got invalid IP packet") {
IpVersion::Ipv4 => {
let packet = smoltcp::wire::Ipv4Packet::new_checked(packet)
.expect("got invalid IPv4 packet");
let src_addr = IpAddress::Ipv4(packet.src_addr());
let dst_addr = IpAddress::Ipv4(packet.dst_addr());
if packet.dst_addr().is_broadcast() {
dispatch_link_local_fanout(
devices,
pending_fanout,
dst_addr,
packet.into_inner(),
)
} else {
dispatch_unicast_packet(
rx_buffer,
devices,
table,
src_addr,
dst_addr,
packet.into_inner(),
sockets,
)
}
}
IpVersion::Ipv6 => {
let packet = smoltcp::wire::Ipv6Packet::new_checked(packet)
.expect("got invalid IPv6 packet");
let src_addr = IpAddress::Ipv6(packet.src_addr());
let dst_addr = IpAddress::Ipv6(packet.dst_addr());
if packet.dst_addr().is_multicast() {
dispatch_link_local_fanout(
devices,
pending_fanout,
dst_addr,
packet.into_inner(),
)
} else {
dispatch_unicast_packet(
rx_buffer,
devices,
table,
src_addr,
dst_addr,
packet.into_inner(),
sockets,
)
}
}
};
match outcome {
DispatchOutcome::Consumed(next) => {
poll_next |= next;
tx_buffer
.dequeue_one()
.expect("the packet was only peeked while dispatching");
}
DispatchOutcome::Retry(next) => {
poll_next |= next;
break;
}
}
}
if tx_buffer.is_empty() {
tx_buffer.clear();
}
poll_next
}
}
fn dispatch_link_local_fanout(
devices: &mut [DeviceHandle],
pending: &mut Vec<usize>,
dst_addr: IpAddress,
packet: &[u8],
) -> DispatchOutcome {
if pending.is_empty() {
pending.extend(devices.iter().enumerate().filter_map(|(index, dev)| {
(dev.interface_id != InterfaceId::LOOPBACK).then_some(index)
}));
}
let mut poll_next = false;
pending.retain(|&index| {
let dev = &mut devices[index];
match dev.try_send(dst_addr, packet, now()) {
Ok(consumed) => {
poll_next |= consumed;
false
}
Err(NetDeviceError::Again) => true,
Err(error) => {
warn!("{}: transmit failed: {error:?}", dev.name);
dev.count_tx_errors(1);
dev.drain_device_counters();
false
}
}
});
if pending.is_empty() {
DispatchOutcome::Consumed(poll_next)
} else {
DispatchOutcome::Retry(poll_next)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum DispatchOutcome {
Consumed(bool),
Retry(bool),
}
fn dispatch_unicast_packet(
rx_buffer: &mut RouterPacketBuffer,
devices: &mut [DeviceHandle],
table: &SharedRouteTable,
src_addr: IpAddress,
dst_addr: IpAddress,
packet: &[u8],
sockets: &mut SocketSet<'_>,
) -> DispatchOutcome {
let route = {
let routes = table.read();
let Some(route) = routes.select_route_for_source(&dst_addr, &src_addr) else {
debug!(
"No route found for source {} destination {}",
src_addr, dst_addr
);
return DispatchOutcome::Consumed(false);
};
route
};
let dev = &mut devices[route.dev];
if dev.interface_id == InterfaceId::LOOPBACK {
let ok = inject_loopback_rx_direct(rx_buffer, dst_addr, packet, sockets);
if ok {
dev.count_tx(packet.len());
dev.count_rx(packet.len());
} else {
dev.count_rx_dropped(1);
}
DispatchOutcome::Consumed(ok)
} else {
match dev.try_send(route.next_hop, packet, now()) {
Ok(consumed) => DispatchOutcome::Consumed(consumed),
Err(NetDeviceError::Again) => DispatchOutcome::Retry(false),
Err(error) => {
warn!("{}: transmit failed: {error:?}", dev.name);
dev.count_tx_errors(1);
dev.drain_device_counters();
DispatchOutcome::Consumed(false)
}
}
}
}
fn inject_loopback_rx_direct(
rx_buffer: &mut RouterPacketBuffer,
dst_addr: IpAddress,
packet: &[u8],
sockets: &mut SocketSet<'_>,
) -> bool {
snoop_tcp_packet(packet, sockets);
let Ok(dst) = rx_buffer.enqueue(packet.len(), rx_metadata(InterfaceId::LOOPBACK, packet))
else {
warn!("Loopback: RX buffer full, dropping packet to {}", dst_addr);
return false;
};
dst.copy_from_slice(packet);
true
}
pub struct TxToken<'a>(&'a mut RingBuffer<'static, TxPacket>);
impl smoltcp::phy::TxToken for TxToken<'_> {
fn consume<R, F>(self, len: usize, f: F) -> R
where
F: FnOnce(&mut [u8]) -> R,
{
let slot = self
.0
.enqueue_one()
.expect("This was checked before creating the TxToken");
slot.len = len;
let packet = &mut slot.bytes[..len];
let result = f(packet);
apply_egress_ip_tos(packet);
result
}
}
fn snoop_tcp_packet(buf: &[u8], sockets: &mut SocketSet<'_>) {
if buf.is_empty() {
return;
}
let (src_addr, dst_addr, payload) = match IpVersion::of_packet(buf) {
Ok(IpVersion::Ipv4) => {
let Ok(packet) = Ipv4Packet::new_checked(buf) else {
return;
};
if packet.next_header() != IpProtocol::Tcp {
return;
}
(
IpAddress::Ipv4(packet.src_addr()),
IpAddress::Ipv4(packet.dst_addr()),
packet.payload(),
)
}
Ok(IpVersion::Ipv6) => {
let Ok(packet) = Ipv6Packet::new_checked(buf) else {
return;
};
if packet.next_header() != IpProtocol::Tcp {
return;
}
(
IpAddress::Ipv6(packet.src_addr()),
IpAddress::Ipv6(packet.dst_addr()),
packet.payload(),
)
}
Err(_) => return,
};
let Ok(tcp_packet) = TcpPacket::new_checked(payload) else {
return;
};
let src_addr = (src_addr, tcp_packet.src_port()).into();
let dst_addr = (dst_addr, tcp_packet.dst_port()).into();
let is_first = tcp_packet.syn() && !tcp_packet.ack();
if is_first {
LISTEN_TABLE.incoming_tcp_packet(src_addr, dst_addr, sockets);
}
}
enum RxTokenPacket<'a> {
Borrowed(&'a [u8]),
Owned(DeviceRxPacket),
}
pub struct RxToken<'a> {
interface_id: InterfaceId,
packet_meta: PacketMeta,
packet: RxTokenPacket<'a>,
}
impl<'a> smoltcp::phy::RxToken for RxToken<'a> {
fn consume<R, F>(self, f: F) -> R
where
F: FnOnce(&[u8]) -> R,
{
let _ingress_if = self.interface_id;
match self.packet {
RxTokenPacket::Borrowed(packet) => f(packet),
RxTokenPacket::Owned(packet) => packet.consume(f),
}
}
fn meta(&self) -> PacketMeta {
self.packet_meta
}
}
impl smoltcp::phy::Device for Router {
type RxToken<'a> = RxToken<'a>;
type TxToken<'a> = TxToken<'a>;
fn receive(&mut self, _timestamp: Instant) -> Option<(Self::RxToken<'_>, Self::TxToken<'_>)> {
if self.tx_buffer.is_full() {
return None;
}
let Self {
rx_buffer,
ready_rx,
tx_buffer,
..
} = self;
let rx_token = if !rx_buffer.is_empty() {
let (metadata, packet) = rx_buffer.dequeue().unwrap();
RxToken {
interface_id: metadata.interface_id,
packet_meta: metadata.packet_meta,
packet: RxTokenPacket::Borrowed(packet),
}
} else {
let packet = ready_rx.pop_front()?;
RxToken {
interface_id: packet.metadata.interface_id,
packet_meta: packet.metadata.packet_meta,
packet: RxTokenPacket::Owned(packet.packet),
}
};
Some((rx_token, TxToken(tx_buffer)))
}
fn transmit(&mut self, _timestamp: Instant) -> Option<Self::TxToken<'_>> {
if self.tx_buffer.is_full() {
None
} else {
Some(TxToken(&mut self.tx_buffer))
}
}
fn capabilities(&self) -> DeviceCapabilities {
let mut caps = DeviceCapabilities::default();
caps.medium = Medium::Ip;
caps.max_transmission_unit = STANDARD_MTU;
caps.max_burst_size = Some(SOCKET_BUFFER_SIZE);
caps
}
}
#[cfg(test)]
mod tests {
use smoltcp::{
phy::{Device as _, TxToken as _},
storage::PacketBuffer,
};
use super::*;
use crate::device::TxChecksumCapabilities;
#[test]
fn stack_tcp_and_udp_emit_complete_software_checksums() {
use smoltcp::{
iface::{Config, Interface},
socket::{tcp, udp},
wire::{HardwareAddress, UdpPacket},
};
let mut router = Router::new(Arc::new(RwLock::new(RouteTable::new())));
router.add_device(
IF0,
Box::new(crate::device::EthernetDevice::new(
"checksum".into(),
Box::new(ChecksumPort),
None,
)),
);
let now = Instant::from_millis(0);
let mut interface = Interface::new(Config::new(HardwareAddress::Ip), &mut router, now);
interface.update_ip_addrs(|addrs| {
addrs
.push(ipv4_cidr(Ipv4Address::new(10, 0, 0, 2), 24))
.unwrap()
});
let destination = IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1));
let mut sockets = SocketSet::new(vec![]);
let mut tcp = tcp::Socket::new(
tcp::SocketBuffer::new(vec![0; 1024]),
tcp::SocketBuffer::new(vec![0; 1024]),
);
tcp.connect(interface.context(), (destination, 4321), 1234)
.unwrap();
sockets.add(tcp);
interface.poll_egress(now, &mut router, &mut sockets);
let packet = router
.tx_buffer
.dequeue_one()
.expect("TCP SYN must be emitted");
let ip = Ipv4Packet::new_checked(packet.as_bytes()).unwrap();
let tcp = TcpPacket::new_checked(ip.payload()).unwrap();
assert!(tcp.syn());
assert!(tcp.verify_checksum(&ip.src_addr().into(), &ip.dst_addr().into()));
let mut udp = udp::Socket::new(
udp::PacketBuffer::new(vec![udp::PacketMetadata::EMPTY; 1], vec![0; 64]),
udp::PacketBuffer::new(vec![udp::PacketMetadata::EMPTY; 1], vec![0; 64]),
);
udp.bind(1235).unwrap();
udp.send_slice(b"checksum", (destination, 4322)).unwrap();
sockets.add(udp);
interface.poll_egress(now, &mut router, &mut sockets);
let packet = router
.tx_buffer
.dequeue_one()
.expect("UDP packet must be emitted");
let ip = Ipv4Packet::new_checked(packet.as_bytes()).unwrap();
let udp = UdpPacket::new_checked(ip.payload()).unwrap();
assert_ne!(udp.checksum(), 0, "ordinary UDP must generate a checksum");
assert!(udp.verify_checksum(&ip.src_addr().into(), &ip.dst_addr().into()));
assert_eq!(udp.payload(), b"checksum");
}
#[test]
fn loopback_preserves_raw_udp_checksum() {
let table = Arc::new(RwLock::new(RouteTable::new()));
let mut router = Router::new(table);
let mut sockets = SocketSet::new(vec![]);
for checksum in [0u16, 0x1234] {
let mut packet = [0u8; 32];
packet[0] = 0x45;
packet[2..4].copy_from_slice(&32u16.to_be_bytes());
packet[8] = 64;
packet[9] = 17;
packet[12..16].copy_from_slice(&[127, 0, 0, 1]);
packet[16..20].copy_from_slice(&[127, 0, 0, 1]);
packet[20..22].copy_from_slice(&1234u16.to_be_bytes());
packet[22..24].copy_from_slice(&4321u16.to_be_bytes());
packet[24..26].copy_from_slice(&12u16.to_be_bytes());
packet[26..28].copy_from_slice(&checksum.to_be_bytes());
assert!(inject_loopback_rx_direct(
&mut router.rx_buffer,
IpAddress::Ipv4(Ipv4Address::LOCALHOST),
&packet,
&mut sockets
));
let (_, received) = router.rx_buffer.dequeue().unwrap();
assert_eq!(
received, &packet,
"loopback rewrote raw UDP transport bytes"
);
}
}
const IF0: InterfaceId = InterfaceId::new(2);
const IF1: InterfaceId = InterfaceId::new(3);
const SRC0: IpAddress = IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 2));
const SRC1: IpAddress = IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 2));
struct EmptyDevice;
impl Device for EmptyDevice {
fn name(&self) -> &str {
"empty"
}
fn recv(
&mut self,
_interface_id: InterfaceId,
_buffer: &mut PacketBuffer<InterfaceId>,
_timestamp: Instant,
_snoop: &mut dyn FnMut(&[u8]),
) -> usize {
0
}
fn send(&mut self, _next_hop: IpAddress, _packet: &[u8], _timestamp: Instant) -> usize {
0
}
}
struct RetryDevice;
impl Device for RetryDevice {
fn name(&self) -> &str {
"retry"
}
fn recv(
&mut self,
_interface_id: InterfaceId,
_buffer: &mut PacketBuffer<InterfaceId>,
_timestamp: Instant,
_snoop: &mut dyn FnMut(&[u8]),
) -> usize {
0
}
fn send(&mut self, _next_hop: IpAddress, _packet: &[u8], _timestamp: Instant) -> usize {
0
}
fn try_send(
&mut self,
_next_hop: IpAddress,
_packet: &[u8],
_timestamp: Instant,
) -> crate::device::NetDeviceResult<usize> {
Err(NetDeviceError::Again)
}
}
struct ChecksumPort;
impl crate::device::EthernetFramePort for ChecksumPort {
fn device_name(&self) -> &str {
"checksum"
}
fn mac_address(&self) -> [u8; 6] {
[2, 0, 0, 0, 0, 1]
}
fn checksum_capabilities(&self) -> TxChecksumCapabilities {
TxChecksumCapabilities::TCP_UDP
}
fn transmit(
&mut self,
_: &crate::device::ProtocolEthernetFrame,
) -> crate::device::NetDeviceResult {
Err(NetDeviceError::Again)
}
fn receive(
&mut self,
) -> crate::device::NetDeviceResult<crate::device::ProtocolEthernetFrame> {
Err(NetDeviceError::Again)
}
}
fn test_device_handle(device: Box<dyn Device>) -> DeviceHandle {
DeviceHandle::new(IF0, device)
}
fn ipv4_cidr(addr: Ipv4Address, prefix_len: u8) -> IpCidr {
Ipv4Cidr::new(addr, prefix_len).into()
}
#[test]
fn route_lookup_uses_longest_prefix() {
let mut table = RouteTable::new();
table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
0,
IF0,
SRC0,
100,
));
table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::new(10, 0, 1, 0), 24),
None,
1,
IF1,
SRC1,
200,
));
let route = table
.select_route_if(&IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 99)), |_| true)
.unwrap();
assert_eq!(route.dev, 1);
assert_eq!(route.interface_id, IF1);
assert_eq!(route.source, SRC1);
assert_eq!(
route.next_hop,
IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 99))
);
}
#[test]
fn transient_tx_backpressure_keeps_the_router_packet_queued() {
let table = Arc::new(RwLock::new(RouteTable::new()));
let mut router = Router::new(Arc::clone(&table));
router.add_device(IF0, Box::new(RetryDevice));
router.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
0,
IF0,
SRC0,
100,
));
router
.transmit(Instant::from_millis(0))
.expect("the empty router TX queue has capacity")
.consume(20, |packet| {
packet[0] = 0x45;
packet[2..4].copy_from_slice(&20u16.to_be_bytes());
packet[12..16].copy_from_slice(&[10, 0, 0, 2]);
packet[16..20].copy_from_slice(&[198, 51, 100, 1]);
});
let mut sockets = SocketSet::new(vec![]);
assert!(!router.dispatch(Instant::from_millis(0), &mut sockets));
assert_eq!(router.tx_buffer.len(), 1);
assert_eq!(router.tx_buffer.get_allocated(0, 1)[0].as_bytes().len(), 20);
assert_eq!(router.devices[0].stats().tx_packets, 0);
assert_eq!(router.devices[0].stats().tx_dropped, 0);
}
#[test]
fn drained_tx_queue_reuses_packet_storage() {
use smoltcp::phy::RxToken as _;
let mut router = Router::new(Arc::new(RwLock::new(RouteTable::new())));
router.add_device(InterfaceId::LOOPBACK, Box::new(EmptyDevice));
router.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::LOCALHOST, 8),
None,
0,
InterfaceId::LOOPBACK,
IpAddress::Ipv4(Ipv4Address::LOCALHOST),
0,
));
let now = Instant::from_millis(0);
let mut sockets = SocketSet::new(vec![]);
let mut first_slot = None;
for len in [64, STANDARD_MTU, 64] {
let mut packet = vec![0; len];
packet[0] = 0x45;
packet[2..4].copy_from_slice(&(len as u16).to_be_bytes());
packet[12..16].copy_from_slice(&[127, 0, 0, 1]);
packet[16..20].copy_from_slice(&[127, 0, 0, 1]);
let address = router.transmit(now).unwrap().consume(len, |dst| {
dst.copy_from_slice(&packet);
dst.as_ptr() as usize
});
assert_eq!(
address,
*first_slot.get_or_insert(address),
"a drained TX queue must reuse the first packet's storage"
);
assert!(router.dispatch(now, &mut sockets));
let (rx, _tx) = router.receive(now).unwrap();
rx.consume(|received| assert_eq!(received, packet));
}
assert!(router.receive(now).is_none());
}
#[test]
fn transmit_token_survives_backpressure_and_payload_wrap() {
check_tx_token_after_payload_wrap(false);
}
#[test]
fn receive_token_survives_backpressure_and_payload_wrap() {
check_tx_token_after_payload_wrap(true);
}
fn check_tx_token_after_payload_wrap(reply_to_rx: bool) {
use ax_sync::SpinLock;
use smoltcp::phy::RxToken as _;
#[derive(Default)]
struct TxProbe {
allowance: usize,
packets: Vec<Vec<u8>>,
}
struct BackpressureDevice(Arc<SpinLock<TxProbe>>);
impl Device for BackpressureDevice {
fn name(&self) -> &str {
"backpressure"
}
fn recv(
&mut self,
_: InterfaceId,
_: &mut PacketBuffer<InterfaceId>,
_: Instant,
_: &mut dyn FnMut(&[u8]),
) -> usize {
0
}
fn send(&mut self, _: IpAddress, _: &[u8], _: Instant) -> usize {
panic!("dispatch must use the fallible TX contract")
}
fn try_send(
&mut self,
_: IpAddress,
packet: &[u8],
_: Instant,
) -> crate::device::NetDeviceResult<usize> {
let mut probe = self.0.lock_irqsave();
if probe.allowance == 0 {
return Err(NetDeviceError::Again);
}
probe.allowance -= 1;
probe.packets.push(packet.to_vec());
Ok(packet.len())
}
}
fn packet(len: usize, id: u8) -> Vec<u8> {
let mut packet = vec![id; len];
packet[0] = 0x45;
packet[2..4].copy_from_slice(&(len as u16).to_be_bytes());
packet[12..16].copy_from_slice(&[10, 0, 0, 2]);
packet[16..20].copy_from_slice(&[198, 51, 100, 1]);
packet
}
let mut router = Router::new(Arc::new(RwLock::new(RouteTable::new())));
let probe = Arc::new(SpinLock::new(TxProbe::default()));
router.add_device(IF0, Box::new(BackpressureDevice(Arc::clone(&probe))));
router.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
0,
IF0,
SRC0,
100,
));
let now = Instant::from_millis(0);
let mut sockets = SocketSet::new(vec![]);
let mut expected = vec![];
for id in 0..2 {
let packet = packet(20, id);
router.transmit(now).unwrap().consume(packet.len(), |dst| {
dst.copy_from_slice(&packet);
});
expected.push(packet);
}
probe.lock_irqsave().allowance = 1;
assert!(router.dispatch(now, &mut sockets));
assert_eq!(probe.lock_irqsave().packets, expected[..1]);
for id in 0..SOCKET_BUFFER_SIZE - 1 {
let packet = packet(STANDARD_MTU, id as u8);
router.transmit(now).unwrap().consume(packet.len(), |dst| {
dst.copy_from_slice(&packet);
});
expected.push(packet);
assert!(!router.dispatch(now, &mut sockets));
}
assert!(router.transmit(now).is_none());
let incoming = packet(20, 0);
router
.rx_buffer
.enqueue(incoming.len(), rx_metadata(IF0, &incoming))
.unwrap()
.copy_from_slice(&incoming);
assert!(router.receive(now).is_none());
probe.lock_irqsave().allowance = 1;
assert!(router.dispatch(now, &mut sockets));
assert_eq!(probe.lock_irqsave().packets, expected[..2]);
let token = if reply_to_rx {
let (rx, tx) = router.receive(now).unwrap();
rx.consume(|packet| assert_eq!(packet, incoming));
tx
} else {
router.transmit(now).unwrap()
};
let final_packet = packet(STANDARD_MTU, 0xfe);
assert_eq!(
token.consume(final_packet.len(), |dst| {
dst.copy_from_slice(&final_packet);
42
}),
42
);
expected.push(final_packet);
probe.lock_irqsave().allowance = expected.len();
assert!(router.dispatch(now, &mut sockets));
assert_eq!(probe.lock_irqsave().packets, expected);
assert!(router.transmit(now).is_some());
assert!(!router.dispatch(now, &mut sockets));
assert_eq!(router.devices[0].stats().tx_packets, expected.len() as u64);
assert_eq!(router.devices[0].stats().tx_errors, 0);
assert_eq!(router.devices[0].stats().tx_dropped, 0);
}
#[test]
fn fanout_retries_only_blocked_ports_without_repeating_accepted_packets() {
use ax_sync::SpinLock;
#[derive(Default)]
struct TxProbe {
failures: VecDeque<NetDeviceError>,
attempts: usize,
packets: Vec<Vec<u8>>,
}
struct FanoutDevice(Arc<SpinLock<TxProbe>>);
impl Device for FanoutDevice {
fn name(&self) -> &str {
"fanout"
}
fn recv(
&mut self,
_: InterfaceId,
_: &mut PacketBuffer<InterfaceId>,
_: Instant,
_: &mut dyn FnMut(&[u8]),
) -> usize {
0
}
fn send(&mut self, _: IpAddress, _: &[u8], _: Instant) -> usize {
panic!("fanout must preserve the fallible TX contract")
}
fn try_send(
&mut self,
_: IpAddress,
packet: &[u8],
_: Instant,
) -> crate::device::NetDeviceResult<usize> {
let mut probe = self.0.lock_irqsave();
probe.attempts += 1;
if let Some(error) = probe.failures.pop_front() {
return Err(error);
}
probe.packets.push(packet.to_vec());
Ok(packet.len())
}
}
for ipv6 in [false, true] {
let mut packet = if ipv6 { vec![0u8; 40] } else { vec![0u8; 20] };
if ipv6 {
packet[0] = 0x60;
packet[24] = 0xff;
packet[25] = 2;
packet[39] = 1;
} else {
packet[0] = 0x45;
packet[2..4].copy_from_slice(&20u16.to_be_bytes());
packet[16..20].fill(0xff);
}
let mut router = Router::new(Arc::new(RwLock::new(RouteTable::new())));
let probes: Vec<_> = [
vec![],
vec![NetDeviceError::Again, NetDeviceError::Again],
vec![NetDeviceError::Again],
vec![NetDeviceError::Io],
vec![],
]
.into_iter()
.enumerate()
.map(|(index, failures)| {
let probe = Arc::new(SpinLock::new(TxProbe {
failures: failures.into(),
..TxProbe::default()
}));
let id = if index == 4 {
InterfaceId::LOOPBACK
} else {
InterfaceId::new(index as u32 + 2)
};
router.add_device(id, Box::new(FanoutDevice(Arc::clone(&probe))));
probe
})
.collect();
let mut next_packet = packet.clone();
next_packet[1] = 1;
for queued in [&packet, &next_packet] {
router
.transmit(Instant::from_millis(0))
.unwrap()
.consume(queued.len(), |dst| dst.copy_from_slice(queued));
}
let mut sockets = SocketSet::new(vec![]);
for completed_port in [0, 2] {
assert!(router.dispatch(Instant::from_millis(0), &mut sockets));
assert_eq!(router.tx_buffer.get_allocated(0, 1)[0].as_bytes(), packet);
assert_eq!(router.tx_buffer.len(), 2);
assert_eq!(
probes[completed_port].lock_irqsave().packets,
vec![packet.clone()]
);
assert_eq!(probes[0].lock_irqsave().attempts, 1);
assert_eq!(probes[3].lock_irqsave().attempts, 1);
assert_eq!(router.devices[3].stats().tx_errors, 1);
}
probes[1]
.lock_irqsave()
.failures
.push_back(NetDeviceError::Again);
assert!(!router.dispatch(Instant::from_millis(0), &mut sockets));
assert_eq!(router.tx_buffer.get_allocated(0, 1)[0].as_bytes(), packet);
assert_eq!(probes[0].lock_irqsave().attempts, 1);
assert_eq!(probes[2].lock_irqsave().attempts, 2);
assert!(router.dispatch(Instant::from_millis(0), &mut sockets));
assert!(router.tx_buffer.is_empty());
for (index, attempts) in [2, 5, 3, 2].into_iter().enumerate() {
let probe = probes[index].lock_irqsave();
assert_eq!(probe.attempts, attempts);
let expected = if index == 3 {
vec![next_packet.clone()]
} else {
vec![packet.clone(), next_packet.clone()]
};
assert_eq!(probe.packets, expected);
assert_eq!(
router.devices[index].stats().tx_packets,
expected.len() as u64
);
assert_eq!(router.devices[index].stats().tx_dropped, 0);
}
assert_eq!(probes[4].lock_irqsave().attempts, 0);
}
}
#[test]
fn router_keeps_software_checksums_even_with_offload_capable_devices() {
let table = Arc::new(RwLock::new(RouteTable::new()));
let mut router = Router::new(table);
router.add_device(
IF0,
Box::new(crate::device::EthernetDevice::new(
"checksum".into(),
Box::new(ChecksumPort),
None,
)),
);
let caps = smoltcp::phy::Device::capabilities(&router);
assert!(caps.checksum.tcp.rx());
assert!(caps.checksum.tcp.tx());
assert!(caps.checksum.udp.rx());
assert!(caps.checksum.udp.tx());
router.add_device(IF1, Box::new(EmptyDevice));
let caps = smoltcp::phy::Device::capabilities(&router);
assert!(caps.checksum.tcp.tx());
assert!(caps.checksum.udp.tx());
}
#[test]
fn route_lookup_uses_metric_for_same_prefix() {
let mut table = RouteTable::new();
let dst = IpAddress::Ipv4(Ipv4Address::new(203, 0, 113, 10));
table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
0,
IF0,
SRC0,
200,
));
table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 1))),
1,
IF1,
SRC1,
100,
));
let route = table.select_route_if(&dst, |_| true).unwrap();
assert_eq!(route.interface_id, IF1);
assert_eq!(route.metric, 100);
assert_eq!(
route.next_hop,
IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 1))
);
}
#[test]
fn route_lookup_keeps_stable_order_for_equal_metric() {
let mut table = RouteTable::new();
let dst = IpAddress::Ipv4(Ipv4Address::new(203, 0, 113, 10));
table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
0,
IF0,
SRC0,
100,
));
table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 1))),
1,
IF1,
SRC1,
100,
));
let route = table.select_route_if(&dst, |_| true).unwrap();
assert_eq!(route.interface_id, IF0);
assert_eq!(
route.next_hop,
IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))
);
}
#[test]
fn route_lookup_skips_unusable_interface() {
let mut table = RouteTable::new();
let dst = IpAddress::Ipv4(Ipv4Address::new(203, 0, 113, 10));
table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
0,
IF0,
SRC0,
100,
));
table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 1, 1))),
1,
IF1,
SRC1,
200,
));
let route = table
.select_route_if(&dst, |interface_id| interface_id != IF0)
.unwrap();
assert_eq!(route.interface_id, IF1);
}
#[test]
fn snoop_tcp_packet_drops_truncated_ip_and_tcp_headers() {
const IPV4_HEADER_LEN: usize = 20;
const IPV6_HEADER_LEN: usize = 40;
const TCP_HEADER_LEN: usize = 20;
let mut sockets = SocketSet::new(vec![]);
let mut ipv4_tcp = [0u8; IPV4_HEADER_LEN + TCP_HEADER_LEN];
ipv4_tcp[0] = 0x45;
let ipv4_tcp_len = ipv4_tcp.len() as u16;
ipv4_tcp[2..4].copy_from_slice(&ipv4_tcp_len.to_be_bytes());
ipv4_tcp[9] = IpProtocol::Tcp.into();
for len in 0..ipv4_tcp.len() {
snoop_tcp_packet(&ipv4_tcp[..len], &mut sockets);
}
let mut ipv6_tcp = [0u8; IPV6_HEADER_LEN + TCP_HEADER_LEN];
ipv6_tcp[0] = 0x60;
ipv6_tcp[4..6].copy_from_slice(&20u16.to_be_bytes());
ipv6_tcp[6] = IpProtocol::Tcp.into();
for len in 0..ipv6_tcp.len() {
snoop_tcp_packet(&ipv6_tcp[..len], &mut sockets);
}
for tcp_len in 0..TCP_HEADER_LEN {
let mut ipv4_tcp = vec![0u8; IPV4_HEADER_LEN + tcp_len];
ipv4_tcp[0] = 0x45;
let ipv4_len = ipv4_tcp.len() as u16;
ipv4_tcp[2..4].copy_from_slice(&ipv4_len.to_be_bytes());
ipv4_tcp[9] = IpProtocol::Tcp.into();
snoop_tcp_packet(&ipv4_tcp, &mut sockets);
let mut ipv6_tcp = vec![0u8; IPV6_HEADER_LEN + tcp_len];
ipv6_tcp[0] = 0x60;
ipv6_tcp[4..6].copy_from_slice(&(tcp_len as u16).to_be_bytes());
ipv6_tcp[6] = IpProtocol::Tcp.into();
snoop_tcp_packet(&ipv6_tcp, &mut sockets);
}
}
#[test]
fn default_routes_only_reports_zero_prefix_ipv4_rules() {
let mut table = RouteTable::new();
table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::UNSPECIFIED, 0),
Some(IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))),
0,
IF0,
SRC0,
100,
));
table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::new(10, 0, 1, 0), 24),
None,
1,
IF1,
SRC1,
100,
));
let routes = table.default_routes();
assert_eq!(routes.len(), 1);
assert_eq!(routes[0].interface_id, IF0);
}
#[test]
fn no_route_does_not_count_interface_tx_dropped() {
use smoltcp::{iface::SocketSet, storage::PacketMetadata};
let dev0 = test_device_handle(Box::new(EmptyDevice));
let dev1 = DeviceHandle::new(IF1, Box::new(EmptyDevice));
let mut devices = vec![dev0, dev1];
let mut route_table = RouteTable::new();
route_table.add_rule(Rule::new(
ipv4_cidr(Ipv4Address::new(10, 0, 0, 0), 24),
Some(SRC0),
0,
IF0, SRC0,
100,
));
let shared_table: SharedRouteTable = Arc::new(RwLock::new(route_table));
let mut rx_buffer: RouterPacketBuffer = PacketBuffer::new(
vec![PacketMetadata::EMPTY; 1],
vec![0u8; super::STANDARD_MTU],
);
let mut sockets = SocketSet::new(vec![]);
let src_addr = SRC0;
let dst_addr = IpAddress::Ipv4(Ipv4Address::new(203, 0, 113, 10));
let packet = [0u8; 64];
let before: Vec<_> = devices.iter().map(|d| d.stats()).collect();
let outcome = dispatch_unicast_packet(
&mut rx_buffer,
&mut devices,
&shared_table,
src_addr,
dst_addr,
&packet,
&mut sockets,
);
assert_eq!(
outcome,
DispatchOutcome::Consumed(false),
"no-route dispatch must consume the packet without scheduling work"
);
for (i, dev) in devices.iter().enumerate() {
let snap = dev.stats();
assert_eq!(
snap.tx_dropped, before[i].tx_dropped,
"device {i} tx_dropped changed from {} to {} after no-route dispatch",
before[i].tx_dropped, snap.tx_dropped,
);
}
}
}
#[cfg(test)]
mod l2_counter_tests {
use smoltcp::{
storage::{PacketBuffer, PacketMetadata},
time::Instant,
wire::{IpAddress, Ipv4Address},
};
use super::*;
const IF0: InterfaceId = InterfaceId::new(2);
struct CountingMockDevice {
name: &'static str,
send_returns: usize,
recv_returns: usize,
deferred_tx_lens: Vec<usize>,
deferred_rx_lens: Vec<usize>,
}
impl Device for CountingMockDevice {
fn name(&self) -> &str {
self.name
}
fn recv(
&mut self,
_interface_id: InterfaceId,
_buffer: &mut PacketBuffer<InterfaceId>,
_timestamp: Instant,
_snoop: &mut dyn FnMut(&[u8]),
) -> usize {
self.recv_returns
}
fn send(&mut self, _next_hop: IpAddress, _packet: &[u8], _timestamp: Instant) -> usize {
self.send_returns
}
fn drain_deferred_tx(&mut self) -> Vec<usize> {
core::mem::take(&mut self.deferred_tx_lens)
}
fn drain_deferred_rx(&mut self) -> Vec<usize> {
core::mem::take(&mut self.deferred_rx_lens)
}
}
fn test_device_handle(device: Box<dyn Device>) -> DeviceHandle {
DeviceHandle::new(IF0, device)
}
fn test_ip() -> IpAddress {
IpAddress::Ipv4(Ipv4Address::new(10, 0, 0, 1))
}
fn test_packet_buffer() -> PacketBuffer<'static, InterfaceId> {
PacketBuffer::new(
vec![PacketMetadata::EMPTY; 1],
vec![0u8; super::STANDARD_MTU],
)
}
#[test]
fn count_rx_accumulates_bytes_and_packets() {
let device = test_device_handle(Box::new(CountingMockDevice {
name: "mock",
send_returns: 0,
deferred_tx_lens: vec![],
deferred_rx_lens: vec![],
recv_returns: 0,
}));
device.count_rx(100);
assert_eq!(device.stats().rx_bytes, 100);
assert_eq!(device.stats().rx_packets, 1);
device.count_rx(200);
assert_eq!(device.stats().rx_bytes, 300);
assert_eq!(device.stats().rx_packets, 2);
}
#[test]
fn count_tx_accumulates_bytes_and_packets() {
let device = test_device_handle(Box::new(CountingMockDevice {
name: "mock",
send_returns: 0,
deferred_tx_lens: vec![],
deferred_rx_lens: vec![],
recv_returns: 0,
}));
device.count_tx(64);
assert_eq!(device.stats().tx_bytes, 64);
assert_eq!(device.stats().tx_packets, 1);
device.count_tx(1500);
assert_eq!(device.stats().tx_bytes, 1564);
assert_eq!(device.stats().tx_packets, 2);
}
#[test]
fn stats_reflects_current_counters_after_counting() {
let device = test_device_handle(Box::new(CountingMockDevice {
name: "mock",
send_returns: 0,
deferred_tx_lens: vec![],
deferred_rx_lens: vec![],
recv_returns: 0,
}));
device.count_rx(100);
device.count_tx(64);
let snap = device.stats();
assert_eq!(snap.rx_bytes, 100);
assert_eq!(snap.rx_packets, 1);
assert_eq!(snap.tx_bytes, 64);
assert_eq!(snap.tx_packets, 1);
}
#[test]
fn send_returns_frame_len_tx_counts_l2_not_ip_payload() {
let mut device = test_device_handle(Box::new(CountingMockDevice {
name: "mock",
send_returns: 1514, deferred_tx_lens: vec![],
deferred_rx_lens: vec![],
recv_returns: 0,
}));
let frame_len = device
.inner
.send(test_ip(), &[0u8; 100], Instant::from_millis(0));
assert_eq!(frame_len, 1514);
if frame_len > 0 {
device.count_tx(frame_len);
}
let snap = device.stats();
assert_eq!(snap.tx_bytes, 1514);
assert_eq!(snap.tx_packets, 1);
}
#[test]
fn send_returns_zero_no_tx_counted() {
let mut device = test_device_handle(Box::new(CountingMockDevice {
name: "mock",
send_returns: 0, deferred_tx_lens: vec![],
deferred_rx_lens: vec![],
recv_returns: 0,
}));
let frame_len = device
.inner
.send(test_ip(), &[0u8; 100], Instant::from_millis(0));
assert_eq!(frame_len, 0);
if frame_len > 0 {
device.count_tx(frame_len);
}
let snap = device.stats();
assert_eq!(snap.tx_bytes, 0);
assert_eq!(snap.tx_packets, 0);
}
#[test]
fn recv_returns_frame_len_rx_counts_it() {
let mut device = test_device_handle(Box::new(CountingMockDevice {
name: "mock",
send_returns: 0,
deferred_tx_lens: vec![],
deferred_rx_lens: vec![],
recv_returns: 1514,
}));
let frame_len = device.inner.recv(
IF0,
&mut test_packet_buffer(),
Instant::from_millis(0),
&mut |_| {},
);
assert_eq!(frame_len, 1514);
if frame_len > 0 {
device.count_rx(frame_len);
}
let snap = device.stats();
assert_eq!(snap.rx_bytes, 1514);
assert_eq!(snap.rx_packets, 1);
}
#[test]
fn recv_returns_zero_no_rx_counted() {
let mut device = test_device_handle(Box::new(CountingMockDevice {
name: "mock",
send_returns: 0,
deferred_tx_lens: vec![],
deferred_rx_lens: vec![],
recv_returns: 0, }));
let frame_len = device.inner.recv(
IF0,
&mut test_packet_buffer(),
Instant::from_millis(0),
&mut |_| {},
);
assert_eq!(frame_len, 0);
if frame_len > 0 {
device.count_rx(frame_len);
}
let snap = device.stats();
assert_eq!(snap.rx_bytes, 0);
assert_eq!(snap.rx_packets, 0);
}
#[test]
fn protocol_executor_three_path_combined_drain() {
let mut device = test_device_handle(Box::new(CountingMockDevice {
name: "mock",
send_returns: 0,
deferred_tx_lens: vec![60, 60], deferred_rx_lens: vec![42], recv_returns: 1514, }));
let frame_len = device.inner.recv(
IF0,
&mut test_packet_buffer(),
Instant::from_millis(0),
&mut |_| {},
);
if frame_len > 0 {
device.count_rx(frame_len);
}
for len in device.inner.drain_deferred_tx() {
device.count_tx(len);
}
for len in device.inner.drain_deferred_rx() {
device.count_rx(len);
}
let snap = device.stats();
assert_eq!(snap.rx_packets, 2);
assert_eq!(snap.rx_bytes, 1556);
assert_eq!(snap.tx_packets, 2);
assert_eq!(snap.tx_bytes, 120);
}
}