use core::{net::IpAddr, pin::Pin};
use std::collections::{BTreeMap, BTreeSet, HashSet};
use futures_util::{Stream, StreamExt, stream};
#[cfg(target_os = "macos")]
pub use crate::darwin::AfRouteMon as PlatformMon;
#[cfg(target_os = "linux")]
pub use crate::linux::RtNetlinkMon as PlatformMon;
#[cfg(windows)]
pub use crate::windows::Winmon as PlatformMon;
use crate::{Event, Family, InterfaceId, Route, id::MonType};
pub type BoxStream<T> = Pin<Box<dyn Stream<Item = T> + Send + 'static>>;
pub const fn platform_mon() -> Option<impl Netmon + 'static> {
cfg_if::cfg_if! {
if #[cfg(any(windows, target_os = "linux"))] {
Some(&PlatformMon)
} else {
struct NoopMon;
impl Netmon for NoopMon {
fn ty(&self) -> MonType {
unimplemented!()
}
fn event_stream(&self) -> std::io::Result<BoxStream<std::io::Result<Event>>> {
unimplemented!()
}
}
Option::<NoopMon>::None
}
}
}
pub trait Netmon: Send + Sync {
fn ty(&self) -> MonType;
fn event_stream(&self) -> std::io::Result<BoxStream<std::io::Result<Event>>>;
fn strong_delete_consistency(&self) -> bool {
true
}
fn interface_unique_addrs(&self) -> bool {
false
}
}
impl<T> Netmon for &T
where
T: Netmon + ?Sized,
{
fn ty(&self) -> MonType {
T::ty(self)
}
fn event_stream(&self) -> std::io::Result<BoxStream<std::io::Result<Event>>> {
T::event_stream(self)
}
fn strong_delete_consistency(&self) -> bool {
T::strong_delete_consistency(self)
}
fn interface_unique_addrs(&self) -> bool {
T::interface_unique_addrs(self)
}
}
impl dyn Netmon {
pub fn with_default_route_events(
&self,
) -> std::io::Result<impl Stream<Item = std::io::Result<Event>> + Send + use<>> {
let strong_delete_consistency = self.strong_delete_consistency();
let s = self
.event_stream()?
.filter(|x| {
let result = !x
.as_ref()
.is_ok_and(|x| matches!(x, Event::DefaultRouteInterface(..)));
async move { result }
})
.scan(
(
DefaultRouteState::new(Family::Ipv4),
DefaultRouteState::new(Family::Ipv6),
),
move |(state_v4, state_v6), x| {
let [e1, e2] = match &x {
Ok(Event::RouteUpsert(interface, route)) => [
state_v4.add_route(interface, route),
state_v6.add_route(interface, route),
],
Ok(Event::RouteRemoved(interface, route)) => [
state_v4.remove_route(interface, route),
state_v6.remove_route(interface, route),
],
Ok(Event::InterfaceUpsert(interface)) => [
state_v4.update_interface_state(&interface.id, interface.up),
state_v6.update_interface_state(&interface.id, interface.up),
],
Ok(Event::InterfaceRemoved(i)) => {
if strong_delete_consistency {
let e1 = state_v4.remove_interface(i);
let e2 = state_v6.remove_interface(i);
[e1, e2]
} else {
[None, None]
}
}
_ => [None, None],
};
async move {
Some(
stream::once(async move { x })
.chain(stream::iter(e1).map(Ok))
.chain(stream::iter(e2).map(Ok)),
)
}
},
)
.flatten();
Ok(s)
}
}
type DefaultRouteUnique = (InterfaceId, smallvec::SmallVec<[IpAddr; 1]>);
#[derive(Debug)]
struct DefaultRouteState {
metrics: BTreeMap<usize, BTreeSet<DefaultRouteUnique>>,
interfaces_up: HashSet<InterfaceId>,
last_id: Option<InterfaceId>,
family: Family,
}
impl DefaultRouteState {
fn new(family: Family) -> Self {
Self {
metrics: Default::default(),
last_id: None,
interfaces_up: Default::default(),
family,
}
}
fn add_route(&mut self, interface_id: &InterfaceId, route: &Route) -> Option<Event> {
if route.family() != self.family || !route.is_default_route() {
return None;
}
let mut gws = route.gateway.clone();
gws.sort();
let modified = self
.metrics
.entry(route.metric)
.or_default()
.insert((interface_id.clone(), gws));
if !modified {
return None;
}
self.update_best()
}
fn remove_route(&mut self, interface_id: &InterfaceId, route: &Route) -> Option<Event> {
if route.family() != self.family || !route.is_default_route() {
return None;
}
let entry = self.metrics.get_mut(&route.metric)?;
let mut gws = route.gateway.clone();
gws.sort();
let modified = entry.remove(&(interface_id.clone(), gws));
if !modified {
return None;
}
if entry.is_empty() {
self.metrics.remove(&route.metric);
}
self.update_best()
}
fn update_interface_state(&mut self, interface_id: &InterfaceId, up: bool) -> Option<Event> {
if up {
self.interfaces_up.insert(interface_id.clone())
} else {
self.interfaces_up.remove(interface_id)
}
.then(|| self.update_best())
.flatten()
}
fn remove_interface(&mut self, interface_id: &InterfaceId) -> Option<Event> {
let mut metrics_modified = false;
self.metrics.retain(|_, e| {
e.retain(|(id, _)| {
let ret = id == interface_id;
metrics_modified = metrics_modified || ret;
ret
});
!e.is_empty()
});
self.interfaces_up.remove(interface_id);
if !metrics_modified {
return None;
}
self.update_best()
}
fn update_best(&mut self) -> Option<Event> {
let new_best_id = self
.metrics
.values()
.flatten()
.find_map(|(id, _)| self.interfaces_up.contains(id).then(|| id.clone()));
if new_best_id == self.last_id {
return None;
}
self.last_id = new_best_id.clone();
Some(Event::DefaultRouteInterface(new_best_id, self.family))
}
}
#[cfg(test)]
mod test {
use super::*;
const MONTYPE: MonType = MonType::new_static("test");
#[test]
fn different_gateway() -> Result<(), Box<dyn core::error::Error>> {
let mut state = DefaultRouteState::new(Family::Ipv4);
let interface = InterfaceId::new(MONTYPE, 0);
let evt1 = state.update_interface_state(&interface, true);
let route1 = Route {
metric: 0,
gateway: smallvec::smallvec!["1.2.3.4".parse()?],
dst: ipnet::Ipv4Net::default().into(),
};
let route2 = Route {
metric: 0,
gateway: smallvec::smallvec!["5.6.7.8".parse()?],
dst: ipnet::Ipv4Net::default().into(),
};
assert!(route1.is_default_route() && route2.is_default_route());
let evt2 = state.add_route(&interface, &route1);
let evt3 = state.add_route(&interface, &route2);
assert_eq!(state.last_id, Some(interface.clone()));
assert_eq!(evt1, None);
assert_eq!(
evt2,
Some(Event::DefaultRouteInterface(
Some(interface.clone()),
Family::Ipv4
))
);
assert_eq!(evt3, None);
let rem1 = state.remove_route(&interface, &route1);
let rem2 = state.remove_route(&interface, &route2);
assert_eq!(state.last_id, None);
assert_eq!(rem1, None);
assert_eq!(rem2, Some(Event::DefaultRouteInterface(None, Family::Ipv4)));
Ok(())
}
}