use std::{net::Ipv4Addr, sync::Arc}; use arc_swap::ArcSwap; use default_net::Interface; use sdlan_sn_rs::utils::{Result, SDLanError}; #[derive(Default, Clone)] pub struct TrieNode { child: [Option>; 2], prefix_len: u8, nexthop: Option, } #[derive(Default, Clone)] struct IpTrie { root: TrieNode, } impl IpTrie { fn new() -> Self { Self { root: TrieNode::default(), } } fn insert(&mut self, prefix: u32, prefix_len: u8, nexthop: Ipv4Addr) { if prefix_len > 32 { return; } let mut node = &mut self.root; for i in 0..prefix_len { let bit = ((prefix >> (31 - i)) & 1) as usize; node = node.child[bit].get_or_insert_with(|| Box::new(TrieNode::default())); } if prefix_len >= node.prefix_len { node.prefix_len = prefix_len; node.nexthop = Some(nexthop); } } fn lookup(&self, ip: u32) -> Option<(u8, Ipv4Addr)> { let mut node = &self.root; let mut best = None; for i in 0..32 { if node.nexthop.is_some() { best = Some((node.prefix_len, node.nexthop.unwrap())); } let bit = ((ip >> (31 - i)) & 1) as usize; // println!("to find bit {} => {}", i, bit); match &node.child[bit] { Some(child) => { node = child; } None => { break; } } } if node.nexthop.is_some() { best = Some((node.prefix_len, node.nexthop.unwrap())); } best } } pub fn get_default_interface() -> Result { match default_net::get_default_interface() { Ok(interface) => Ok(interface), Err(e) => Err(SDLanError::IOError(e)), } } #[cfg(test)] mod test { use std::net::Ipv4Addr; use crate::{network::RouteTable2, TrieNode}; #[test] fn test_trie() { let trie = RouteTable2::new(); let mut ips: Vec<(Ipv4Addr, Ipv4Addr)> = vec![]; for i in 1..254u32 { ips.push(( format!("192.168.{}.0", i).parse::().unwrap(), format!("192.168.{}.1", i).parse::().unwrap(), )); } for ip in &ips { trie.route_table .insert(u32::from_be_bytes(ip.0.octets()), 24, ip.1); } for ip in &ips { let origin = ip.0; let mut query = origin.octets(); query[3] = 123; let query = Ipv4Addr::from_octets(query); println!("query for {}", query); let result = ip.1; let Some((prefix, target)) = trie.route_table.lookup(u32::from_be_bytes(query.octets())) else { panic!("failed to lookup: {} for {}", query, result); }; if prefix != 24 { panic!("prefix is not 24: {}", prefix); } if target != result { panic!("gateway is not match, {} expected {}", target, result); } } } } pub struct RouteTableTrie { trie: ArcSwap, } impl RouteTableTrie { pub fn new() -> Self { Self { trie: ArcSwap::new(Arc::new(IpTrie::default())), } } pub fn clear(&self) { self.trie.store(Arc::new(IpTrie::default())); } pub(crate) fn lookup(&self, ip: u32) -> Option<(u8, Ipv4Addr)> { let trie = self.trie.load(); trie.lookup(ip) } pub(crate) fn insert(&self, prefix: u32, prefix_len: u8, nexthop: Ipv4Addr) { let old = self.trie.load(); let mut new_trie = (*(*old)).clone(); new_trie.insert(prefix, prefix_len, nexthop); self.trie.store(Arc::new(new_trie)); } }