use crate::{Frame, L2Device, Result};
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, 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;
static PORT_ID_COUNTER: AtomicU64 = AtomicU64::new(0);
#[inline]
fn next_port_id() -> u64 {
PORT_ID_COUNTER.fetch_add(1, Ordering::Relaxed) + 1
}
struct Port {
dev: Arc<dyn L2Device>,
id: u64,
}
#[derive(Copy, Clone)]
struct MacEntry {
port_id: u64,
expires: Instant,
}
pub struct L2Hub {
ports: RwLock<Vec<Arc<Port>>>,
mac_table: Mutex<HashMap<[u8; 6], MacEntry>>,
}
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.len()).unwrap_or(0);
f.debug_struct("L2Hub").field("ports", &n).finish()
}
}
impl L2Hub {
pub fn new() -> L2Hub {
L2Hub {
ports: RwLock::new(Vec::new()),
mac_table: Mutex::new(HashMap::new()),
}
}
pub fn connect<D>(self: &Arc<Self>, dev: D) -> L2HubHandle
where
D: L2Device + 'static,
{
let id = next_port_id();
let dev_arc: Arc<dyn L2Device> = Arc::new(dev);
let port = Arc::new(Port {
dev: dev_arc.clone(),
id,
});
self.ports.write().unwrap().push(port);
let hub = Arc::downgrade(self);
dev_arc.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),
}
}
pub fn connect_arc(self: &Arc<Self>, dev: Arc<dyn L2Device>) -> L2HubHandle {
let id = next_port_id();
let port = Arc::new(Port {
dev: dev.clone(),
id,
});
self.ports.write().unwrap().push(port);
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),
}
}
fn forward(&self, f: &Frame, source_id: u64) {
let bytes = f.as_bytes();
if bytes.len() < 14 {
return;
}
let mut src = [0u8; 6];
src.copy_from_slice(&bytes[6..12]);
{
let mut table = self.mac_table.lock().unwrap();
match table.get(&src) {
Some(entry) if entry.port_id == source_id => {}
Some(_) => {
table.insert(
src,
MacEntry {
port_id: source_id,
expires: Instant::now() + MAC_AGING,
},
);
}
None => {
if table.len() < MAC_TABLE_MAX_SIZE {
table.insert(
src,
MacEntry {
port_id: source_id,
expires: Instant::now() + MAC_AGING,
},
);
}
}
}
}
let ports: Vec<Arc<Port>> = self.ports.read().unwrap().clone();
if bytes[0] & 1 != 0 {
for p in &ports {
if p.id != source_id {
let _ = p.dev.send(f);
}
}
return;
}
let mut dst = [0u8; 6];
dst.copy_from_slice(&bytes[0..6]);
let target = {
let mut table = self.mac_table.lock().unwrap();
match table.get(&dst).copied() {
Some(entry) if entry.expires > Instant::now() => Some(entry.port_id),
Some(_) => {
table.remove(&dst);
None
}
None => None,
}
};
if let Some(pid) = target {
for p in &ports {
if p.id == pid && p.id != source_id {
let _ = p.dev.send(f);
return;
}
}
}
for p in &ports {
if p.id != source_id {
let _ = p.dev.send(f);
}
}
}
fn disconnect(&self, id: u64) {
self.ports.write().unwrap().retain(|p| p.id != id);
self.mac_table
.lock()
.unwrap()
.retain(|_, e| e.port_id != 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::{build_frame, EtherType, L2Handler, MacAddr};
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);
}
}