148 lines
3.8 KiB
Rust
148 lines
3.8 KiB
Rust
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<Box<TrieNode>>; 2],
|
|
prefix_len: u8,
|
|
nexthop: Option<Ipv4Addr>,
|
|
}
|
|
|
|
#[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<Interface> {
|
|
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::<Ipv4Addr>().unwrap(),
|
|
format!("192.168.{}.1", i).parse::<Ipv4Addr>().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<IpTrie>,
|
|
}
|
|
|
|
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));
|
|
}
|
|
}
|