use crate::{Frame, HubCounters, HubStats, L2Device, Result};
use std::cell::Cell;
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, RwLock};
use std::time::{Duration, Instant};
const MAC_AGING: Duration = Duration::from_secs(5 * 60);
const MAC_TABLE_MAX_SIZE: usize = 8192;
const DEFAULT_PORT_MAC_LIMIT: usize = 1024;
const DEFAULT_MAX_FORWARD_DEPTH: u32 = 16;
static PORT_ID_COUNTER: AtomicU64 = AtomicU64::new(0);
#[inline]
fn next_port_id() -> u64 {
PORT_ID_COUNTER.fetch_add(1, Ordering::Relaxed) + 1
}
thread_local! {
static FORWARD_DEPTH: Cell<u32> = const { Cell::new(0) };
}
struct DepthGuard;
impl DepthGuard {
fn enter(max: u32) -> Option<DepthGuard> {
FORWARD_DEPTH.with(|d| {
if d.get() >= max {
None
} else {
d.set(d.get() + 1);
Some(DepthGuard)
}
})
}
}
impl Drop for DepthGuard {
fn drop(&mut self) {
FORWARD_DEPTH.with(|d| d.set(d.get().saturating_sub(1)));
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct VlanSet(VlanSetInner);
#[derive(Clone, PartialEq, Eq)]
enum VlanSetInner {
All,
Bits(Box<[u64; 64]>),
}
impl VlanSet {
pub fn all() -> VlanSet {
VlanSet(VlanSetInner::All)
}
pub fn none() -> VlanSet {
VlanSet(VlanSetInner::Bits(Box::new([0; 64])))
}
fn bits_mut(&mut self) -> &mut [u64; 64] {
if matches!(self.0, VlanSetInner::All) {
let mut bits = Box::new([0u64; 64]);
bits.fill(u64::MAX);
self.0 = VlanSetInner::Bits(bits);
}
match &mut self.0 {
VlanSetInner::Bits(b) => b,
VlanSetInner::All => unreachable!("just replaced"),
}
}
pub fn from_ids(ids: impl IntoIterator<Item = u16>) -> VlanSet {
let mut set = VlanSet::none();
for id in ids {
set.insert(id);
}
set
}
pub fn insert(&mut self, vlan: u16) {
if vlan > 4095 {
return;
}
self.bits_mut()[(vlan / 64) as usize] |= 1 << (vlan % 64);
}
pub fn remove(&mut self, vlan: u16) {
if vlan > 4095 {
return;
}
self.bits_mut()[(vlan / 64) as usize] &= !(1 << (vlan % 64));
}
#[inline]
pub fn contains(&self, vlan: u16) -> bool {
if vlan > 4095 {
return false;
}
match &self.0 {
VlanSetInner::All => true,
VlanSetInner::Bits(b) => b[(vlan / 64) as usize] & (1 << (vlan % 64)) != 0,
}
}
}
impl core::fmt::Debug for VlanSet {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
if matches!(self.0, VlanSetInner::All) {
return f.write_str("VlanSet(all)");
}
let ids: Vec<u16> = (0u16..=4095).filter(|v| self.contains(*v)).collect();
write!(f, "VlanSet({ids:?})")
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PortMode {
Access { vlan: u16 },
Trunk {
allowed: VlanSet,
native: Option<u16>,
},
}
impl PortMode {
pub fn transparent() -> PortMode {
PortMode::Trunk {
allowed: VlanSet::all(),
native: Some(0),
}
}
fn ingress_vlan(&self, tagged: Option<u16>) -> Option<u16> {
match (self, tagged) {
(PortMode::Access { vlan }, None) => Some(*vlan),
(PortMode::Access { vlan }, Some(t)) if t == *vlan => Some(*vlan),
(PortMode::Access { .. }, Some(_)) => None,
(PortMode::Trunk { allowed, .. }, Some(t)) if allowed.contains(t) => Some(t),
(PortMode::Trunk { .. }, Some(_)) => None,
(PortMode::Trunk { allowed, native }, None) => native.filter(|n| allowed.contains(*n)),
}
}
fn egress(&self, vlan: u16) -> Option<TagAction> {
match self {
PortMode::Access { vlan: v } if *v == vlan => Some(TagAction::Untagged),
PortMode::Access { .. } => None,
PortMode::Trunk { allowed, native } => {
if !allowed.contains(vlan) {
return None;
}
if *native == Some(vlan) {
Some(TagAction::Untagged)
} else {
Some(TagAction::Tagged(vlan))
}
}
}
}
}
impl Default for PortMode {
fn default() -> PortMode {
PortMode::transparent()
}
}
#[derive(Debug, Copy, Clone, PartialEq, Eq)]
enum TagAction {
Untagged,
Tagged(u16),
}
struct Egress<'a> {
original: &'a Frame,
pcp: u8,
untagged: Option<Vec<u8>>,
tagged: Option<(u16, Vec<u8>)>,
}
impl<'a> Egress<'a> {
fn new(original: &'a Frame) -> Egress<'a> {
Egress {
pcp: original.vlan_pcp(),
original,
untagged: None,
tagged: None,
}
}
fn apply(&mut self, action: TagAction) -> &Frame {
match action {
TagAction::Untagged => {
if !self.original.has_vlan() {
return self.original;
}
let buf = self
.untagged
.get_or_insert_with(|| crate::build::pop_vlan(self.original));
Frame::from_slice(buf)
}
TagAction::Tagged(vlan) => {
if self.original.has_vlan() && self.original.vlan_id() == vlan {
return self.original;
}
if !matches!(&self.tagged, Some((v, _)) if *v == vlan) {
let base = if self.original.has_vlan() {
crate::build::pop_vlan(self.original)
} else {
self.original.to_vec()
};
let out = crate::build::push_vlan(Frame::from_slice(&base), vlan, self.pcp);
self.tagged = Some((vlan, out));
}
let (_, buf) = self.tagged.as_ref().expect("just built");
Frame::from_slice(buf)
}
}
}
}
struct Port {
dev: Arc<dyn L2Device>,
id: u64,
mac_limit: Option<usize>,
mode: PortMode,
}
struct PortSettings {
mac_limit: Option<usize>,
mode: PortMode,
}
struct PortTable {
list: Vec<Arc<Port>>,
by_id: HashMap<u64, Arc<Port>>,
}
impl PortTable {
fn from_list(list: Vec<Arc<Port>>) -> Arc<PortTable> {
let by_id = list.iter().map(|p| (p.id, p.clone())).collect();
Arc::new(PortTable { list, by_id })
}
#[inline]
fn get(&self, id: u64) -> Option<&Arc<Port>> {
self.by_id.get(&id)
}
}
#[derive(Clone)]
struct MacEntry {
port_id: u64,
expires: Instant,
}
type MacKey = (u16, [u8; 6]);
#[derive(Default)]
struct MacTable {
entries: HashMap<MacKey, MacEntry>,
per_port: HashMap<u64, usize>,
}
impl MacTable {
fn insert(&mut self, key: MacKey, entry: MacEntry) {
let new_port = entry.port_id;
if let Some(old) = self.entries.insert(key, entry) {
decrement(&mut self.per_port, old.port_id);
}
*self.per_port.entry(new_port).or_insert(0) += 1;
}
fn remove(&mut self, key: &MacKey) {
if let Some(old) = self.entries.remove(key) {
decrement(&mut self.per_port, old.port_id);
}
}
fn purge_port(&mut self, port_id: u64) {
self.entries.retain(|_, e| e.port_id != port_id);
self.per_port.remove(&port_id);
}
fn port_count(&self, port_id: u64) -> usize {
self.per_port.get(&port_id).copied().unwrap_or(0)
}
}
fn decrement(counts: &mut HashMap<u64, usize>, port_id: u64) {
if let Some(n) = counts.get_mut(&port_id) {
*n = n.saturating_sub(1);
if *n == 0 {
counts.remove(&port_id);
}
}
}
pub struct L2Hub {
ports: RwLock<Arc<PortTable>>,
mac_table: RwLock<MacTable>,
stats: HubStats,
loop_drops: AtomicU64,
max_depth: AtomicUsize,
}
impl Default for L2Hub {
fn default() -> Self {
Self::new()
}
}
impl core::fmt::Debug for L2Hub {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
let n = self.ports.read().map(|p| p.list.len()).unwrap_or(0);
f.debug_struct("L2Hub").field("ports", &n).finish()
}
}
impl L2Hub {
pub fn new() -> L2Hub {
L2Hub {
ports: RwLock::new(PortTable::from_list(Vec::new())),
mac_table: RwLock::new(MacTable::default()),
stats: HubStats::new(),
loop_drops: AtomicU64::new(0),
max_depth: AtomicUsize::new(DEFAULT_MAX_FORWARD_DEPTH as usize),
}
}
pub fn stats(&self) -> HubCounters {
self.stats.snapshot()
}
pub fn loop_drops(&self) -> u64 {
self.loop_drops.load(Ordering::Relaxed)
}
pub fn set_max_forward_depth(&self, depth: u32) {
self.max_depth
.store(depth.max(1) as usize, Ordering::Relaxed);
}
pub fn set_port_mac_limit(&self, handle: &L2HubHandle, limit: Option<usize>) {
self.reconfigure(handle.id, |p| p.mac_limit = limit);
}
pub fn set_port_mode(&self, handle: &L2HubHandle, mode: PortMode) {
self.reconfigure(handle.id, |p| p.mode = mode);
self.mac_table.write().unwrap().purge_port(handle.id);
}
pub fn port_mode(&self, handle: &L2HubHandle) -> Option<PortMode> {
self.ports().get(handle.id).map(|p| p.mode.clone())
}
fn reconfigure(&self, port_id: u64, edit: impl FnOnce(&mut PortSettings)) {
let mut guard = self.ports.write().unwrap();
let mut list = guard.list.clone();
let mut edit = Some(edit);
for slot in list.iter_mut() {
if slot.id != port_id {
continue;
}
let mut cfg = PortSettings {
mac_limit: slot.mac_limit,
mode: slot.mode.clone(),
};
if let Some(edit) = edit.take() {
edit(&mut cfg);
}
*slot = Arc::new(Port {
dev: slot.dev.clone(),
id: slot.id,
mac_limit: cfg.mac_limit,
mode: cfg.mode,
});
break;
}
*guard = PortTable::from_list(list);
}
pub fn mac_table_len(&self) -> usize {
self.mac_table.read().unwrap().entries.len()
}
pub fn connect<D>(self: &Arc<Self>, dev: D) -> L2HubHandle
where
D: L2Device + 'static,
{
self.connect_arc(Arc::new(dev))
}
pub fn connect_arc(self: &Arc<Self>, dev: Arc<dyn L2Device>) -> L2HubHandle {
let id = next_port_id();
{
let mut guard = self.ports.write().unwrap();
let mut list = guard.list.clone();
list.push(Arc::new(Port {
dev: dev.clone(),
id,
mac_limit: Some(DEFAULT_PORT_MAC_LIMIT),
mode: PortMode::transparent(),
}));
*guard = PortTable::from_list(list);
}
let hub = Arc::downgrade(self);
dev.set_handler(Arc::new(move |f: &Frame| {
if let Some(hub) = hub.upgrade() {
hub.forward(f, id);
}
Ok(())
}));
L2HubHandle {
hub: Arc::downgrade(self),
id,
closed: Mutex::new(false),
}
}
#[inline]
fn ports(&self) -> Arc<PortTable> {
self.ports.read().unwrap().clone()
}
fn forward(&self, f: &Frame, source_id: u64) {
self.stats.record_received();
let max = self.max_depth.load(Ordering::Relaxed) as u32;
let _depth = match DepthGuard::enter(max) {
Some(g) => g,
None => {
self.loop_drops.fetch_add(1, Ordering::Relaxed);
self.stats.record_dropped();
return;
}
};
let bytes = f.as_bytes();
if bytes.len() < 14 {
self.stats.record_dropped();
return;
}
let ports = self.ports();
let Some(source) = ports.get(source_id) else {
self.stats.record_dropped();
return;
};
let tag = f.has_vlan().then(|| f.vlan_id());
let Some(vlan) = source.mode.ingress_vlan(tag) else {
self.stats.record_dropped();
return;
};
let mut mac = [0u8; 6];
mac.copy_from_slice(&bytes[6..12]);
self.learn((vlan, mac), source);
let mut egress = Egress::new(f);
if bytes[0] & 1 != 0 {
self.flood(&ports, &mut egress, vlan, source_id);
return;
}
let mut dst_mac = [0u8; 6];
dst_mac.copy_from_slice(&bytes[0..6]);
if let Some(dst) = self.lookup(&ports, (vlan, dst_mac), source_id) {
if let Some(action) = dst.mode.egress(vlan) {
let _ = dst.dev.send(egress.apply(action));
self.stats.record_forwarded(1);
} else {
self.stats.record_dropped();
}
return;
}
self.flood(&ports, &mut egress, vlan, source_id);
}
fn learn(&self, key: MacKey, source: &Arc<Port>) {
let now = Instant::now();
{
let table = self.mac_table.read().unwrap();
if let Some(e) = table.entries.get(&key)
&& e.port_id == source.id
&& e.expires.saturating_duration_since(now) > MAC_AGING / 2
{
return;
}
}
let mut table = self.mac_table.write().unwrap();
let known = table.entries.get(&key).map(|e| e.port_id);
if known.is_none() {
if table.entries.len() >= MAC_TABLE_MAX_SIZE {
return;
}
if let Some(limit) = source.mac_limit
&& table.port_count(source.id) >= limit
{
return;
}
}
table.insert(
key,
MacEntry {
port_id: source.id,
expires: now + MAC_AGING,
},
);
}
fn lookup<'a>(
&self,
ports: &'a PortTable,
key: MacKey,
source_id: u64,
) -> Option<&'a Arc<Port>> {
let now = Instant::now();
let stale = {
let table = self.mac_table.read().unwrap();
match table.entries.get(&key) {
Some(e) if e.port_id == source_id => return None,
Some(e) if e.expires <= now => true,
Some(e) => match ports.get(e.port_id) {
Some(port) => return Some(port),
None => true,
},
None => false,
}
};
if stale {
self.mac_table.write().unwrap().remove(&key);
}
None
}
fn flood(&self, ports: &PortTable, egress: &mut Egress<'_>, vlan: u16, source_id: u64) {
let mut sent = 0u64;
for p in ports.list.iter() {
if p.id == source_id {
continue;
}
if let Some(action) = p.mode.egress(vlan) {
let _ = p.dev.send(egress.apply(action));
sent += 1;
}
}
if sent == 0 {
self.stats.record_dropped();
} else {
self.stats.record_flooded();
}
}
#[cfg(test)]
fn forward_from(&self, f: &Frame, port_id: u64) {
self.forward(f, port_id);
}
fn disconnect(&self, id: u64) {
{
let mut guard = self.ports.write().unwrap();
let mut list = guard.list.clone();
list.retain(|p| p.id != id);
*guard = PortTable::from_list(list);
}
self.mac_table.write().unwrap().purge_port(id);
}
}
impl crate::L2Connector for Arc<L2Hub> {
fn connect_l2(&self, dev: Arc<dyn L2Device>) -> Result<crate::Cleanup> {
let handle = self.connect_arc(dev);
let mut taken = Some(handle);
Ok(Box::new(move || {
if let Some(h) = taken.take() {
h.close();
}
Ok(())
}))
}
}
pub struct L2HubHandle {
hub: std::sync::Weak<L2Hub>,
id: u64,
closed: Mutex<bool>,
}
impl core::fmt::Debug for L2HubHandle {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("L2HubHandle").field("id", &self.id).finish()
}
}
impl L2HubHandle {
pub fn close(&self) {
let mut closed = self.closed.lock().unwrap();
if *closed {
return;
}
if let Some(hub) = self.hub.upgrade() {
hub.disconnect(self.id);
}
*closed = true;
}
}
impl Drop for L2HubHandle {
fn drop(&mut self) {
self.close();
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{EtherType, L2Handler, MacAddr, build_frame};
use std::sync::Mutex;
#[derive(Default, Clone)]
struct Sink {
inner: Arc<Mutex<Vec<Vec<u8>>>>,
mac: MacAddr,
}
impl L2Device for Sink {
fn set_handler(&self, _h: L2Handler) {}
fn send(&self, f: &Frame) -> Result<()> {
self.inner.lock().unwrap().push(f.as_bytes().to_vec());
Ok(())
}
fn hw_addr(&self) -> MacAddr {
self.mac
}
fn close(&self) -> Result<()> {
Ok(())
}
}
#[test]
fn broadcast_floods_to_all_except_source() {
let hub = Arc::new(L2Hub::new());
let a_mac: MacAddr = "02:00:00:00:00:01".parse().unwrap();
let b = Sink {
mac: "02:00:00:00:00:02".parse().unwrap(),
..Default::default()
};
let c = Sink {
mac: "02:00:00:00:00:03".parse().unwrap(),
..Default::default()
};
let a = Arc::new(crate::PipeL2::new(a_mac));
let _ha = hub.connect_arc(a.clone() as Arc<dyn L2Device>);
let _hb = hub.connect(b.clone());
let _hc = hub.connect(c.clone());
let buf = build_frame(MacAddr::broadcast(), a_mac, EtherType::IPV4, &[1, 2, 3]);
a.inject(Frame::from_slice(&buf)).unwrap();
assert_eq!(b.inner.lock().unwrap().len(), 1);
assert_eq!(c.inner.lock().unwrap().len(), 1);
}
#[derive(Default, Clone)]
struct Spy {
inner: Arc<Mutex<Vec<Vec<u8>>>>,
handler: Arc<Mutex<Option<L2Handler>>>,
mac: MacAddr,
}
impl L2Device for Spy {
fn set_handler(&self, h: L2Handler) {
*self.handler.lock().unwrap() = Some(h);
}
fn send(&self, f: &Frame) -> Result<()> {
self.inner.lock().unwrap().push(f.as_bytes().to_vec());
Ok(())
}
fn hw_addr(&self) -> MacAddr {
self.mac
}
fn close(&self) -> Result<()> {
Ok(())
}
}
impl Spy {
fn inject(&self, f: &Frame) {
let h = self.handler.lock().unwrap().clone();
if let Some(h) = h {
let _ = h(f);
}
}
fn count(&self) -> usize {
self.inner.lock().unwrap().len()
}
}
#[test]
fn learned_unicast_goes_to_one_port() {
let hub = Arc::new(L2Hub::new());
let a_mac: MacAddr = "02:00:00:00:00:01".parse().unwrap();
let b_mac: MacAddr = "02:00:00:00:00:02".parse().unwrap();
let c_mac: MacAddr = "02:00:00:00:00:03".parse().unwrap();
let a = Spy {
mac: a_mac,
..Default::default()
};
let b = Spy {
mac: b_mac,
..Default::default()
};
let c = Spy {
mac: c_mac,
..Default::default()
};
let _ha = hub.connect(a.clone());
let _hb = hub.connect(b.clone());
let _hc = hub.connect(c.clone());
let bf = build_frame(MacAddr::broadcast(), b_mac, EtherType::IPV4, &[0]);
b.inject(Frame::from_slice(&bf));
assert_eq!(c.count(), 1);
assert_eq!(a.count(), 1);
let ab = build_frame(b_mac, a_mac, EtherType::IPV4, &[1]);
a.inject(Frame::from_slice(&ab));
assert_eq!(b.count(), 1);
assert_eq!(c.count(), 1); }
#[test]
fn disconnect_removes_port() {
let hub = Arc::new(L2Hub::new());
let a = Arc::new(crate::PipeL2::new("02:00:00:00:00:01".parse().unwrap()));
let b = Sink {
mac: "02:00:00:00:00:02".parse().unwrap(),
..Default::default()
};
let _ha = hub.connect_arc(a.clone() as Arc<dyn L2Device>);
let hb = hub.connect(b.clone());
hb.close();
let bf = build_frame(MacAddr::broadcast(), MacAddr::zero(), EtherType::IPV4, &[]);
a.inject(Frame::from_slice(&bf)).unwrap();
assert_eq!(b.inner.lock().unwrap().len(), 0);
}
#[test]
fn stats_track_flood_forward_and_drop() {
let hub = Arc::new(L2Hub::new());
let a_mac: MacAddr = "02:00:00:00:00:01".parse().unwrap();
let b_mac: MacAddr = "02:00:00:00:00:02".parse().unwrap();
let a = Sink {
mac: a_mac,
..Default::default()
};
let b = Sink {
mac: b_mac,
..Default::default()
};
let ha = hub.connect(a.clone());
let hb = hub.connect(b.clone());
let bcast = build_frame(MacAddr::broadcast(), a_mac, EtherType::IPV4, &[0; 20]);
hub.forward_from(Frame::from_slice(&bcast), ha.id);
let s = hub.stats();
assert_eq!((s.received, s.flooded, s.forwarded), (1, 1, 0));
let unicast = build_frame(a_mac, b_mac, EtherType::IPV4, &[0; 20]);
hub.forward_from(Frame::from_slice(&unicast), hb.id);
let s = hub.stats();
assert_eq!((s.received, s.flooded, s.forwarded), (2, 1, 1));
hub.forward_from(Frame::from_slice(&[0u8; 4]), ha.id);
assert_eq!(hub.stats().dropped, 1);
}
#[test]
fn learning_is_vlan_aware() {
let hub = Arc::new(L2Hub::new());
let station: MacAddr = "02:00:00:00:00:aa".parse().unwrap();
let other: MacAddr = "02:00:00:00:00:bb".parse().unwrap();
let a = Sink {
mac: "02:00:00:00:00:01".parse().unwrap(),
..Default::default()
};
let b = Sink {
mac: "02:00:00:00:00:02".parse().unwrap(),
..Default::default()
};
let ha = hub.connect(a.clone());
let hb = hub.connect(b.clone());
let (a_id, b_id) = (ha.id, hb.id);
let untagged = build_frame(other, station, EtherType::IPV4, &[0; 20]);
hub.forward_from(Frame::from_slice(&untagged), a_id);
let tagged = crate::build::push_vlan(Frame::from_slice(&untagged), 5, 0);
hub.forward_from(Frame::from_slice(&tagged), b_id);
assert_eq!(hub.mac_table_len(), 2, "two VLANs, two entries");
a.inner.lock().unwrap().clear();
b.inner.lock().unwrap().clear();
let reply = build_frame(station, other, EtherType::IPV4, &[0; 20]);
hub.forward_from(Frame::from_slice(&reply), b_id);
assert_eq!(a.inner.lock().unwrap().len(), 1);
assert_eq!(b.inner.lock().unwrap().len(), 0);
let tagged_reply = crate::build::push_vlan(Frame::from_slice(&reply), 5, 0);
a.inner.lock().unwrap().clear();
hub.forward_from(Frame::from_slice(&tagged_reply), a_id);
assert_eq!(a.inner.lock().unwrap().len(), 0);
assert_eq!(b.inner.lock().unwrap().len(), 1);
}
#[derive(Default)]
struct Cable {
handler: Mutex<Option<L2Handler>>,
peer: Mutex<Option<Arc<Cable>>>,
mac: MacAddr,
}
impl Cable {
fn pair() -> (Arc<Cable>, Arc<Cable>) {
let a = Arc::new(Cable {
mac: "02:00:00:00:0c:01".parse().unwrap(),
..Default::default()
});
let b = Arc::new(Cable {
mac: "02:00:00:00:0c:02".parse().unwrap(),
..Default::default()
});
*a.peer.lock().unwrap() = Some(b.clone());
*b.peer.lock().unwrap() = Some(a.clone());
(a, b)
}
}
impl L2Device for Cable {
fn set_handler(&self, h: L2Handler) {
*self.handler.lock().unwrap() = Some(h);
}
fn send(&self, f: &Frame) -> Result<()> {
let peer = self.peer.lock().unwrap().clone();
if let Some(peer) = peer {
let h = peer.handler.lock().unwrap().clone();
if let Some(h) = h {
let _ = h(f);
}
}
Ok(())
}
fn hw_addr(&self) -> MacAddr {
self.mac
}
fn close(&self) -> Result<()> {
Ok(())
}
}
#[test]
fn frames_cross_a_chain_of_switches() {
let (a, b, c) = (
Arc::new(L2Hub::new()),
Arc::new(L2Hub::new()),
Arc::new(L2Hub::new()),
);
let (ab, ba) = Cable::pair();
let (bc, cb) = Cable::pair();
let _h1 = a.connect_arc(ab);
let _h2 = b.connect_arc(ba);
let _h3 = b.connect_arc(bc);
let _h4 = c.connect_arc(cb);
let left = Sink {
mac: "02:00:00:00:00:0a".parse().unwrap(),
..Default::default()
};
let right = Sink {
mac: "02:00:00:00:00:0c".parse().unwrap(),
..Default::default()
};
let hl = a.connect(left.clone());
let hr = c.connect(right.clone());
let announce = build_frame(MacAddr::broadcast(), right.mac, EtherType::IPV4, &[0; 40]);
c.forward_from(Frame::from_slice(&announce), hr.id);
assert!(
!left.inner.lock().unwrap().is_empty(),
"a broadcast must reach across three switches"
);
right.inner.lock().unwrap().clear();
let unicast = build_frame(right.mac, left.mac, EtherType::IPV4, &[0; 40]);
a.forward_from(Frame::from_slice(&unicast), hl.id);
assert_eq!(right.inner.lock().unwrap().len(), 1);
assert_eq!(a.loop_drops(), 0, "a chain is not a loop");
assert_eq!(b.loop_drops(), 0);
assert_eq!(c.loop_drops(), 0);
}
#[test]
fn a_topology_cycle_is_bounded_instead_of_overflowing_the_stack() {
let (a, b) = (Arc::new(L2Hub::new()), Arc::new(L2Hub::new()));
for _ in 0..2 {
let (x, y) = Cable::pair();
let _hx = a.connect_arc(x);
let _hy = b.connect_arc(y);
std::mem::forget((_hx, _hy));
}
let victim = Sink {
mac: "02:00:00:00:00:99".parse().unwrap(),
..Default::default()
};
let hv = a.connect(victim.clone());
let bcast = build_frame(MacAddr::broadcast(), victim.mac, EtherType::IPV4, &[0; 40]);
a.forward_from(Frame::from_slice(&bcast), hv.id);
assert!(
a.loop_drops() + b.loop_drops() > 0,
"the cycle should have been detected and cut"
);
}
#[test]
fn forward_depth_is_configurable_and_restored_after_each_frame() {
let hub = Arc::new(L2Hub::new());
hub.set_max_forward_depth(1);
let s = Sink {
mac: "02:00:00:00:00:01".parse().unwrap(),
..Default::default()
};
let h = hub.connect(s.clone());
let f = build_frame(MacAddr::broadcast(), s.mac, EtherType::IPV4, &[0; 40]);
for _ in 0..3 {
hub.forward_from(Frame::from_slice(&f), h.id);
}
assert_eq!(hub.loop_drops(), 0, "a depth of one is enough for one hop");
}
fn spray(hub: &Arc<L2Hub>, port: u64, n: u16) {
for i in 0..n {
let b = i.to_be_bytes();
let src = MacAddr::new([0x02, 0xff, 0, 0, b[0], b[1]]);
let f = build_frame(MacAddr::broadcast(), src, EtherType::IPV4, &[0; 40]);
hub.forward_from(Frame::from_slice(&f), port);
}
}
#[test]
fn one_port_flooding_addresses_cannot_starve_another() {
let hub = Arc::new(L2Hub::new());
let noisy = Sink {
mac: "02:00:00:00:00:01".parse().unwrap(),
..Default::default()
};
let quiet = Sink {
mac: "02:00:00:00:00:02".parse().unwrap(),
..Default::default()
};
let hn = hub.connect(noisy.clone());
let hq = hub.connect(quiet.clone());
spray(&hub, hn.id, 4000);
let after_spray = hub.mac_table_len();
assert!(
after_spray <= DEFAULT_PORT_MAC_LIMIT + 1,
"the per-port cap should have held the table to ~{DEFAULT_PORT_MAC_LIMIT}, got {after_spray}"
);
let f = build_frame(MacAddr::broadcast(), quiet.mac, EtherType::IPV4, &[0; 40]);
hub.forward_from(Frame::from_slice(&f), hq.id);
let unicast = build_frame(quiet.mac, noisy.mac, EtherType::IPV4, &[0; 40]);
quiet.inner.lock().unwrap().clear();
hub.forward_from(Frame::from_slice(&unicast), hn.id);
assert_eq!(
quiet.inner.lock().unwrap().len(),
1,
"the quiet station was never learned"
);
assert_eq!(hub.stats().forwarded, 1, "should be a forward, not a flood");
}
#[test]
fn an_uplink_can_learn_without_limit() {
let hub = Arc::new(L2Hub::new());
let uplink = Sink {
mac: "02:00:00:00:00:01".parse().unwrap(),
..Default::default()
};
let h = hub.connect(uplink.clone());
hub.set_port_mac_limit(&h, None);
spray(&hub, h.id, 3000);
assert!(
hub.mac_table_len() > DEFAULT_PORT_MAC_LIMIT,
"an unlimited port should learn past the per-port cap, got {}",
hub.mac_table_len()
);
}
#[test]
fn disconnecting_a_port_drops_its_learned_addresses() {
let hub = Arc::new(L2Hub::new());
let a = Sink {
mac: "02:00:00:00:00:01".parse().unwrap(),
..Default::default()
};
let b = Sink {
mac: "02:00:00:00:00:02".parse().unwrap(),
..Default::default()
};
let ha = hub.connect(a.clone());
let _hb = hub.connect(b.clone());
spray(&hub, ha.id, 50);
assert!(hub.mac_table_len() >= 50);
ha.close();
assert_eq!(
hub.mac_table_len(),
0,
"entries pointing at a removed port must go with it"
);
}
#[test]
fn a_station_that_moves_ports_is_relearned() {
let hub = Arc::new(L2Hub::new());
let left = Sink {
mac: "02:00:00:00:00:01".parse().unwrap(),
..Default::default()
};
let right = Sink {
mac: "02:00:00:00:00:02".parse().unwrap(),
..Default::default()
};
let observer = Sink {
mac: "02:00:00:00:00:03".parse().unwrap(),
..Default::default()
};
let hl = hub.connect(left.clone());
let hr = hub.connect(right.clone());
let ho = hub.connect(observer.clone());
let station: MacAddr = "02:00:00:00:aa:aa".parse().unwrap();
let announce = build_frame(MacAddr::broadcast(), station, EtherType::IPV4, &[0; 40]);
hub.forward_from(Frame::from_slice(&announce), hl.id);
hub.forward_from(Frame::from_slice(&announce), hr.id);
assert_eq!(hub.mac_table_len(), 1, "a move replaces, it does not add");
left.inner.lock().unwrap().clear();
right.inner.lock().unwrap().clear();
let to_station = build_frame(station, observer.mac, EtherType::IPV4, &[0; 40]);
hub.forward_from(Frame::from_slice(&to_station), ho.id);
assert_eq!(right.inner.lock().unwrap().len(), 1);
assert_eq!(left.inner.lock().unwrap().len(), 0);
}
#[test]
fn a_frame_is_never_sent_back_out_of_the_port_it_arrived_on() {
let hub = Arc::new(L2Hub::new());
let a = Sink {
mac: "02:00:00:00:00:01".parse().unwrap(),
..Default::default()
};
let ha = hub.connect(a.clone());
let station: MacAddr = "02:00:00:00:aa:aa".parse().unwrap();
let announce = build_frame(MacAddr::broadcast(), station, EtherType::IPV4, &[0; 40]);
hub.forward_from(Frame::from_slice(&announce), ha.id);
a.inner.lock().unwrap().clear();
let f = build_frame(station, MacAddr::broadcast(), EtherType::IPV4, &[0; 40]);
hub.forward_from(Frame::from_slice(&f), ha.id);
assert_eq!(a.inner.lock().unwrap().len(), 0);
}
fn access(vlan: u16) -> PortMode {
PortMode::Access { vlan }
}
fn trunk(ids: &[u16], native: Option<u16>) -> PortMode {
PortMode::Trunk {
allowed: VlanSet::from_ids(ids.iter().copied()),
native,
}
}
fn sinks(hub: &Arc<L2Hub>, n: u16) -> Vec<(Sink, L2HubHandle)> {
(0..n)
.map(|i| {
let b = i.to_be_bytes();
let s = Sink {
mac: MacAddr::new([0x02, 0, 0, 0, b[0], b[1]]),
..Default::default()
};
let h = hub.connect(s.clone());
(s, h)
})
.collect()
}
#[test]
fn vlan_set_membership() {
let all = VlanSet::all();
assert!(all.contains(0) && all.contains(4095));
assert!(!all.contains(4096), "4095 is the largest 802.1Q id");
let mut set = VlanSet::from_ids([1, 100, 4095]);
assert!(set.contains(1) && set.contains(100) && set.contains(4095));
assert!(!set.contains(2) && !set.contains(0));
set.remove(100);
assert!(!set.contains(100));
set.insert(9999);
assert!(!set.contains(9999 & 4095));
let mut all = VlanSet::all();
all.remove(7);
assert!(!all.contains(7));
assert!(all.contains(6) && all.contains(8) && all.contains(4095));
}
#[test]
fn access_ports_on_different_vlans_are_isolated() {
let hub = Arc::new(L2Hub::new());
let ports = sinks(&hub, 3);
hub.set_port_mode(&ports[0].1, access(10));
hub.set_port_mode(&ports[1].1, access(10));
hub.set_port_mode(&ports[2].1, access(20));
let f = build_frame(
MacAddr::broadcast(),
ports[0].0.mac,
EtherType::IPV4,
&[0; 40],
);
hub.forward_from(Frame::from_slice(&f), ports[0].1.id);
assert_eq!(
ports[1].0.inner.lock().unwrap().len(),
1,
"same VLAN should receive"
);
assert_eq!(
ports[2].0.inner.lock().unwrap().len(),
0,
"a different VLAN must not"
);
}
#[test]
fn an_access_port_receives_untagged_and_a_trunk_receives_tagged() {
let hub = Arc::new(L2Hub::new());
let ports = sinks(&hub, 3);
hub.set_port_mode(&ports[0].1, access(10));
hub.set_port_mode(&ports[1].1, access(10));
hub.set_port_mode(&ports[2].1, trunk(&[10, 20], None));
let f = build_frame(
MacAddr::broadcast(),
ports[0].0.mac,
EtherType::IPV4,
&[0; 40],
);
hub.forward_from(Frame::from_slice(&f), ports[0].1.id);
let at_access = ports[1].0.inner.lock().unwrap()[0].clone();
let seen = Frame::from_slice(&at_access);
assert!(!seen.has_vlan(), "an access port must never see a tag");
assert_eq!(seen.ether_type(), EtherType::IPV4);
assert_eq!(seen.payload(), &[0u8; 40]);
let at_trunk = ports[2].0.inner.lock().unwrap()[0].clone();
let seen = Frame::from_slice(&at_trunk);
assert!(seen.has_vlan(), "a trunk carries the tag");
assert_eq!(seen.vlan_id(), 10);
assert_eq!(seen.ether_type(), EtherType::IPV4);
assert_eq!(seen.payload(), &[0u8; 40]);
}
#[test]
fn a_trunk_drops_vlans_it_does_not_carry() {
let hub = Arc::new(L2Hub::new());
let ports = sinks(&hub, 2);
hub.set_port_mode(&ports[0].1, trunk(&[10, 20], None));
hub.set_port_mode(&ports[1].1, trunk(&[10], None));
let base = build_frame(
MacAddr::broadcast(),
ports[0].0.mac,
EtherType::IPV4,
&[0; 40],
);
let on_20 = crate::build::push_vlan(Frame::from_slice(&base), 20, 0);
hub.forward_from(Frame::from_slice(&on_20), ports[0].1.id);
assert_eq!(
ports[1].0.inner.lock().unwrap().len(),
0,
"VLAN 20 is not allowed on that trunk"
);
let on_10 = crate::build::push_vlan(Frame::from_slice(&base), 10, 0);
hub.forward_from(Frame::from_slice(&on_10), ports[0].1.id);
assert_eq!(ports[1].0.inner.lock().unwrap().len(), 1);
}
#[test]
fn untagged_frames_on_a_trunk_need_a_native_vlan() {
let hub = Arc::new(L2Hub::new());
let ports = sinks(&hub, 2);
hub.set_port_mode(&ports[1].1, access(10));
hub.set_port_mode(&ports[0].1, trunk(&[10], None));
let f = build_frame(
MacAddr::broadcast(),
ports[0].0.mac,
EtherType::IPV4,
&[0; 40],
);
hub.forward_from(Frame::from_slice(&f), ports[0].1.id);
assert_eq!(ports[1].0.inner.lock().unwrap().len(), 0);
assert!(hub.stats().dropped >= 1);
hub.set_port_mode(&ports[0].1, trunk(&[10], Some(10)));
hub.forward_from(Frame::from_slice(&f), ports[0].1.id);
assert_eq!(ports[1].0.inner.lock().unwrap().len(), 1);
let back = build_frame(
MacAddr::broadcast(),
ports[1].0.mac,
EtherType::IPV4,
&[0; 40],
);
hub.forward_from(Frame::from_slice(&back), ports[1].1.id);
let at_trunk = ports[0].0.inner.lock().unwrap().last().unwrap().clone();
assert!(
!Frame::from_slice(&at_trunk).has_vlan(),
"native is untagged"
);
}
#[test]
fn an_access_port_rejects_a_foreign_tag() {
let hub = Arc::new(L2Hub::new());
let ports = sinks(&hub, 2);
hub.set_port_mode(&ports[0].1, access(10));
hub.set_port_mode(&ports[1].1, trunk(&[10, 20], None));
let base = build_frame(
MacAddr::broadcast(),
ports[0].0.mac,
EtherType::IPV4,
&[0; 40],
);
let tagged_20 = crate::build::push_vlan(Frame::from_slice(&base), 20, 0);
hub.forward_from(Frame::from_slice(&tagged_20), ports[0].1.id);
assert_eq!(
ports[1].0.inner.lock().unwrap().len(),
0,
"an access port claiming VLAN 10 cannot inject VLAN 20"
);
let tagged_10 = crate::build::push_vlan(Frame::from_slice(&base), 10, 0);
hub.forward_from(Frame::from_slice(&tagged_10), ports[0].1.id);
assert_eq!(ports[1].0.inner.lock().unwrap().len(), 1);
}
#[test]
fn two_switches_trunked_together_keep_vlans_apart() {
let (left, right) = (Arc::new(L2Hub::new()), Arc::new(L2Hub::new()));
let (cable_l, cable_r) = Cable::pair();
let hl_trunk = left.connect_arc(cable_l);
let hr_trunk = right.connect_arc(cable_r);
left.set_port_mode(&hl_trunk, trunk(&[10, 20], None));
right.set_port_mode(&hr_trunk, trunk(&[10, 20], None));
left.set_port_mac_limit(&hl_trunk, None);
right.set_port_mac_limit(&hr_trunk, None);
let l10 = Sink {
mac: "02:00:00:00:10:01".parse().unwrap(),
..Default::default()
};
let r10 = Sink {
mac: "02:00:00:00:10:02".parse().unwrap(),
..Default::default()
};
let r20 = Sink {
mac: "02:00:00:00:20:02".parse().unwrap(),
..Default::default()
};
let h_l10 = left.connect(l10.clone());
let h_r10 = right.connect(r10.clone());
let h_r20 = right.connect(r20.clone());
left.set_port_mode(&h_l10, access(10));
right.set_port_mode(&h_r10, access(10));
right.set_port_mode(&h_r20, access(20));
let f = build_frame(MacAddr::broadcast(), l10.mac, EtherType::IPV4, &[0; 40]);
left.forward_from(Frame::from_slice(&f), h_l10.id);
assert_eq!(r10.inner.lock().unwrap().len(), 1, "VLAN 10 must cross");
let arrived = r10.inner.lock().unwrap()[0].clone();
assert!(!Frame::from_slice(&arrived).has_vlan());
assert_eq!(
r20.inner.lock().unwrap().len(),
0,
"VLAN 20 must not see it"
);
r10.inner.lock().unwrap().clear();
let reply = build_frame(l10.mac, r10.mac, EtherType::IPV4, &[0; 40]);
let before = right.stats().forwarded;
right.forward_from(Frame::from_slice(&reply), h_r10.id);
assert_eq!(
right.stats().forwarded,
before + 1,
"the trunk should be a learned destination, not a flood"
);
assert_eq!(l10.inner.lock().unwrap().len(), 1, "the reply came back");
}
#[test]
fn the_same_address_on_two_vlans_is_two_stations() {
let hub = Arc::new(L2Hub::new());
let ports = sinks(&hub, 4);
hub.set_port_mode(&ports[0].1, access(10));
hub.set_port_mode(&ports[1].1, access(20));
hub.set_port_mode(&ports[2].1, access(10));
hub.set_port_mode(&ports[3].1, access(20));
let station: MacAddr = "02:00:00:00:aa:aa".parse().unwrap();
let announce = build_frame(MacAddr::broadcast(), station, EtherType::IPV4, &[0; 40]);
hub.forward_from(Frame::from_slice(&announce), ports[0].1.id);
hub.forward_from(Frame::from_slice(&announce), ports[1].1.id);
assert_eq!(hub.mac_table_len(), 2, "one per VLAN, not a port flap");
for (sender, expect, other) in [(2usize, 0usize, 1usize), (3, 1, 0)] {
for (s, _) in ports.iter() {
s.inner.lock().unwrap().clear();
}
let f = build_frame(station, ports[sender].0.mac, EtherType::IPV4, &[0; 40]);
hub.forward_from(Frame::from_slice(&f), ports[sender].1.id);
assert_eq!(ports[expect].0.inner.lock().unwrap().len(), 1);
assert_eq!(ports[other].0.inner.lock().unwrap().len(), 0);
}
}
#[test]
fn transparent_ports_pass_tags_through_untouched() {
let hub = Arc::new(L2Hub::new());
let ports = sinks(&hub, 2);
assert_eq!(hub.port_mode(&ports[0].1), Some(PortMode::transparent()));
let base = build_frame(
MacAddr::broadcast(),
ports[0].0.mac,
EtherType::IPV4,
&[0; 40],
);
let tagged = crate::build::push_vlan(Frame::from_slice(&base), 77, 5);
hub.forward_from(Frame::from_slice(&tagged), ports[0].1.id);
let out = ports[1].0.inner.lock().unwrap()[0].clone();
assert_eq!(out, tagged, "byte-identical, tag and priority included");
hub.forward_from(Frame::from_slice(&base), ports[0].1.id);
let out = ports[1].0.inner.lock().unwrap()[1].clone();
assert_eq!(out, base);
}
#[test]
fn a_reconfigured_port_still_forwards() {
let hub = Arc::new(L2Hub::new());
let a = Spy {
mac: "02:00:00:00:00:01".parse().unwrap(),
..Default::default()
};
let b = Spy {
mac: "02:00:00:00:00:02".parse().unwrap(),
..Default::default()
};
let ha = hub.connect(a.clone());
let hb = hub.connect(b.clone());
hub.set_port_mode(&ha, access(10));
hub.set_port_mode(&hb, access(10));
hub.set_port_mac_limit(&ha, None);
let f = build_frame(MacAddr::broadcast(), a.mac, EtherType::IPV4, &[0; 40]);
a.inject(Frame::from_slice(&f));
assert_eq!(
b.count(),
1,
"a port that has been reconfigured must still forward"
);
}
#[test]
fn changing_a_port_mode_forgets_what_it_learned() {
let hub = Arc::new(L2Hub::new());
let ports = sinks(&hub, 2);
let f = build_frame(
MacAddr::broadcast(),
ports[0].0.mac,
EtherType::IPV4,
&[0; 40],
);
hub.forward_from(Frame::from_slice(&f), ports[0].1.id);
assert_eq!(hub.mac_table_len(), 1);
hub.set_port_mode(&ports[0].1, access(10));
assert_eq!(hub.mac_table_len(), 0);
assert_eq!(hub.port_mode(&ports[0].1), Some(access(10)));
}
}