#![doc = include_str!("../README.md")]
use std::{collections::HashMap, sync::Arc, time::Instant};
use ts_bart::RoutingTable;
use ts_overlay_router as or;
use ts_packet::PacketMut;
use ts_packetfilter::{FilterExt, IpProto};
use ts_time::{Handle, Scheduler};
use ts_transport::{OverlayTransportId, PeerId, UnderlayTransportId};
use ts_tunnel::{Endpoint, NodeKeyPair};
use ts_underlay_router as ur;
pub mod async_tokio;
const ALLOWED_LINK_LOCAL_V4: std::net::Ipv4Addr = std::net::Ipv4Addr::new(169, 254, 169, 254);
fn drop_before_rules(dst: std::net::IpAddr) -> bool {
if dst.is_multicast() {
return true;
}
match dst {
std::net::IpAddr::V4(v4) => v4.is_link_local() && v4 != ALLOWED_LINK_LOCAL_V4,
std::net::IpAddr::V6(v6) => (v6.segments()[0] & 0xffc0) == 0xfe80,
}
}
#[derive(Debug, Clone, Copy)]
struct Ipv4Fragment {
offset_blocks: u16,
more_fragments: bool,
}
const MIN_FRAG_BLKS: u16 = (60 + 20) / 8;
fn inbound_filter_verdict(
filter: &(dyn ts_packetfilter::Filter + Send + Sync),
proto: IpProto,
src: std::net::IpAddr,
dst: std::net::IpAddr,
dst_port: u16,
frag: Option<Ipv4Fragment>,
) -> bool {
if drop_before_rules(dst) {
tracing::trace!(?dst, "dropping multicast/link-local dst (pre-rule)");
return false;
}
if let Some(frag) = frag {
if frag.offset_blocks > 0 {
if frag.offset_blocks < MIN_FRAG_BLKS {
tracing::trace!(?dst, "dropping low-offset IPv4 fragment (RFC 1858)");
return false;
}
tracing::trace!(
?dst,
"accepting later IPv4 fragment (Go pre() pass-through)"
);
return true;
}
if proto == IpProto::TSMP && frag.more_fragments {
tracing::trace!(?dst, "dropping fragmented TSMP (Go parity)");
return false;
}
}
if proto == IpProto::TSMP {
tracing::trace!(?dst, "accepting TSMP inbound (bypasses ACL, Go parity)");
return true;
}
let info = ts_packetfilter::PacketInfo {
ip_proto: proto,
port: dst_port,
src,
dst,
};
let caps = [];
let verdict = filter.can_access(&info, caps);
tracing::trace!(?info, ?caps, verdict);
verdict
}
fn filter_inbound_from_peer(
filter: &(dyn ts_packetfilter::Filter + Send + Sync),
peer_id: PeerId,
packets: &mut Vec<PacketMut>,
learned_disco_keys: &mut Vec<(PeerId, ts_packet::tsmp::DiscoKeyAdvertisement)>,
) {
packets.retain(|packet| {
let bytes = packet.as_ref();
let Ok(pkt) = etherparse::SlicedPacket::from_ip(bytes) else {
tracing::trace!("does not look like ip packet");
return false;
};
let (proto, src, dst, frag) = match pkt.net {
Some(etherparse::NetSlice::Ipv4(ipv4)) => {
let hdr = ipv4.header();
(
IpProto::new(ipv4.payload().ip_number.0 as _),
hdr.source_addr().into(),
hdr.destination_addr().into(),
Some(Ipv4Fragment {
offset_blocks: hdr.fragments_offset().value(),
more_fragments: hdr.more_fragments(),
}),
)
}
Some(etherparse::NetSlice::Ipv6(ipv6)) => (
IpProto::new(ipv6.payload().ip_number.0 as _),
ipv6.header().source_addr().into(),
ipv6.header().destination_addr().into(),
None,
),
_ => {
tracing::trace!("parsed packet is neither IPv4 nor IPv6; dropping");
return false;
}
};
let (_src_port, dst_port) = match pkt.transport {
Some(etherparse::TransportSlice::Udp(udp)) => {
(udp.source_port(), udp.destination_port())
}
Some(etherparse::TransportSlice::Tcp(tcp)) => {
(tcp.source_port(), tcp.destination_port())
}
_ => (0, 0),
};
if proto == IpProto::TSMP
&& let Some(advert) = ts_packet::tsmp::DiscoKeyAdvertisement::parse(bytes)
{
if advert.key_is_zero() {
tracing::debug!(
?peer_id,
"TSMP disco-key advertisement carried the zero key; ignoring"
);
} else {
tracing::debug!(?peer_id, %src, "learned peer disco key over TSMP");
learned_disco_keys.push((peer_id, advert));
}
return false;
}
inbound_filter_verdict(filter, proto, src, dst, dst_port, frag)
});
}
#[derive(Debug, Clone, Default)]
pub struct DiscoAdvertisementState {
pub disco_key: [u8; ts_packet::tsmp::DISCO_KEY_LEN],
pub self_addrs: Vec<std::net::IpAddr>,
pub peers: HashMap<PeerId, AdvertisementTarget>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AdvertisementTarget {
pub node_addr: std::net::IpAddr,
pub wireguard_only: bool,
}
impl DiscoAdvertisementState {
pub fn advertisement_for(&self, peer: PeerId) -> Option<Vec<u8>> {
if self.disco_key == [0u8; ts_packet::tsmp::DISCO_KEY_LEN] {
tracing::debug!(?peer, "no disco key of our own; not advertising");
return None;
}
let target = self.peers.get(&peer)?;
if target.wireguard_only {
return None;
}
let src = self_ip_matching_family(&self.self_addrs, target.node_addr)?;
ts_packet::tsmp::DiscoKeyAdvertisement {
src,
dst: target.node_addr,
key: self.disco_key,
}
.marshal()
.inspect_err(|e| tracing::debug!(?peer, error = %e, "not advertising our disco key"))
.ok()
}
}
fn self_ip_matching_family(
addrs: &[std::net::IpAddr],
want: std::net::IpAddr,
) -> Option<std::net::IpAddr> {
addrs
.iter()
.copied()
.find(|addr| addr.is_ipv4() == want.is_ipv4())
}
pub enum Subsystem {
Wireguard,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CapturePath {
FromLocal = 0,
FromPeer = 1,
SynthesizedToLocal = 2,
SynthesizedToPeer = 3,
}
impl CapturePath {
pub fn code(self) -> u16 {
self as u16
}
}
pub type CaptureHook = std::sync::Arc<dyn Fn(CapturePath, &[u8]) + Send + Sync>;
pub struct DataPlane {
pub wireguard: Endpoint,
pub or_out: or::outbound::Router,
pub ur_out: ur::outbound::Router,
pub src_filter_in: Arc<ts_bart::Table<PeerId>>,
pub or_in: or::inbound::Router,
pub packet_filter: Arc<dyn ts_packetfilter::Filter + Send + Sync>,
pub events: Scheduler<Subsystem>,
pub wg_next: Option<Handle<Subsystem>>,
pub capture: Option<CaptureHook>,
pub disco_advertisement: Option<Arc<DiscoAdvertisementState>>,
}
impl DataPlane {
pub fn new(my_key: NodeKeyPair) -> Self {
DataPlane {
wireguard: Endpoint::new(my_key),
or_out: Default::default(),
ur_out: Default::default(),
src_filter_in: Default::default(),
or_in: Default::default(),
events: Default::default(),
packet_filter: Arc::new(ts_packetfilter::DropAllFilter),
wg_next: None,
capture: None,
disco_advertisement: None,
}
}
#[tracing::instrument(skip_all, fields(n_packets = packets.len()))]
pub fn process_outbound(&mut self, packets: Vec<PacketMut>) -> OutboundResult {
if let Some(hook) = &self.capture {
for p in &packets {
hook(CapturePath::FromLocal, p.as_ref());
}
}
let or::outbound::Result {
to_wireguard,
loopback,
} = self.or_out.route(packets);
let to_wireguard = to_wireguard
.into_iter()
.map(|(k, v)| (ts_tunnel::PeerId(k.0), v))
.collect::<Vec<_>>();
let ts_tunnel::SendResult {
to_peers: encrypted,
} = self.wireguard.send(to_wireguard);
let to_peers = self
.ur_out
.route(encrypted.into_iter().map(|(k, v)| (PeerId(k.0), v)));
if let Some(next) = self.wireguard.next_event()
&& let Some(prev) = self
.wg_next
.replace(self.events.add(next, Subsystem::Wireguard))
{
prev.cancel();
}
OutboundResult { to_peers, loopback }
}
pub fn process_inbound(
&mut self,
packets: impl IntoIterator<Item = PacketMut>,
) -> InboundResult {
let ts_tunnel::RecvResult {
to_local,
to_peers,
sessions_established,
} = self.wireguard.recv(packets);
if let Some(hook) = &self.capture {
for packets in to_local.values() {
for p in packets {
hook(CapturePath::FromPeer, p.as_ref());
}
}
}
let mut learned_disco_keys: Vec<(PeerId, ts_packet::tsmp::DiscoKeyAdvertisement)> =
Vec::new();
let to_local = to_local
.into_iter()
.map(|(peer_id, mut packets)| -> (PeerId, Vec<PacketMut>) {
let _span = tracing::trace_span!(
"src_filter_inbound",
peer_id = ?peer_id,
n_packet = packets.len(),
)
.entered();
packets.retain(|packet| {
let Some(src) = packet.get_src_addr() else {
tracing::trace!("does not look like ip packet");
return false;
};
let verdict = if let Some(allowed_peer) = self.src_filter_in.lookup(src) {
*allowed_peer == PeerId(peer_id.0)
} else {
tracing::trace!(remote_ip = %src, "unknown peer address");
false
};
tracing::trace!(?src, verdict);
verdict
});
(PeerId(peer_id.0), packets)
})
.map(|(peer_id, mut v)| {
let _span = tracing::trace_span!(
"packet_filter_inbound",
peer_id = ?peer_id,
n_packet = v.len()
)
.entered();
filter_inbound_from_peer(
self.packet_filter.as_ref(),
peer_id,
&mut v,
&mut learned_disco_keys,
);
v
});
let mut to_peers = to_peers;
if let Some(advert) = self.disco_advertisement.clone() {
let mut priority: HashMap<ts_tunnel::PeerId, Vec<PacketMut>> = HashMap::new();
for peer in sessions_established {
let Some(msg) = advert.advertisement_for(PeerId(peer.0)) else {
continue;
};
tracing::debug!(peer_id = ?peer, "advertising our disco key over TSMP");
for (peer, packets) in self.wireguard.send_priority_message(peer, &msg).to_peers {
priority.entry(peer).or_default().extend(packets);
}
}
for (peer, mut packets) in priority {
let queued = to_peers.entry(peer).or_default();
packets.append(queued);
*queued = packets;
}
}
let to_peers = to_peers
.into_iter()
.map(|(k, v)| (ts_transport::PeerId(k.0), v));
let to_local = self.or_in.route(to_local.flatten());
let to_peers = self.ur_out.route(to_peers);
if let Some(next) = self.wireguard.next_event()
&& let Some(prev) = self
.wg_next
.replace(self.events.add(next, Subsystem::Wireguard))
{
prev.cancel();
}
InboundResult {
to_local,
to_peers,
learned_disco_keys,
}
}
pub fn next_event(&self) -> Option<Instant> {
self.events.next_dispatch()
}
pub fn process_events(&mut self) -> EventResult {
let mut to_peers = HashMap::new();
let now = Instant::now();
for event in self.events.dispatch(now) {
match event {
Subsystem::Wireguard => {
let res = self.wireguard.dispatch_events(now);
to_peers.extend(
res.to_peers
.into_iter()
.map(|(id, pkts)| (ts_transport::PeerId(id.0), pkts)),
);
}
}
}
let to_peers = self.ur_out.route(to_peers);
if let Some(next) = self.wireguard.next_event()
&& let Some(prev) = self
.wg_next
.replace(self.events.add(next, Subsystem::Wireguard))
{
prev.cancel();
}
EventResult { to_peers }
}
}
pub struct OutboundResult {
pub to_peers: HashMap<(UnderlayTransportId, PeerId), Vec<PacketMut>>,
pub loopback: HashMap<OverlayTransportId, Vec<PacketMut>>,
}
pub struct InboundResult {
pub to_local: HashMap<OverlayTransportId, Vec<PacketMut>>,
pub to_peers: HashMap<(UnderlayTransportId, PeerId), Vec<PacketMut>>,
pub learned_disco_keys: Vec<(PeerId, ts_packet::tsmp::DiscoKeyAdvertisement)>,
}
#[derive(Default)]
pub struct EventResult {
pub to_peers: HashMap<(UnderlayTransportId, PeerId), Vec<PacketMut>>,
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use super::*;
type CaptureLog = Arc<Mutex<Vec<(CapturePath, Vec<u8>)>>>;
#[test]
fn capture_path_codes() {
assert_eq!(CapturePath::FromLocal.code(), 0);
assert_eq!(CapturePath::FromPeer.code(), 1);
assert_eq!(CapturePath::SynthesizedToLocal.code(), 2);
assert_eq!(CapturePath::SynthesizedToPeer.code(), 3);
}
#[test]
fn pre_rule_drop_matches_go() {
let ip = |s: &str| s.parse::<std::net::IpAddr>().unwrap();
assert!(drop_before_rules(ip("224.0.0.1")), "IPv4 multicast dropped");
assert!(
drop_before_rules(ip("239.255.255.250")),
"IPv4 multicast (SSDP) dropped"
);
assert!(
drop_before_rules(ip("169.254.1.1")),
"IPv4 link-local dropped"
);
assert!(drop_before_rules(ip("ff02::1")), "IPv6 multicast dropped");
assert!(drop_before_rules(ip("fe80::1")), "IPv6 link-local dropped");
assert!(
drop_before_rules(ip("febf:ffff::1")),
"top of fe80::/10 dropped (locks the 0xffc0/0xfe80 mask)"
);
assert!(
!drop_before_rules(ip("fec0::1")),
"just past fe80::/10 passes (locks the 0xffc0/0xfe80 mask)"
);
assert!(
!drop_before_rules(ip("::ffff:224.0.0.1")),
"4in6-mapped multicast falls through to the ACL, matching Go"
);
assert!(
!drop_before_rules(ip("::ffff:169.254.1.1")),
"4in6-mapped link-local falls through to the ACL, matching Go"
);
assert!(
!drop_before_rules(ip("100.64.0.5")),
"ordinary tailnet unicast passes"
);
assert!(
!drop_before_rules(ip("8.8.8.8")),
"ordinary public unicast passes"
);
assert!(
!drop_before_rules(ip("169.254.169.254")),
"the cloud-metadata link-local address is the Go-allowlisted exception"
);
assert!(
!drop_before_rules(ip("fd7a:115c:a1e0::1")),
"IPv6 ULA (tailnet) passes"
);
}
struct DenyAll;
impl ts_packetfilter::Filter for DenyAll {
fn match_for(
&self,
_info: &ts_packetfilter::PacketInfo,
_caps: ts_packetfilter::filter::CapIter,
) -> Option<&str> {
None
}
}
#[test]
fn tsmp_bypasses_acl_matches_go() {
let ip = |s: &str| s.parse::<std::net::IpAddr>().unwrap();
let src = ip("100.64.0.9");
let dst = ip("100.64.0.1");
let tsmp = IpProto::new(99);
assert!(
inbound_filter_verdict(&DenyAll, tsmp, src, dst, 0, None),
"TSMP admitted by bypassing the (deny-all) ACL"
);
assert!(
!inbound_filter_verdict(&DenyAll, IpProto::TCP, src, dst, 443, None),
"TCP still consults the ACL (deny-all → dropped)"
);
assert!(
!inbound_filter_verdict(&DenyAll, tsmp, src, ip("224.0.0.1"), 0, None),
"TSMP to a multicast dst is still dropped (pre() before the switch)"
);
assert!(
!inbound_filter_verdict(&DenyAll, tsmp, src, ip("169.254.1.1"), 0, None),
"TSMP to a link-local dst is still dropped (pre() before the switch)"
);
assert_eq!(IpProto::TSMP, tsmp, "IpProto::TSMP == 99");
}
#[test]
fn ipv4_fragment_handling_matches_go_decode4() {
let ip = |s: &str| s.parse::<std::net::IpAddr>().unwrap();
let src = ip("100.64.0.9");
let dst = ip("100.64.0.1");
let frag = |offset_blocks: u16, more_fragments: bool| {
Some(Ipv4Fragment {
offset_blocks,
more_fragments,
})
};
assert!(
inbound_filter_verdict(
&DenyAll,
IpProto::TCP,
src,
dst,
0,
frag(MIN_FRAG_BLKS, false)
),
"a valid later fragment (offset >= MIN_FRAG_BLKS) is accepted ahead of the ACL"
);
assert!(
inbound_filter_verdict(
&DenyAll,
IpProto::UDP,
src,
dst,
0,
frag(MIN_FRAG_BLKS + 50, true)
),
"a later fragment well past the floor (MF set) is also accepted"
);
assert!(
!inbound_filter_verdict(
&DenyAll,
IpProto::TCP,
src,
dst,
0,
frag(MIN_FRAG_BLKS - 1, false)
),
"a low-offset later fragment is dropped (RFC 1858)"
);
assert!(
!inbound_filter_verdict(&DenyAll, IpProto::TCP, src, dst, 0, frag(1, false)),
"the smallest non-zero offset is dropped"
);
assert!(
!inbound_filter_verdict(&DenyAll, IpProto::TCP, src, dst, 443, frag(0, true)),
"a first fragment defers to the ACL (deny-all -> dropped) on its parsed port"
);
assert!(
!inbound_filter_verdict(&DenyAll, IpProto::TSMP, src, dst, 0, frag(0, true)),
"a fragmented TSMP first fragment is dropped (Go parity)"
);
assert!(
inbound_filter_verdict(&DenyAll, IpProto::TSMP, src, dst, 0, frag(0, false)),
"a non-fragmented TSMP (offset 0, MF clear) still bypasses the ACL"
);
assert!(
inbound_filter_verdict(
&DenyAll,
IpProto::TSMP,
src,
dst,
0,
frag(MIN_FRAG_BLKS, true)
),
"a later TSMP fragment is accepted via the fragment path (proto-independent)"
);
}
fn tsmp_packet4(src: [u8; 4], dst: [u8; 4], body: &[u8]) -> PacketMut {
let mut buf = vec![0u8; 20 + body.len()];
buf[20..].copy_from_slice(body);
buf[0] = 0x45;
let total_len = buf.len() as u16;
buf[2..4].copy_from_slice(&total_len.to_be_bytes());
buf[8] = 64;
buf[9] = 99;
buf[12..16].copy_from_slice(&src);
buf[16..20].copy_from_slice(&dst);
PacketMut::from(buf)
}
fn advertisement_body(key: [u8; 32]) -> Vec<u8> {
let mut body = vec![ts_packet::tsmp::TSMP_TYPE_DISCO_ADVERTISEMENT];
body.extend_from_slice(&key);
body
}
#[test]
fn tsmp_disco_key_advertisement_is_learned_and_dropped() {
let peer = PeerId(7);
let src = [100, 64, 0, 2];
let dst = [100, 64, 0, 1];
let key = [0xa5u8; 32];
let mut packets = vec![tsmp_packet4(src, dst, &advertisement_body(key))];
let mut learned = Vec::new();
filter_inbound_from_peer(&DenyAll, peer, &mut packets, &mut learned);
assert!(
packets.is_empty(),
"a consumed advertisement must not be delivered to the local stack"
);
assert_eq!(learned.len(), 1, "the advertisement must be harvested");
assert_eq!(
learned[0].0, peer,
"attributed to the sending wireguard peer"
);
assert_eq!(learned[0].1.key, key, "the advertised disco key is learned");
assert_eq!(learned[0].1.src, std::net::IpAddr::from(src));
let mut ping = vec![ts_packet::tsmp::TSMP_TYPE_PING];
ping.extend_from_slice(&[1, 2, 3, 4, 5, 6, 7, 8]);
let mut packets = vec![tsmp_packet4(src, dst, &ping)];
let mut learned = Vec::new();
filter_inbound_from_peer(&DenyAll, peer, &mut packets, &mut learned);
assert_eq!(packets.len(), 1, "a TSMP ping still bypasses the ACL");
assert!(learned.is_empty(), "a ping advertises no disco key");
}
#[test]
fn malformed_tsmp_disco_key_advertisements_teach_nothing() {
let peer = PeerId(7);
let src = [100, 64, 0, 2];
let dst = [100, 64, 0, 1];
let mut truncated = advertisement_body([0xa5u8; 32]);
truncated.truncate(32);
for (name, body, still_delivered) in [
("truncated advertisement", truncated, true),
(
"unknown TSMP type byte",
{
let mut b = advertisement_body([0xa5u8; 32]);
b[0] = b'Z';
b
},
true,
),
(
"zero-key advertisement",
advertisement_body([0u8; 32]),
false,
),
] {
let mut packets = vec![tsmp_packet4(src, dst, &body)];
let mut learned = Vec::new();
filter_inbound_from_peer(&DenyAll, peer, &mut packets, &mut learned);
assert!(
learned.is_empty(),
"a {name} must not be half-parsed into a learned disco key"
);
assert_eq!(
packets.len(),
usize::from(still_delivered),
"a {name} must {} be delivered",
if still_delivered { "still" } else { "not" }
);
}
}
const SELF_DISCO_KEY: [u8; 32] = [
0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xaa, 0xbb, 0xcc, 0xdd, 0xee, 0xff,
0x00, 0x9c, 0x5f, 0x3a, 0x01, 0x7d, 0xe2, 0x44, 0xb8, 0x0f, 0x1e, 0x2d, 0x3c, 0x4b, 0x5a,
0x69, 0x78,
];
fn advertisement_state(peer: PeerId, target: AdvertisementTarget) -> DiscoAdvertisementState {
DiscoAdvertisementState {
disco_key: SELF_DISCO_KEY,
self_addrs: vec![
std::net::IpAddr::from([100, 64, 0, 1]),
std::net::IpAddr::from([
0xfd, 0x7a, 0x11, 0x5c, 0xa1, 0xe0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1,
]),
],
peers: HashMap::from([(peer, target)]),
}
}
#[test]
fn disco_advertisement_matches_priority_message_for_peer() {
let peer = PeerId(3);
let peer_v4 = std::net::IpAddr::from([100, 64, 0, 2]);
let target = AdvertisementTarget {
node_addr: peer_v4,
wireguard_only: false,
};
let state = advertisement_state(peer, target);
let msg = state
.advertisement_for(peer)
.expect("a Tailscale peer with a matching-family address must be advertised to");
let parsed = ts_packet::tsmp::DiscoKeyAdvertisement::parse(&msg)
.expect("what we emit must parse as an advertisement");
assert_eq!(parsed.key, SELF_DISCO_KEY, "we advertise OUR disco key");
assert_eq!(parsed.src, std::net::IpAddr::from([100, 64, 0, 1]));
assert_eq!(parsed.dst, peer_v4);
assert_eq!(
msg,
ts_packet::tsmp::DiscoKeyAdvertisement {
src: std::net::IpAddr::from([100, 64, 0, 1]),
dst: peer_v4,
key: SELF_DISCO_KEY,
}
.marshal()
.unwrap(),
"the emitted bytes are exactly what Marshal produces"
);
let peer_v6 = std::net::IpAddr::from([
0xfd, 0x7a, 0x11, 0x5c, 0xa1, 0xe0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2,
]);
let v6_state = advertisement_state(
peer,
AdvertisementTarget {
node_addr: peer_v6,
wireguard_only: false,
},
);
let parsed = v6_state
.advertisement_for(peer)
.and_then(|m| ts_packet::tsmp::DiscoKeyAdvertisement::parse(&m))
.expect("a v6 peer must be advertised to over v6");
assert!(parsed.src.is_ipv6(), "source must match the peer's family");
assert_eq!(parsed.dst, peer_v6);
let mut no_key = advertisement_state(peer, target);
no_key.disco_key = [0u8; 32];
assert!(
no_key.advertisement_for(peer).is_none(),
"the zero disco key must never be advertised"
);
assert!(
state.advertisement_for(PeerId(0xbad)).is_none(),
"an unknown peer must not be advertised to"
);
let mut no_self = advertisement_state(peer, target);
no_self.self_addrs.clear();
assert!(
no_self.advertisement_for(peer).is_none(),
"a node with no tailnet address of its own has no source to advertise from"
);
let wg_only = advertisement_state(
peer,
AdvertisementTarget {
node_addr: peer_v4,
wireguard_only: true,
},
);
assert!(
wg_only.advertisement_for(peer).is_none(),
"a WireGuard-only peer must never be sent TSMP"
);
let mut v4_only = advertisement_state(
peer,
AdvertisementTarget {
node_addr: peer_v6,
wireguard_only: false,
},
);
v4_only.self_addrs = vec![std::net::IpAddr::from([100, 64, 0, 1])];
assert!(
v4_only.advertisement_for(peer).is_none(),
"no self address in the peer's family means no advertisement"
);
}
#[test]
fn session_establishment_advertises_our_disco_key_to_the_peer() {
let underlay: UnderlayTransportId = 0.into();
let wg_peer = ts_tunnel::PeerId(1);
let peer = PeerId(1);
let a_addr = std::net::IpAddr::from([100, 64, 0, 1]);
let b_addr = std::net::IpAddr::from([100, 64, 0, 2]);
let (a_static, b_static) = (NodeKeyPair::new(), NodeKeyPair::new());
let (mut a, mut b) = (
DataPlane::new(a_static.clone()),
DataPlane::new(b_static.clone()),
);
for (dp, key) in [(&mut a, b_static.public), (&mut b, a_static.public)] {
dp.wireguard.upsert_peer(
wg_peer,
ts_tunnel::PeerConfig {
key,
psk: [0u8; 32].into(),
persistent_keepalive_interval: None,
},
);
dp.ur_out.table.insert(peer, underlay);
}
a.disco_advertisement = Some(Arc::new(advertisement_state(
peer,
AdvertisementTarget {
node_addr: b_addr,
wireguard_only: false,
},
)));
let mut src_filter = ts_bart::Table::default();
src_filter.insert(ipnet::IpNet::from(a_addr), peer);
b.src_filter_in = Arc::new(src_filter);
let take = |out: HashMap<(UnderlayTransportId, PeerId), Vec<PacketMut>>| {
out.into_values().flatten().collect::<Vec<_>>()
};
let init = a
.wireguard
.send([(wg_peer, vec![PacketMut::from(&b"hello"[..])])])
.to_peers
.remove(&wg_peer)
.expect("handshake initiation");
let resp = take(b.process_inbound(init).to_peers);
assert!(!resp.is_empty(), "B must answer the handshake initiation");
let from_a = take(a.process_inbound(resp).to_peers);
assert_eq!(
from_a.len(),
2,
"A must emit the queued data AND its disco-key advertisement"
);
let inbound = b.process_inbound(from_a);
assert_eq!(
inbound
.learned_disco_keys
.iter()
.map(|(peer, advert)| (*peer, advert.key))
.collect::<Vec<_>>(),
vec![(peer, SELF_DISCO_KEY)],
"B must learn exactly the disco key A holds, attributed to A's wireguard peer"
);
assert!(
inbound.to_peers.is_empty(),
"B has no advertisement state, so it advertises nothing back"
);
}
#[test]
fn the_advertisement_leads_the_traffic_released_by_the_same_establishment() {
let underlay: UnderlayTransportId = 0.into();
let wg_peer = ts_tunnel::PeerId(1);
let peer = PeerId(1);
let a_addr = std::net::IpAddr::from([100, 64, 0, 1]);
let b_addr = std::net::IpAddr::from([100, 64, 0, 2]);
let (a_static, b_static) = (NodeKeyPair::new(), NodeKeyPair::new());
let (mut a, mut b) = (
DataPlane::new(a_static.clone()),
DataPlane::new(b_static.clone()),
);
for (dp, key) in [(&mut a, b_static.public), (&mut b, a_static.public)] {
dp.wireguard.upsert_peer(
wg_peer,
ts_tunnel::PeerConfig {
key,
psk: [0u8; 32].into(),
persistent_keepalive_interval: None,
},
);
dp.ur_out.table.insert(peer, underlay);
}
a.disco_advertisement = Some(Arc::new(advertisement_state(
peer,
AdvertisementTarget {
node_addr: b_addr,
wireguard_only: false,
},
)));
let mut src_filter = ts_bart::Table::default();
src_filter.insert(ipnet::IpNet::from(a_addr), peer);
b.src_filter_in = Arc::new(src_filter);
let recorded: CaptureLog = Arc::new(Mutex::new(Vec::new()));
let sink = recorded.clone();
b.capture = Some(Arc::new(move |path: CapturePath, bytes: &[u8]| {
sink.lock().unwrap().push((path, bytes.to_vec()));
}));
let take = |out: HashMap<(UnderlayTransportId, PeerId), Vec<PacketMut>>| {
out.into_values().flatten().collect::<Vec<_>>()
};
const QUEUED: &[u8] = b"staged while the session was still coming up";
let init = a
.wireguard
.send([(wg_peer, vec![PacketMut::from(QUEUED)])])
.to_peers
.remove(&wg_peer)
.expect("handshake initiation");
let resp = take(b.process_inbound(init).to_peers);
let from_a = take(a.process_inbound(resp).to_peers);
assert_eq!(
from_a.len(),
2,
"A must emit the queued data AND its disco-key advertisement"
);
let learned = b.process_inbound(from_a).learned_disco_keys;
assert_eq!(
learned
.iter()
.map(|(peer, advert)| (*peer, advert.key))
.collect::<Vec<_>>(),
vec![(peer, SELF_DISCO_KEY)],
"B must still learn A's disco key"
);
let advertisement = ts_packet::tsmp::DiscoKeyAdvertisement {
src: a_addr,
dst: b_addr,
key: SELF_DISCO_KEY,
}
.marshal()
.expect("a v4 advertisement between two v4 addresses marshals");
let captured = recorded.lock().unwrap();
let from_peer = captured
.iter()
.filter(|(path, _)| *path == CapturePath::FromPeer)
.map(|(_, bytes)| bytes.as_slice())
.collect::<Vec<_>>();
assert_eq!(from_peer.len(), 2, "B must decrypt both of A's packets");
assert!(
from_peer[0].starts_with(&advertisement),
"the advertisement must reach the peer FIRST, ahead of the traffic the same \
establishment released"
);
assert!(
from_peer[1].starts_with(QUEUED),
"the queued traffic follows the advertisement"
);
}
#[test]
fn capture_hook_fires_on_outbound() {
let mut dp = DataPlane::new(NodeKeyPair::new());
let recorded: CaptureLog = Arc::new(Mutex::new(Vec::new()));
let sink = recorded.clone();
dp.capture = Some(Arc::new(move |path: CapturePath, bytes: &[u8]| {
sink.lock().unwrap().push((path, bytes.to_vec()));
}));
let payload: Vec<u8> = vec![0xde, 0xad, 0xbe, 0xef];
let packet = PacketMut::from(payload.clone());
drop(dp.process_outbound(vec![packet]));
let captured = recorded.lock().unwrap();
assert_eq!(captured.len(), 1, "hook must fire exactly once per packet");
assert_eq!(captured[0].0, CapturePath::FromLocal);
assert_eq!(captured[0].1, payload);
}
}