sdlan-lib-rs/src/utils/system_action.rs
2026-04-22 23:27:56 +08:00

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));
}
}