use std::collections::HashSet;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::Mutex;
use std::time::{Duration, Instant};
use ipnetwork::IpNetwork;
use log::{debug, info, warn};
pub const CONTROL_MARKER: u8 = 0x00;
const TYPE_ROUTE_ADVERT: u8 = 0x01;
const TYPE_ROUTE_PUSH: u8 = 0x02;
const FLAG_ACCEPT_ROUTES: u8 = 0x01;
pub const MAX_ROUTES: usize = 64;
const ADVERT_HEADER_LEN: usize = 3 + 4 + 16 + 1;
const PUSH_HEADER_LEN: usize = 4;
pub fn is_control(payload: &[u8]) -> bool {
payload.first() == Some(&CONTROL_MARKER)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Control {
Keepalive(Option<Ipv4Addr>),
RouteAdvert(RouteAdvert),
RoutePush(RoutePush),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RouteAdvert {
pub tunnel_ip: Ipv4Addr,
pub tunnel_ip6: Option<Ipv6Addr>,
pub accept_routes: bool,
pub routes: Vec<IpNetwork>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RoutePush {
pub routes: Vec<IpNetwork>,
}
pub fn canonical(net: IpNetwork) -> IpNetwork {
IpNetwork::new(net.network(), net.prefix()).expect("network address + same prefix is valid")
}
fn push_route(buf: &mut Vec<u8>, net: &IpNetwork) {
match canonical(*net) {
IpNetwork::V4(n) => {
buf.push(4);
buf.push(n.prefix());
buf.extend_from_slice(&n.ip().octets());
}
IpNetwork::V6(n) => {
buf.push(6);
buf.push(n.prefix());
buf.extend_from_slice(&n.ip().octets());
}
}
}
fn parse_routes(mut buf: &[u8], count: usize) -> Option<Vec<IpNetwork>> {
let mut routes = Vec::with_capacity(count);
for _ in 0..count {
let (&family, rest) = buf.split_first()?;
let (&plen, rest) = rest.split_first()?;
let (addr, rest): (IpAddr, &[u8]) = match family {
4 => {
let (bytes, rest) = rest.split_first_chunk::<4>()?;
(Ipv4Addr::from(*bytes).into(), rest)
}
6 => {
let (bytes, rest) = rest.split_first_chunk::<16>()?;
(Ipv6Addr::from(*bytes).into(), rest)
}
_ => return None,
};
routes.push(canonical(IpNetwork::new(addr, plen).ok()?));
buf = rest;
}
buf.is_empty().then_some(routes)
}
impl RouteAdvert {
pub fn encode(&self) -> Vec<u8> {
let routes = &self.routes[..self.routes.len().min(MAX_ROUTES)];
let mut buf = Vec::with_capacity(ADVERT_HEADER_LEN + routes.len() * 18);
buf.push(CONTROL_MARKER);
buf.push(TYPE_ROUTE_ADVERT);
buf.push(if self.accept_routes {
FLAG_ACCEPT_ROUTES
} else {
0
});
buf.extend_from_slice(&self.tunnel_ip.octets());
buf.extend_from_slice(&self.tunnel_ip6.unwrap_or(Ipv6Addr::UNSPECIFIED).octets());
buf.push(routes.len() as u8);
for net in routes {
push_route(&mut buf, net);
}
buf
}
}
impl RoutePush {
pub fn encode(&self) -> Vec<u8> {
let routes = &self.routes[..self.routes.len().min(MAX_ROUTES)];
let mut buf = Vec::with_capacity(PUSH_HEADER_LEN + routes.len() * 18);
buf.push(CONTROL_MARKER);
buf.push(TYPE_ROUTE_PUSH);
buf.push(0); buf.push(routes.len() as u8);
for net in routes {
push_route(&mut buf, net);
}
buf
}
}
pub fn parse_control(payload: &[u8]) -> Option<Control> {
if !is_control(payload) {
return None;
}
match payload {
[_] => return Some(Control::Keepalive(None)),
[_, a, b, c, d] => return Some(Control::Keepalive(Some(Ipv4Addr::new(*a, *b, *c, *d)))),
_ => {}
}
match *payload.get(1)? {
TYPE_ROUTE_ADVERT => {
let body = payload.get(2..)?;
if body.len() < ADVERT_HEADER_LEN - 2 {
return None;
}
let flags = body[0];
let tunnel_ip = Ipv4Addr::new(body[1], body[2], body[3], body[4]);
let ip6 = Ipv6Addr::from(<[u8; 16]>::try_from(&body[5..21]).expect("16 bytes"));
let count = body[21] as usize;
if count > MAX_ROUTES {
return None;
}
let routes = parse_routes(&body[22..], count)?;
Some(Control::RouteAdvert(RouteAdvert {
tunnel_ip,
tunnel_ip6: (!ip6.is_unspecified()).then_some(ip6),
accept_routes: flags & FLAG_ACCEPT_ROUTES != 0,
routes,
}))
}
TYPE_ROUTE_PUSH => {
let count = *payload.get(3)? as usize;
if count > MAX_ROUTES {
return None;
}
let routes = parse_routes(payload.get(4..)?, count)?;
Some(Control::RoutePush(RoutePush { routes }))
}
_ => None,
}
}
#[derive(Debug, Clone, Default)]
pub struct RouteApproval {
pub auto: bool,
pub allowlist: Vec<IpNetwork>,
}
impl RouteApproval {
pub fn approves(&self, net: &IpNetwork) -> bool {
if self.auto {
return true;
}
self.allowlist.iter().any(|allow| match (allow, net) {
(IpNetwork::V4(a), IpNetwork::V4(n)) => a.is_supernet_of(*n),
(IpNetwork::V6(a), IpNetwork::V6(n)) => a.is_supernet_of(*n),
_ => false,
})
}
}
#[derive(Debug, Clone)]
struct SubnetRoute {
net: IpNetwork,
peer: SocketAddr,
approved: bool,
last_seen: Instant,
}
#[derive(Debug, Default, PartialEq, Eq)]
pub struct AdvertOutcome {
pub approved: Vec<IpNetwork>,
pub awaiting: Vec<IpNetwork>,
pub moved: Vec<IpNetwork>,
pub withdrawn: Vec<IpNetwork>,
}
impl AdvertOutcome {
pub fn is_quiet(&self) -> bool {
self.approved.is_empty()
&& self.awaiting.is_empty()
&& self.moved.is_empty()
&& self.withdrawn.is_empty()
}
}
#[derive(Debug, Default)]
pub struct SubnetTable {
routes: Vec<SubnetRoute>,
}
impl SubnetTable {
pub fn advertise(
&mut self,
peer: SocketAddr,
advertised: &[IpNetwork],
approval: &RouteApproval,
now: Instant,
) -> AdvertOutcome {
let mut outcome = AdvertOutcome::default();
let advertised: Vec<IpNetwork> = advertised.iter().copied().map(canonical).collect();
self.routes.retain(|r| {
let keep = r.peer != peer || advertised.contains(&r.net);
if !keep {
outcome.withdrawn.push(r.net);
}
keep
});
for net in advertised {
if let Some(existing) = self.routes.iter_mut().find(|r| r.net == net) {
if existing.peer != peer {
existing.peer = peer;
outcome.moved.push(net);
}
existing.approved = approval.approves(&net);
existing.last_seen = now;
} else {
let approved = approval.approves(&net);
self.routes.push(SubnetRoute {
net,
peer,
approved,
last_seen: now,
});
if approved {
outcome.approved.push(net);
} else {
outcome.awaiting.push(net);
}
}
}
outcome
}
pub fn lookup(&self, dst: IpAddr) -> Option<SocketAddr> {
self.routes
.iter()
.filter(|r| r.approved && r.net.contains(dst))
.max_by_key(|r| r.net.prefix())
.map(|r| r.peer)
}
pub fn routes_for(&self, peer: SocketAddr) -> Vec<IpNetwork> {
self.routes
.iter()
.filter(|r| r.approved && r.peer != peer)
.map(|r| r.net)
.collect()
}
pub fn expire(&mut self, ttl: Duration, now: Instant) -> Vec<IpNetwork> {
let mut expired = Vec::new();
self.routes.retain(|r| {
let live = now.duration_since(r.last_seen) <= ttl;
if !live {
expired.push(r.net);
}
live
});
expired
}
pub fn len(&self) -> usize {
self.routes.len()
}
pub fn is_empty(&self) -> bool {
self.routes.is_empty()
}
}
pub struct RouteInstaller {
ifindex: u32,
tun_ip: Ipv4Addr,
server_ip: IpAddr,
installed: Mutex<HashSet<IpNetwork>>,
}
impl RouteInstaller {
pub fn new(tun_name: &str, tun_ip: Ipv4Addr, server_ip: IpAddr) -> std::io::Result<Self> {
Ok(Self {
ifindex: crate::policy::route::interface_index(tun_name)?,
tun_ip,
server_ip,
installed: Mutex::new(HashSet::new()),
})
}
pub fn apply(&self, pushed: &[IpNetwork]) {
let mut installed = self.installed.lock().unwrap();
let (add, remove) = diff_routes(&installed, pushed);
for net in add {
if net.contains(self.server_ip) {
warn!(
"refusing pushed route {net}: it covers the VPN server {} \
(would loop the tunnel into itself)",
self.server_ip
);
continue;
}
match crate::policy::route::modify_route(
self.ifindex,
self.tun_ip,
net.network(),
net.prefix(),
true,
) {
Ok(()) => {
info!("installed subnet route {net} via the tunnel");
installed.insert(net);
}
Err(e) => warn!("failed to install subnet route {net}: {e}"),
}
}
for net in remove {
match crate::policy::route::modify_route(
self.ifindex,
self.tun_ip,
net.network(),
net.prefix(),
false,
) {
Ok(()) => info!("removed withdrawn subnet route {net}"),
Err(e) => warn!("failed to remove subnet route {net}: {e}"),
}
installed.remove(&net);
}
}
pub fn remove_all(&self) {
let nets: Vec<IpNetwork> = {
let mut installed = self.installed.lock().unwrap();
installed.drain().collect()
};
for net in nets {
if let Err(e) = crate::policy::route::modify_route(
self.ifindex,
self.tun_ip,
net.network(),
net.prefix(),
false,
) {
debug!("failed to remove subnet route {net} on shutdown: {e}");
}
}
}
}
impl Drop for RouteInstaller {
fn drop(&mut self) {
self.remove_all();
}
}
pub struct InstallerGuard(std::sync::Arc<RouteInstaller>);
impl InstallerGuard {
pub fn new(installer: std::sync::Arc<RouteInstaller>) -> Self {
Self(installer)
}
}
impl Drop for InstallerGuard {
fn drop(&mut self) {
self.0.remove_all();
}
}
pub fn diff_routes(
installed: &HashSet<IpNetwork>,
pushed: &[IpNetwork],
) -> (Vec<IpNetwork>, Vec<IpNetwork>) {
let pushed: HashSet<IpNetwork> = pushed.iter().copied().map(canonical).collect();
let add = pushed.difference(installed).copied().collect();
let remove = installed.difference(&pushed).copied().collect();
(add, remove)
}
#[cfg(test)]
mod tests {
use super::*;
fn net(s: &str) -> IpNetwork {
s.parse().expect("valid test network")
}
fn peer(port: u16) -> SocketAddr {
format!("192.0.2.1:{port}").parse().unwrap()
}
#[test]
fn control_marker_never_collides_with_ip() {
assert!(!is_control(&[0x45, 0, 0, 20]));
assert!(!is_control(&[0x60, 0, 0, 0]));
assert!(is_control(&[0x00]));
assert!(!is_control(&[]));
}
#[test]
fn legacy_keepalives_still_parse() {
assert_eq!(parse_control(&[0]), Some(Control::Keepalive(None)));
assert_eq!(
parse_control(&[0, 10, 7, 0, 2]),
Some(Control::Keepalive(Some(Ipv4Addr::new(10, 7, 0, 2))))
);
}
#[test]
fn advert_roundtrips() {
let advert = RouteAdvert {
tunnel_ip: Ipv4Addr::new(10, 77, 0, 2),
tunnel_ip6: Some("fd07:7::2".parse().unwrap()),
accept_routes: true,
routes: vec![net("192.168.200.0/24"), net("fd42:cafe::/64")],
};
let bytes = advert.encode();
assert!(is_control(&bytes));
assert_eq!(parse_control(&bytes), Some(Control::RouteAdvert(advert)));
}
#[test]
fn advert_without_ip6_or_routes_roundtrips() {
let advert = RouteAdvert {
tunnel_ip: Ipv4Addr::new(10, 77, 0, 3),
tunnel_ip6: None,
accept_routes: true,
routes: vec![],
};
assert_eq!(
parse_control(&advert.encode()),
Some(Control::RouteAdvert(advert))
);
}
#[test]
fn push_roundtrips_and_canonicalizes() {
let push = RoutePush {
routes: vec![net("10.1.2.3/16"), net("fd42:cafe::1/64")],
};
let parsed = parse_control(&push.encode());
assert_eq!(
parsed,
Some(Control::RoutePush(RoutePush {
routes: vec![net("10.1.0.0/16"), net("fd42:cafe::/64")],
}))
);
}
#[test]
fn malformed_messages_are_rejected() {
assert_eq!(parse_control(&[0, 0x7f, 0, 0]), None);
assert_eq!(parse_control(&[0, 1, 0]), None);
let mut advert = RouteAdvert {
tunnel_ip: Ipv4Addr::UNSPECIFIED,
tunnel_ip6: None,
accept_routes: false,
routes: vec![],
}
.encode();
let count_at = advert.len() - 1;
advert[count_at] = 3;
assert_eq!(parse_control(&advert), None);
let mut push = RoutePush {
routes: vec![net("10.0.0.0/8")],
}
.encode();
push.push(0xaa);
assert_eq!(parse_control(&push), None);
assert_eq!(parse_control(&[0, 2, 0, 1, 5, 24, 1, 2, 3, 4]), None);
assert_eq!(parse_control(&[0, 2, 0, 1, 4, 33, 1, 2, 3, 4]), None);
}
#[test]
fn approval_allowlist_is_supernet_scoped() {
let policy = RouteApproval {
auto: false,
allowlist: vec![net("192.168.0.0/16"), net("fd42::/16")],
};
assert!(policy.approves(&net("192.168.200.0/24")));
assert!(policy.approves(&net("192.168.0.0/16")));
assert!(!policy.approves(&net("10.0.0.0/8")));
assert!(policy.approves(&net("fd42:cafe::/64")));
assert!(!policy.approves(&net("fd07::/64")));
assert!(!policy.approves(&net("172.16.0.0/12")));
let auto = RouteApproval {
auto: true,
allowlist: vec![],
};
assert!(auto.approves(&net("10.0.0.0/8")));
}
#[test]
fn table_advertise_lookup_and_split_horizon() {
let mut table = SubnetTable::default();
let approval = RouteApproval {
auto: false,
allowlist: vec![net("192.168.200.0/24"), net("fd42:cafe::/64")],
};
let now = Instant::now();
let outcome = table.advertise(
peer(1),
&[
net("192.168.200.0/24"),
net("fd42:cafe::/64"),
net("10.99.0.0/16"),
],
&approval,
now,
);
assert_eq!(outcome.approved.len(), 2);
assert_eq!(outcome.awaiting, vec![net("10.99.0.0/16")]);
assert_eq!(
table.lookup("192.168.200.7".parse().unwrap()),
Some(peer(1))
);
assert_eq!(table.lookup("fd42:cafe::1".parse().unwrap()), Some(peer(1)));
assert_eq!(table.lookup("10.99.1.1".parse().unwrap()), None);
assert!(table.routes_for(peer(1)).is_empty());
let for_other = table.routes_for(peer(2));
assert_eq!(for_other.len(), 2);
assert!(!for_other.contains(&net("10.99.0.0/16")));
assert!(table
.advertise(
peer(1),
&[
net("192.168.200.0/24"),
net("fd42:cafe::/64"),
net("10.99.0.0/16")
],
&approval,
now,
)
.is_quiet());
}
#[test]
fn table_longest_prefix_wins() {
let mut table = SubnetTable::default();
let auto = RouteApproval {
auto: true,
allowlist: vec![],
};
let now = Instant::now();
table.advertise(peer(1), &[net("10.0.0.0/8")], &auto, now);
table.advertise(peer(2), &[net("10.5.0.0/16")], &auto, now);
assert_eq!(table.lookup("10.5.1.1".parse().unwrap()), Some(peer(2)));
assert_eq!(table.lookup("10.9.1.1".parse().unwrap()), Some(peer(1)));
}
#[test]
fn table_moves_withdraws_and_expires() {
let mut table = SubnetTable::default();
let auto = RouteApproval {
auto: true,
allowlist: vec![],
};
let now = Instant::now();
table.advertise(peer(1), &[net("192.168.200.0/24")], &auto, now);
let outcome = table.advertise(peer(2), &[net("192.168.200.0/24")], &auto, now);
assert_eq!(outcome.moved, vec![net("192.168.200.0/24")]);
assert_eq!(
table.lookup("192.168.200.1".parse().unwrap()),
Some(peer(2))
);
let outcome = table.advertise(peer(2), &[], &auto, now);
assert_eq!(outcome.withdrawn, vec![net("192.168.200.0/24")]);
assert!(table.is_empty());
table.advertise(peer(1), &[net("10.1.0.0/16")], &auto, now);
let expired = table.expire(Duration::from_secs(60), now + Duration::from_secs(120));
assert_eq!(expired, vec![net("10.1.0.0/16")]);
assert_eq!(table.len(), 0);
}
#[test]
fn diff_routes_computes_add_and_remove() {
let installed: HashSet<IpNetwork> = [net("192.168.200.0/24"), net("10.1.0.0/16")]
.into_iter()
.collect();
let (add, remove) = diff_routes(
&installed,
&[net("192.168.200.0/24"), net("fd42:cafe::/64")],
);
assert_eq!(add, vec![net("fd42:cafe::/64")]);
assert_eq!(remove, vec![net("10.1.0.0/16")]);
}
}