#![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;
const IP4_HEADER_LEN: usize = 20;
const IP6_HEADER_LEN: usize = 40;
const IP6_FRAG_HEADER: u8 = 44;
const IPPROTO_UNKNOWN: IpProto = IpProto::new(0);
const IPPROTO_FRAGMENT_SENTINEL: IpProto = IpProto::new(0xff);
const IP6_FRAG_HEADER_LEN: usize = 8;
const SCTP_HEADER_LEN: usize = 12;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Ipv6Fragment {
Unknown,
Later,
First {
proto: IpProto,
dst_port: u16,
},
}
fn decode6_fragment(b: &[u8]) -> Ipv6Fragment {
if b.len() < IP6_HEADER_LEN {
return Ipv6Fragment::Unknown;
}
let length = usize::from(u16::from_be_bytes([b[4], b[5]])) + IP6_HEADER_LEN;
if b.len() < length {
return Ipv6Fragment::Unknown;
}
let Some(frag) = b.get(IP6_HEADER_LEN..) else {
return Ipv6Fragment::Unknown;
};
if frag.len() < IP6_FRAG_HEADER_LEN {
return Ipv6Fragment::Unknown;
}
let next_header = frag[0];
let frag_ofs = u16::from_be_bytes([frag[2], frag[3]]) >> 3;
let sub = &frag[IP6_FRAG_HEADER_LEN..];
if frag_ofs == 0 {
return decode6_first_fragment(IpProto::new(i64::from(next_header)), sub);
}
if frag_ofs < MIN_FRAG_BLKS {
return Ipv6Fragment::Unknown;
}
Ipv6Fragment::Later
}
fn decode6_first_fragment(proto: IpProto, sub: &[u8]) -> Ipv6Fragment {
const ICMP6_HEADER_LEN: usize = 4;
const TCP_HEADER_LEN: usize = 20;
const UDP_HEADER_LEN: usize = 8;
const MIN_TSMP_SIZE: usize = 7;
let ported = |min_len: usize| {
if sub.len() < min_len {
return Ipv6Fragment::Unknown;
}
Ipv6Fragment::First {
proto,
dst_port: u16::from_be_bytes([sub[2], sub[3]]),
}
};
let portless = |min_len: usize| {
if sub.len() < min_len {
return Ipv6Fragment::Unknown;
}
Ipv6Fragment::First { proto, dst_port: 0 }
};
match proto {
IpProto::ICMPV6 => portless(ICMP6_HEADER_LEN),
IpProto::TCP => ported(TCP_HEADER_LEN),
IpProto::UDP => ported(UDP_HEADER_LEN),
IpProto::SCTP => ported(SCTP_HEADER_LEN),
IpProto::TSMP => portless(MIN_TSMP_SIZE),
IPPROTO_FRAGMENT_SENTINEL => Ipv6Fragment::Unknown,
_ => Ipv6Fragment::First { proto, dst_port: 0 },
}
}
fn sctp_dst_port(sub: &[u8]) -> Option<u16> {
if sub.len() < SCTP_HEADER_LEN {
return None;
}
Some(u16::from_be_bytes([sub[2], sub[3]]))
}
fn fragment_header_is_chained(ipv6: ðerparse::Ipv6Slice<'_>) -> bool {
ipv6.extensions()
.clone()
.into_iter()
.any(|ext| matches!(ext, etherparse::Ipv6ExtensionSlice::Fragment(_)))
}
#[derive(Debug, Clone, Copy)]
enum Fragment {
V4(Ipv4Fragment),
V6(Ipv6Fragment),
}
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<Fragment>,
) -> bool {
if drop_before_rules(dst) {
tracing::trace!(?dst, "dropping multicast/link-local dst (pre-rule)");
return false;
}
match frag {
Some(Fragment::V4(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;
}
}
Some(Fragment::V6(Ipv6Fragment::Unknown)) => {
tracing::trace!(
?dst,
"dropping IPv6 fragment classified unknown (Go pre() drop)"
);
return false;
}
Some(Fragment::V6(Ipv6Fragment::Later)) => {
tracing::trace!(
?dst,
"accepting later IPv6 fragment (Go pre() pass-through)"
);
return true;
}
Some(Fragment::V6(Ipv6Fragment::First { .. })) | None => {}
}
if proto == IPPROTO_UNKNOWN {
tracing::trace!(?dst, "dropping unknown-proto packet (Go pre() drop)");
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 sub = match &pkt.net {
Some(etherparse::NetSlice::Ipv4(ipv4)) => ipv4.payload().payload,
Some(etherparse::NetSlice::Ipv6(ipv6)) => ipv6.payload().payload,
_ => &[][..],
};
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(Fragment::V4(Ipv4Fragment {
offset_blocks: hdr.fragments_offset().value(),
more_fragments: hdr.more_fragments(),
})),
)
}
Some(etherparse::NetSlice::Ipv6(ipv6)) => {
let hdr = ipv6.header();
let base_proto = match IpProto::new(i64::from(hdr.next_header().0)) {
IPPROTO_FRAGMENT_SENTINEL => IPPROTO_UNKNOWN,
other => other,
};
let frag = if hdr.next_header().0 == IP6_FRAG_HEADER {
Some(decode6_fragment(bytes))
} else if fragment_header_is_chained(&ipv6) {
Some(Ipv6Fragment::Unknown)
} else {
None
};
let proto = match frag {
Some(Ipv6Fragment::First { proto, .. }) => proto,
Some(Ipv6Fragment::Later | Ipv6Fragment::Unknown) => IPPROTO_UNKNOWN,
None => base_proto,
};
(
proto,
hdr.source_addr().into(),
hdr.destination_addr().into(),
frag.map(Fragment::V6),
)
}
_ => {
tracing::trace!("parsed packet is neither IPv4 nor IPv6; dropping");
return false;
}
};
let dst_port = match frag {
Some(Fragment::V6(Ipv6Fragment::First { dst_port, .. })) => dst_port,
_ if !proto.is_port_ful() => 0,
Some(Fragment::V4(v4)) if v4.offset_blocks > 0 => 0,
_ if proto == IpProto::SCTP => {
let Some(port) = sctp_dst_port(sub) else {
tracing::trace!(?dst, "dropping SCTP packet shorter than its own header");
return false;
};
port
}
_ => match pkt.transport {
Some(etherparse::TransportSlice::Udp(udp)) => udp.destination_port(),
Some(etherparse::TransportSlice::Tcp(tcp)) => tcp.destination_port(),
_ => 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())
}
fn metric_out_to_wg_drop_tsmp() -> &'static ts_metrics::Metric {
static M: std::sync::OnceLock<&'static ts_metrics::Metric> = std::sync::OnceLock::new();
M.get_or_init(|| ts_metrics::Metric::new_counter("tstun_out_to_wg_drop_tsmp"))
}
fn outbound_packet_carries_tsmp(b: &[u8]) -> bool {
match b.first().map(|first| first >> 4) {
Some(4) => b.len() >= IP4_HEADER_LEN && b[9] == ts_packet::tsmp::IP_PROTO_TSMP,
Some(6) => {
if b.len() < IP6_HEADER_LEN {
return false;
}
match b[6] {
ts_packet::tsmp::IP_PROTO_TSMP => true,
IP6_FRAG_HEADER => b
.get(IP6_HEADER_LEN)
.is_some_and(|next| *next == ts_packet::tsmp::IP_PROTO_TSMP),
_ => false,
}
}
_ => false,
}
}
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, mut packets: Vec<PacketMut>) -> OutboundResult {
if let Some(hook) = &self.capture {
for p in &packets {
hook(CapturePath::FromLocal, p.as_ref());
}
}
packets.retain(|p| {
if outbound_packet_carries_tsmp(p.as_ref()) {
tracing::debug!("[unexpected] TSMP packet written into the tun; dropping");
metric_out_to_wg_drop_tsmp().inc();
return false;
}
true
});
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 {
self.process_inbound_from(None, packets)
}
pub fn process_inbound_from(
&mut self,
from: Option<PeerId>,
packets: impl IntoIterator<Item = PacketMut>,
) -> InboundResult {
let ts_tunnel::RecvResult {
to_local,
to_peers,
sessions_established,
} = self
.wireguard
.recv_from(from.map(|p| ts_tunnel::PeerId(p.0)), 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(Fragment::V4(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)"
);
}
struct AllowAll;
impl ts_packetfilter::Filter for AllowAll {
fn match_for(
&self,
_info: &ts_packetfilter::PacketInfo,
_caps: ts_packetfilter::filter::CapIter,
) -> Option<&str> {
Some("allow-all")
}
}
struct AllowPort(u16);
impl ts_packetfilter::Filter for AllowPort {
fn match_for(
&self,
info: &ts_packetfilter::PacketInfo,
_caps: ts_packetfilter::filter::CapIter,
) -> Option<&str> {
(info.port == self.0).then_some("allow-port")
}
}
const IPV6_FIXTURE_SRC: std::net::Ipv6Addr =
std::net::Ipv6Addr::new(0x2001, 0xdb8, 0, 0, 0, 0, 0, 5);
const IPV6_FIXTURE_DST: std::net::Ipv6Addr =
std::net::Ipv6Addr::new(0x2001, 0xdb8, 0, 0, 0, 0, 0, 1);
fn ipv6_fragment_packet(
next_header: u8,
offset_blocks: u16,
more_fragments: bool,
rest: &[u8],
) -> Vec<u8> {
let mut buf = vec![0u8; IP6_HEADER_LEN + IP6_FRAG_HEADER_LEN + rest.len()];
buf[0] = 0x60; let payload_len = u16::try_from(IP6_FRAG_HEADER_LEN + rest.len()).unwrap();
buf[4..6].copy_from_slice(&payload_len.to_be_bytes());
buf[6] = IP6_FRAG_HEADER;
buf[7] = 64; buf[8..24].copy_from_slice(&IPV6_FIXTURE_SRC.octets());
buf[24..40].copy_from_slice(&IPV6_FIXTURE_DST.octets());
buf[40] = next_header;
let offset_field = (offset_blocks << 3) | u16::from(more_fragments);
buf[42..44].copy_from_slice(&offset_field.to_be_bytes());
buf[44..48].copy_from_slice(&[0xde, 0xad, 0xbe, 0xef]);
buf[48..].copy_from_slice(rest);
buf
}
fn ipv6_packet(next_header: u8, payload: &[u8]) -> Vec<u8> {
let mut buf = vec![0u8; IP6_HEADER_LEN + payload.len()];
buf[0] = 0x60; buf[4..6].copy_from_slice(&u16::try_from(payload.len()).unwrap().to_be_bytes());
buf[6] = next_header;
buf[7] = 64; buf[8..24].copy_from_slice(&IPV6_FIXTURE_SRC.octets());
buf[24..40].copy_from_slice(&IPV6_FIXTURE_DST.octets());
buf[IP6_HEADER_LEN..].copy_from_slice(payload);
buf
}
fn ipv6_udp_packet(udp: &[u8]) -> Vec<u8> {
let mut buf = ipv6_packet(17, udp);
let udp_len = u16::try_from(udp.len()).unwrap();
buf[IP6_HEADER_LEN + 4..IP6_HEADER_LEN + 6].copy_from_slice(&udp_len.to_be_bytes());
buf
}
fn ipv6_with_prepended_ext_header(ext_proto: u8, inner: &[u8]) -> Vec<u8> {
let mut buf = Vec::with_capacity(inner.len() + 8);
buf.extend_from_slice(&inner[..IP6_HEADER_LEN]);
let displaced = buf[6];
buf[6] = ext_proto;
let payload_len = u16::try_from(inner.len() - IP6_HEADER_LEN + 8).unwrap();
buf[4..6].copy_from_slice(&payload_len.to_be_bytes());
buf.extend_from_slice(&[displaced, 0, 1, 0, 0, 0, 0, 0]);
buf.extend_from_slice(&inner[IP6_HEADER_LEN..]);
buf
}
fn udp_header(dst_port: u16) -> Vec<u8> {
let mut hdr = vec![0u8; 8];
hdr[0..2].copy_from_slice(&54276u16.to_be_bytes());
hdr[2..4].copy_from_slice(&dst_port.to_be_bytes());
hdr[4..6].copy_from_slice(&16u16.to_be_bytes());
hdr
}
#[test]
fn ipv6_fragment_classification_matches_go_decode6() {
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(17, 0, true, &udp_header(443))),
Ipv6Fragment::First {
proto: IpProto::UDP,
dst_port: 443,
},
"a first fragment is decoded past the Fragment header, ports and all"
);
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(17, 185, false, &[0x61; 8])),
Ipv6Fragment::Later,
"a later fragment at a safe offset classifies as a pass-through fragment"
);
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(17, MIN_FRAG_BLKS, false, &[0x61; 8])),
Ipv6Fragment::Later,
"offset == MIN_FRAG_BLKS is the first accepted later fragment"
);
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(17, 1, false, &[0x61; 8])),
Ipv6Fragment::Unknown,
"a later fragment at offset 1 block is rejected (RFC 1858)"
);
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(
17,
MIN_FRAG_BLKS - 1,
false,
&[0x61; 8]
)),
Ipv6Fragment::Unknown,
"one block below the floor is still rejected (RFC 1858)"
);
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(17, 0, true, &udp_header(443)[..4])),
Ipv6Fragment::Unknown,
"a first fragment with only half a UDP header is rejected"
);
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(6, 0, true, &[0u8; 19])),
Ipv6Fragment::Unknown,
"a first fragment one byte short of a TCP header is rejected"
);
let mut tcp = vec![0u8; 20];
tcp[2..4].copy_from_slice(&443u16.to_be_bytes());
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(6, 0, true, &tcp)),
Ipv6Fragment::First {
proto: IpProto::TCP,
dst_port: 443,
},
"a complete TCP header in the first fragment is read normally"
);
let mut short = ipv6_fragment_packet(17, 0, true, &[]);
short.truncate(IP6_HEADER_LEN + 4);
short[4..6].copy_from_slice(&4u16.to_be_bytes());
assert_eq!(
decode6_fragment(&short),
Ipv6Fragment::Unknown,
"a truncated Fragment extension header is rejected"
);
let mut cut = ipv6_fragment_packet(17, 0, true, &udp_header(443));
cut.truncate(cut.len() - 1);
assert_eq!(
decode6_fragment(&cut),
Ipv6Fragment::Unknown,
"a packet cut off before its declared IPv6 length is rejected"
);
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(58, 0, true, &[0u8; 4])),
Ipv6Fragment::First {
proto: IpProto::ICMPV6,
dst_port: 0,
},
"a first ICMPv6 fragment keeps port 0 and is matched IPs-only"
);
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(58, 0, true, &[0u8; 3])),
Ipv6Fragment::Unknown,
"a first ICMPv6 fragment shorter than the ICMPv6 header is rejected"
);
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(0xff, 0, true, &[0u8; 8])),
Ipv6Fragment::Unknown,
"Go's internal Fragment sentinel seen on the wire maps back to unknown"
);
}
#[test]
fn ipv6_fragment_verdict_matches_go_pre() {
let src = std::net::IpAddr::V6(IPV6_FIXTURE_SRC);
let dst = std::net::IpAddr::V6(IPV6_FIXTURE_DST);
let v6 = |class| Some(Fragment::V6(class));
assert!(
!inbound_filter_verdict(
&AllowAll,
IpProto::new(0),
src,
dst,
0,
v6(Ipv6Fragment::Unknown)
),
"an unknown IPv6 fragment is dropped even under an allow-all ACL"
);
assert!(
inbound_filter_verdict(&AllowAll, IpProto::UDP, src, dst, 443, None),
"the allow-all ACL does admit an ordinary packet"
);
assert!(
inbound_filter_verdict(
&DenyAll,
IpProto::new(0),
src,
dst,
0,
v6(Ipv6Fragment::Later)
),
"a later IPv6 fragment is accepted ahead of a deny-all ACL"
);
let first = |dst_port| {
v6(Ipv6Fragment::First {
proto: IpProto::UDP,
dst_port,
})
};
assert!(
inbound_filter_verdict(&AllowPort(443), IpProto::UDP, src, dst, 443, first(443)),
"a first IPv6 fragment is matched on the port behind the Fragment header"
);
assert!(
!inbound_filter_verdict(&AllowPort(443), IpProto::UDP, src, dst, 444, first(444)),
"a first IPv6 fragment on a disallowed port is dropped by the ACL"
);
assert!(
inbound_filter_verdict(&AllowPort(443), IpProto::UDP, src, dst, 443, None),
"control: the port-scoped ACL admits an unfragmented packet to 443"
);
assert!(
!inbound_filter_verdict(&AllowPort(443), IpProto::UDP, src, dst, 0, None),
"control: port 0 - what a v6 fragment used to read as - is not admitted"
);
assert!(
!inbound_filter_verdict(
&AllowAll,
IpProto::new(0),
src,
"ff02::1".parse().unwrap(),
0,
v6(Ipv6Fragment::Later)
),
"a later fragment to a multicast dst is still dropped by pre()"
);
assert!(
!inbound_filter_verdict(
&AllowAll,
IpProto::new(0),
src,
"fe80::1".parse().unwrap(),
0,
v6(Ipv6Fragment::Later)
),
"a later fragment to a link-local dst is still dropped by pre()"
);
}
#[test]
fn ipv6_fragments_are_filtered_end_to_end() {
let keep = |filter: &(dyn ts_packetfilter::Filter + Send + Sync), packet: Vec<u8>| {
let mut packets = vec![PacketMut::from(packet)];
let mut learned = Vec::new();
filter_inbound_from_peer(filter, PeerId(3), &mut packets, &mut learned);
assert!(
learned.is_empty(),
"no TSMP advertisement in these fixtures"
);
!packets.is_empty()
};
assert!(
!keep(&AllowAll, ipv6_fragment_packet(17, 1, false, &[0x61; 8])),
"a low-offset later IPv6 fragment is dropped even by an allow-all ACL (RFC 1858)"
);
assert!(
!keep(
&AllowAll,
ipv6_fragment_packet(17, 0, true, &udp_header(443)[..4])
),
"a first IPv6 fragment too short to hold its UDP header is dropped by an allow-all ACL"
);
assert!(
keep(&AllowAll, ipv6_fragment_packet(17, 185, false, &[0x61; 8])),
"a legitimate later IPv6 fragment is delivered"
);
assert!(
keep(&DenyAll, ipv6_fragment_packet(17, 185, false, &[0x61; 8])),
"a legitimate later IPv6 fragment slides through a deny-all ACL (Go pre())"
);
assert!(
!keep(&DenyAll, ipv6_fragment_packet(17, 1, false, &[0x61; 8])),
"a low-offset later IPv6 fragment is dropped under a deny-all ACL too"
);
assert!(
keep(
&AllowPort(443),
ipv6_fragment_packet(17, 0, true, &udp_header(443))
),
"a first IPv6 fragment to an allowed port is delivered"
);
assert!(
!keep(
&AllowPort(443),
ipv6_fragment_packet(17, 0, true, &udp_header(444))
),
"a first IPv6 fragment to a disallowed port is dropped"
);
}
#[derive(Default)]
struct RecordingAllowAll(std::sync::atomic::AtomicBool);
impl RecordingAllowAll {
fn consulted(&self) -> bool {
self.0.load(std::sync::atomic::Ordering::Relaxed)
}
}
impl ts_packetfilter::Filter for RecordingAllowAll {
fn match_for(
&self,
_info: &ts_packetfilter::PacketInfo,
_caps: ts_packetfilter::filter::CapIter,
) -> Option<&str> {
self.0.store(true, std::sync::atomic::Ordering::Relaxed);
Some("allow-all")
}
}
#[test]
fn first_ipv6_fragment_with_unknown_next_header_is_dropped_before_the_acl() {
assert_eq!(
decode6_fragment(&ipv6_fragment_packet(0, 0, true, &udp_header(443))),
Ipv6Fragment::First {
proto: IPPROTO_UNKNOWN,
dst_port: 0,
},
"a first fragment carries its Fragment header's Next Header, 0 included"
);
let src = std::net::IpAddr::V6(IPV6_FIXTURE_SRC);
let dst = std::net::IpAddr::V6(IPV6_FIXTURE_DST);
assert!(
!inbound_filter_verdict(
&AllowAll,
IPPROTO_UNKNOWN,
src,
dst,
0,
Some(Fragment::V6(Ipv6Fragment::First {
proto: IPPROTO_UNKNOWN,
dst_port: 0,
})),
),
"a first IPv6 fragment declaring protocol 0 is dropped under an allow-all ACL"
);
let keep = |filter: &(dyn ts_packetfilter::Filter + Send + Sync), packet: Vec<u8>| {
let mut packets = vec![PacketMut::from(packet)];
let mut learned = Vec::new();
filter_inbound_from_peer(filter, PeerId(5), &mut packets, &mut learned);
assert!(
learned.is_empty(),
"no TSMP advertisement in these fixtures"
);
!packets.is_empty()
};
let acl = RecordingAllowAll::default();
assert!(
!keep(&acl, ipv6_fragment_packet(0, 0, true, &udp_header(443))),
"a crafted first IPv6 fragment naming protocol 0 is dropped by an allow-all ACL"
);
assert!(
!acl.consulted(),
"and it is dropped ahead of the rules: the ACL is never asked about it"
);
let control = RecordingAllowAll::default();
assert!(
keep(
&control,
ipv6_fragment_packet(17, 0, true, &udp_header(443))
),
"control: the same fragment naming UDP is delivered"
);
assert!(
control.consulted(),
"control: and it got there by being matched against the rules"
);
}
#[test]
fn chained_extension_header_cannot_bypass_the_ipv6_fragment_rules() {
let keep = |filter: &(dyn ts_packetfilter::Filter + Send + Sync), packet: Vec<u8>| {
let mut packets = vec![PacketMut::from(packet)];
let mut learned = Vec::new();
filter_inbound_from_peer(filter, PeerId(4), &mut packets, &mut learned);
assert!(
learned.is_empty(),
"no TSMP advertisement in these fixtures"
);
!packets.is_empty()
};
for ext in [0u8, 43, 60] {
assert!(
!keep(
&AllowAll,
ipv6_with_prepended_ext_header(
ext,
&ipv6_fragment_packet(17, 1, false, &[0x61; 8])
)
),
"a low-offset later fragment behind extension header {ext} is dropped (RFC 1858)"
);
assert!(
!keep(
&AllowAll,
ipv6_with_prepended_ext_header(
ext,
&ipv6_fragment_packet(17, 0, true, &udp_header(443)[..4])
)
),
"a short first fragment behind extension header {ext} is dropped"
);
assert!(
!keep(
&AllowAll,
ipv6_with_prepended_ext_header(
ext,
&ipv6_fragment_packet(17, 185, false, &[0x61; 8])
)
),
"a chained later fragment behind extension header {ext} gets no pass-through"
);
assert!(
!keep(
&AllowAll,
ipv6_with_prepended_ext_header(
ext,
&ipv6_fragment_packet(17, 0, true, &udp_header(443))
)
),
"a chained first fragment behind extension header {ext} is dropped"
);
let plain = ipv6_with_prepended_ext_header(ext, &ipv6_udp_packet(&udp_header(443)));
let parsed = etherparse::SlicedPacket::from_ip(&plain)
.unwrap_or_else(|e| panic!("extension header {ext} fixture must parse: {e:?}"));
assert!(
matches!(parsed.transport, Some(etherparse::TransportSlice::Udp(_))),
"extension header {ext} fixture must chain to a UDP header the parser can reach"
);
}
assert!(
keep(&AllowAll, ipv6_fragment_packet(17, 185, false, &[0x61; 8])),
"an unchained later fragment is still delivered"
);
}
#[derive(Default)]
struct Recording(Mutex<Vec<ts_packetfilter::PacketInfo>>);
impl ts_packetfilter::Filter for Recording {
fn match_for(
&self,
info: &ts_packetfilter::PacketInfo,
_caps: ts_packetfilter::filter::CapIter,
) -> Option<&str> {
self.0.lock().unwrap().push(*info);
Some("recording")
}
}
fn ipv6_acl(
protos: &[i64],
ports: std::ops::RangeInclusive<u16>,
) -> std::collections::BTreeMap<String, ts_packetfilter::Ruleset> {
acl("2001:db8::/32", protos, ports)
}
fn acl(
net: &str,
protos: &[i64],
ports: std::ops::RangeInclusive<u16>,
) -> std::collections::BTreeMap<String, ts_packetfilter::Ruleset> {
let net: ipnet::IpNet = net.parse().unwrap();
std::collections::BTreeMap::from([(
ts_packetfilter::DEFAULT_RULESET_NAME.to_string(),
vec![ts_packetfilter::Rule {
src: ts_packetfilter::SrcMatch {
pfxs: vec![net],
caps: Vec::new(),
},
protos: protos.iter().copied().map(IpProto::new).collect(),
dst: vec![ts_packetfilter::DstMatch {
ports,
ips: vec![net],
}],
}],
)])
}
#[test]
fn ipv6_extension_header_chain_is_matched_on_the_base_next_header() {
let keep = |filter: &(dyn ts_packetfilter::Filter + Send + Sync), packet: Vec<u8>| {
let mut packets = vec![PacketMut::from(packet)];
let mut learned = Vec::new();
filter_inbound_from_peer(filter, PeerId(5), &mut packets, &mut learned);
assert!(
learned.is_empty(),
"no TSMP advertisement in these fixtures"
);
!packets.is_empty()
};
let seen = |packet: Vec<u8>| {
let recording = Recording::default();
keep(&recording, packet);
let seen = recording.0.into_inner().unwrap();
assert!(seen.len() <= 1, "one packet in, at most one ACL question");
seen.into_iter().next()
};
let unchained = ipv6_udp_packet(&udp_header(443));
let info = seen(unchained.clone()).expect("an unchained UDP datagram reaches the ACL");
assert_eq!(info.ip_proto, IpProto::UDP, "unchained: protocol is UDP");
assert_eq!(
info.port, 443,
"unchained: the UDP destination port is read"
);
for ext in [43u8, 60] {
let chained = ipv6_with_prepended_ext_header(ext, &unchained);
let info = seen(chained.clone()).unwrap_or_else(|| {
panic!("a packet behind extension header {ext} reaches the ACL")
});
assert_eq!(
info.ip_proto,
IpProto::new(i64::from(ext)),
"behind extension header {ext}: the ACL sees the base Next Header, not the transport"
);
assert_eq!(
info.port, 0,
"behind extension header {ext}: no port is read past the chain"
);
let udp443 = ipv6_acl(&[i64::from(IpProto::UDP)], 443..=443);
assert!(
keep(&udp443, unchained.clone()),
"a udp:443 rule admits the unchained datagram"
);
assert!(
!keep(&udp443, chained.clone()),
"a udp:443 rule does not admit a packet behind extension header {ext}"
);
assert!(
keep(&ipv6_acl(&[i64::from(ext)], 0..=u16::MAX), chained.clone()),
"an all-ports rule naming protocol {ext} admits it IPs-only"
);
assert!(
!keep(&ipv6_acl(&[i64::from(ext)], 443..=443), chained),
"a port-scoped rule naming protocol {ext} opens nothing (matchProtoAndIPsOnlyIfAllPorts)"
);
}
let hop_by_hop = ipv6_with_prepended_ext_header(0, &unchained);
assert!(
seen(hop_by_hop.clone()).is_none(),
"a hop-by-hop-led packet never reaches the ACL"
);
assert!(
!keep(&AllowAll, hop_by_hop),
"a hop-by-hop-led packet is dropped by an allow-all ACL (Go pre() unknown-proto drop)"
);
let mut sentinel = unchained.clone();
sentinel[6] = 0xff;
assert!(
seen(sentinel.clone()).is_none(),
"a packet whose base Next Header is the 0xff sentinel never reaches the ACL"
);
assert!(
!keep(&AllowAll, sentinel),
"...and is dropped by an allow-all ACL"
);
}
const IPV4_FIXTURE_SRC: std::net::Ipv4Addr = std::net::Ipv4Addr::new(100, 64, 0, 9);
const IPV4_FIXTURE_DST: std::net::Ipv4Addr = std::net::Ipv4Addr::new(100, 64, 0, 1);
const IPV4_FIXTURE_NET: &str = "100.64.0.0/10";
fn v4_packet(proto: u8, offset_blocks: u16, more_fragments: bool, payload: &[u8]) -> Vec<u8> {
let total_len = u16::try_from(IP4_HEADER_LEN + payload.len()).unwrap();
let mut buf = vec![0u8; usize::from(total_len)];
buf[0] = 0x45; buf[2..4].copy_from_slice(&total_len.to_be_bytes());
let frag_field = (offset_blocks & 0x1fff) | if more_fragments { 0x2000 } else { 0 };
buf[6..8].copy_from_slice(&frag_field.to_be_bytes());
buf[8] = 64; buf[9] = proto;
buf[12..16].copy_from_slice(&IPV4_FIXTURE_SRC.octets());
buf[16..20].copy_from_slice(&IPV4_FIXTURE_DST.octets());
buf[IP4_HEADER_LEN..].copy_from_slice(payload);
buf
}
fn sctp_header(dst_port: u16, len: usize) -> Vec<u8> {
let mut hdr = vec![0u8; SCTP_HEADER_LEN];
hdr[0..2].copy_from_slice(&54276u16.to_be_bytes()); hdr[2..4].copy_from_slice(&dst_port.to_be_bytes());
hdr[4..8].copy_from_slice(&[0xde, 0xad, 0xbe, 0xef]); hdr.truncate(len);
hdr
}
#[test]
fn sctp_destination_port_is_read_before_the_acl() {
let keep = |filter: &(dyn ts_packetfilter::Filter + Send + Sync), packet: Vec<u8>| {
let mut packets = vec![PacketMut::from(packet)];
let mut learned = Vec::new();
filter_inbound_from_peer(filter, PeerId(11), &mut packets, &mut learned);
assert!(
learned.is_empty(),
"no TSMP advertisement in these fixtures"
);
!packets.is_empty()
};
let seen = |packet: Vec<u8>| {
let recording = Recording::default();
keep(&recording, packet);
let seen = recording.0.into_inner().unwrap();
assert!(seen.len() <= 1, "one packet in, at most one ACL question");
seen.into_iter().next()
};
let sctp = i64::from(IpProto::SCTP);
let whole = sctp_header(443, SCTP_HEADER_LEN);
let v4 = v4_packet(132, 0, false, &whole);
let v6 = ipv6_packet(132, &whole);
for (family, packet) in [("IPv4", &v4), ("IPv6", &v6)] {
let info = seen(packet.clone())
.unwrap_or_else(|| panic!("{family}: an SCTP packet reaches the ACL"));
assert_eq!(
info.ip_proto,
IpProto::SCTP,
"{family}: the protocol is SCTP"
);
assert_eq!(
info.port, 443,
"{family}: the SCTP destination port is read off the wire"
);
}
assert!(
keep(&acl(IPV4_FIXTURE_NET, &[sctp], 443..=443), v4.clone()),
"IPv4: an sctp:443 rule admits an SCTP packet to port 443"
);
assert!(
!keep(&acl(IPV4_FIXTURE_NET, &[sctp], 0..=442), v4.clone()),
"IPv4: an sctp:0-442 rule does not admit an SCTP packet to port 443"
);
assert!(
keep(&ipv6_acl(&[sctp], 443..=443), v6.clone()),
"IPv6: an sctp:443 rule admits an SCTP packet to port 443"
);
assert!(
!keep(&ipv6_acl(&[sctp], 0..=442), v6),
"IPv6: an sctp:0-442 rule does not admit an SCTP packet to port 443"
);
let info = seen(v4_packet(132, 0, true, &whole))
.expect("IPv4: a first SCTP fragment reaches the ACL");
assert_eq!(
info.port, 443,
"IPv4: a first fragment's SCTP port is read, as decode4 does"
);
assert!(
keep(
&DenyAll,
v4_packet(132, MIN_FRAG_BLKS, false, &[0x01, 0x02, 0x03, 0x04])
),
"IPv4: a valid later SCTP fragment is passed through ahead of the ACL"
);
let short = sctp_header(443, SCTP_HEADER_LEN - 1);
for (family, packet) in [
("IPv4", v4_packet(132, 0, false, &short)),
("IPv6", ipv6_packet(132, &short)),
] {
assert!(
seen(packet.clone()).is_none(),
"{family}: an SCTP header too short to hold its ports never reaches the ACL"
);
assert!(
!keep(&AllowAll, packet),
"{family}: ...and an allow-all ACL does not admit it"
);
}
}
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);
}
fn v4_udp_packet(src: std::net::IpAddr, dst: std::net::IpAddr, payload: &[u8]) -> Vec<u8> {
let (std::net::IpAddr::V4(src), std::net::IpAddr::V4(dst)) = (src, dst) else {
panic!("v4_udp_packet needs two IPv4 addresses");
};
let total_len = u16::try_from(IP4_HEADER_LEN + 8 + payload.len()).unwrap();
let mut buf = vec![0u8; usize::from(total_len)];
buf[0] = 0x45; buf[2..4].copy_from_slice(&total_len.to_be_bytes());
buf[8] = 64; buf[9] = 17; buf[12..16].copy_from_slice(&src.octets());
buf[16..20].copy_from_slice(&dst.octets());
buf[20..22].copy_from_slice(&4242u16.to_be_bytes()); buf[22..24].copy_from_slice(&4343u16.to_be_bytes()); let udp_len = u16::try_from(8 + payload.len()).unwrap();
buf[24..26].copy_from_slice(&udp_len.to_be_bytes());
buf[IP4_HEADER_LEN + 8..].copy_from_slice(payload);
buf
}
#[test]
fn outbound_tsmp_classification_matches_go_decode() {
let v4_src = std::net::IpAddr::from([100, 64, 0, 1]);
let v4_dst = std::net::IpAddr::from([100, 64, 0, 2]);
let v4 = ts_packet::tsmp::DiscoKeyAdvertisement {
src: v4_src,
dst: v4_dst,
key: SELF_DISCO_KEY,
}
.marshal()
.expect("a v4 advertisement between two v4 addresses marshals");
let v6 = ts_packet::tsmp::DiscoKeyAdvertisement {
src: std::net::IpAddr::V6(IPV6_FIXTURE_SRC),
dst: std::net::IpAddr::V6(IPV6_FIXTURE_DST),
key: SELF_DISCO_KEY,
}
.marshal()
.expect("a v6 advertisement between two v6 addresses marshals");
assert!(
outbound_packet_carries_tsmp(&v4),
"an IPv4 TSMP packet from the host is refused"
);
assert!(
outbound_packet_carries_tsmp(&v6),
"an IPv6 TSMP packet from the host is refused"
);
assert!(
!outbound_packet_carries_tsmp(&v4_udp_packet(v4_src, v4_dst, b"hello")),
"IPv4 UDP passes"
);
assert!(
!outbound_packet_carries_tsmp(&ipv6_udp_packet(&udp_header(53))),
"IPv6 UDP passes"
);
let mut fragmented = v4.clone();
fragmented[6] = 0x20; assert!(
outbound_packet_carries_tsmp(&fragmented),
"a fragmented IPv4 TSMP packet is refused too"
);
assert!(
outbound_packet_carries_tsmp(&ipv6_fragment_packet(
ts_packet::tsmp::IP_PROTO_TSMP,
0,
true,
&[b'a'; 33],
)),
"the head fragment of an IPv6 TSMP datagram is refused"
);
assert!(
outbound_packet_carries_tsmp(&ipv6_fragment_packet(
ts_packet::tsmp::IP_PROTO_TSMP,
MIN_FRAG_BLKS,
false,
&[0u8; 8],
)),
"so are its later fragments"
);
assert!(
!outbound_packet_carries_tsmp(&ipv6_fragment_packet(17, 0, true, &udp_header(53))),
"a fragmented IPv6 UDP datagram is not TSMP and still passes"
);
assert!(
!outbound_packet_carries_tsmp(&[]),
"the empty buffer passes"
);
assert!(
!outbound_packet_carries_tsmp(&[0xde, 0xad, 0xbe, 0xef]),
"a non-IP buffer passes (the router drops it for want of a destination)"
);
assert!(
!outbound_packet_carries_tsmp(&v4[..IP4_HEADER_LEN - 1]),
"an IPv4 packet cut off inside its header passes"
);
assert!(
!outbound_packet_carries_tsmp(&v6[..IP6_HEADER_LEN - 1]),
"an IPv6 packet cut off inside its header passes"
);
}
#[test]
fn host_written_tsmp_is_dropped_while_our_own_advertisement_still_goes_out() {
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 routes = ts_bart::Table::default();
routes.insert(
ipnet::IpNet::from(b_addr),
or::outbound::RouteAction::Wireguard(peer),
);
a.or_out.swap(routes);
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<_>>()
};
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);
let from_a = take(a.process_inbound(resp).to_peers);
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)],
"our own advertisement must still reach the peer: it is injected below process_outbound"
);
const FORGED_KEY: [u8; 32] = [0xff; 32];
let forged = ts_packet::tsmp::DiscoKeyAdvertisement {
src: a_addr,
dst: b_addr,
key: FORGED_KEY,
}
.marshal()
.expect("a v4 advertisement between two v4 addresses marshals");
const CARRIED: &[u8] = b"ordinary traffic in the same batch";
let control = v4_udp_packet(a_addr, b_addr, CARRIED);
let counted_before = metric_out_to_wg_drop_tsmp().value();
let out = a.process_outbound(vec![
PacketMut::from(&forged[..]),
PacketMut::from(&control[..]),
]);
let mark = recorded.lock().unwrap().len();
let inbound = b.process_inbound(take(out.to_peers));
assert!(
inbound.learned_disco_keys.is_empty(),
"the forged advertisement must never reach the peer, or it binds the forger's key for us"
);
let captured = recorded.lock().unwrap();
let delivered = captured[mark..]
.iter()
.filter(|(path, _)| *path == CapturePath::FromPeer)
.map(|(_, bytes)| bytes.as_slice())
.collect::<Vec<_>>();
assert_eq!(
delivered.len(),
1,
"exactly the one non-TSMP packet of the batch crosses the tunnel"
);
assert!(
delivered[0].starts_with(&control),
"and it is the ordinary traffic, unaltered"
);
assert_eq!(
metric_out_to_wg_drop_tsmp().value(),
counted_before + 1,
"the drop is counted in tstun_out_to_wg_drop_tsmp (Go metricPacketOutDropTSMP)"
);
}
}