use std::io;
use std::net::{IpAddr, Ipv6Addr};
use std::os::fd::AsRawFd;
use std::sync::Mutex;
use super::insn::{
Asm, BPF_FUNC_MAP_LOOKUP_ELEM, BPF_FUNC_REDIRECT_MAP, Insn, Jmp, Label, R0, R1, R2, R3, R4, R5,
R6, R7, R8, R10, Size, host_be16, ld_map_fd,
};
use super::map::{Map, UpdateFlags, lpm_key};
use super::prog::{Action, Link, Mode, Program, TestRun};
use crate::{EtherType, IpPrefix, Protocol, Result};
const ETH_HLEN: i32 = 14;
const ETH_TYPE: i16 = 12;
const IPV4_FRAG: i16 = ETH_HLEN as i16 + 6;
const IPV4_PROTO: i16 = ETH_HLEN as i16 + 9;
const IPV4_SRC: i16 = ETH_HLEN as i16 + 12;
const IPV4_DST: i16 = ETH_HLEN as i16 + 16;
const IPV4_MIN: i32 = ETH_HLEN + 20;
const IPV4_FRAG_OFF_MASK: u16 = 0x1fff;
const IPV6_NEXT: i16 = ETH_HLEN as i16 + 6;
const IPV6_SRC: i16 = ETH_HLEN as i16 + 8;
const IPV6_DST: i16 = ETH_HLEN as i16 + 24;
const IPV6_MIN: i32 = ETH_HLEN + 40;
const ARP_PTYPE: i16 = ETH_HLEN as i16 + 2;
const ARP_SPA: i16 = ETH_HLEN as i16 + 14;
const ARP_TPA: i16 = ETH_HLEN as i16 + 24;
const ARP_MIN: i32 = ETH_HLEN + 28;
const XDP_MD_RX_QUEUE_INDEX: i16 = 16;
const V4_DST_KEY: i16 = -8;
const V4_SRC_KEY: i16 = -16;
const V6_DST_KEY: i16 = -40;
const V6_SRC_KEY: i16 = -64;
const L4_SPORT: i16 = 0;
const L4_DPORT: i16 = 2;
const RULE_PROTO: u8 = R3;
const RULE_PORT: u8 = R4;
const NO_PORT: i32 = 0x1_0000;
const _: () = {
assert!(V4_DST_KEY % 4 == 0 && V4_SRC_KEY % 4 == 0);
assert!(V6_DST_KEY % 4 == 0 && V6_SRC_KEY % 4 == 0);
assert!(V4_SRC_KEY + 8 <= V4_DST_KEY, "v4 keys overlap");
assert!(
V6_DST_KEY + 20 <= V4_SRC_KEY,
"v6 dst key overlaps a v4 key"
);
assert!(V6_SRC_KEY + 20 <= V6_DST_KEY, "v6 keys overlap");
assert!(V6_SRC_KEY > -512, "keys exceed the BPF stack");
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Rule {
Any,
Proto(Protocol),
Port(Protocol, u16),
}
const KIND_END: u8 = 0;
const KIND_ANY: u8 = 1;
const KIND_PROTO: u8 = 2;
const KIND_PORT: u8 = 3;
const RULE_SIZE: usize = 4;
pub const MAX_RULES_PER_PREFIX: u8 = 64;
impl Rule {
pub fn validate(&self) -> Result<()> {
match self {
Rule::Port(proto, _) if *proto != Protocol::TCP && *proto != Protocol::UDP => {
Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"xdp: a port rule needs TCP or UDP, got protocol {}",
proto.as_u8()
),
))
}
_ => Ok(()),
}
}
fn encode(self) -> [u8; RULE_SIZE] {
match self {
Rule::Any => [KIND_ANY, 0, 0, 0],
Rule::Proto(p) => [KIND_PROTO, p.as_u8(), 0, 0],
Rule::Port(p, port) => {
let b = port.to_be_bytes();
[KIND_PORT, p.as_u8(), b[0], b[1]]
}
}
}
fn decode(b: &[u8; RULE_SIZE]) -> Option<Rule> {
match b[0] {
KIND_ANY => Some(Rule::Any),
KIND_PROTO => Some(Rule::Proto(Protocol(b[1]))),
KIND_PORT => Some(Rule::Port(Protocol(b[1]), u16::from_be_bytes([b[2], b[3]]))),
_ => None,
}
}
}
fn encode_rules(rules: &[Rule], max_rules: u8) -> Vec<u8> {
let mut v = vec![KIND_END; value_size(max_rules) as usize];
let any = rules.iter().filter(|r| **r == Rule::Any);
let rest = rules.iter().filter(|r| **r != Rule::Any);
for (slot, rule) in any.chain(rest).take(max_rules as usize).enumerate() {
v[slot * RULE_SIZE..(slot + 1) * RULE_SIZE].copy_from_slice(&rule.encode());
}
v
}
fn decode_rules(value: &[u8]) -> Vec<Rule> {
value
.as_chunks::<RULE_SIZE>()
.0
.iter()
.map_while(Rule::decode)
.collect()
}
#[inline]
fn value_size(max_rules: u8) -> u32 {
u32::from(max_rules) * RULE_SIZE as u32
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum MatchField {
#[default]
Dst,
Src,
Either,
}
impl MatchField {
#[inline]
fn wants_dst(self) -> bool {
matches!(self, MatchField::Dst | MatchField::Either)
}
#[inline]
fn wants_src(self) -> bool {
matches!(self, MatchField::Src | MatchField::Either)
}
}
#[derive(Debug, Clone)]
pub struct CaptureConfig {
pub match_field: MatchField,
pub arp: bool,
pub neighbor_discovery: bool,
pub default_action: Action,
pub min_prefix_v4: u8,
pub min_prefix_v6: u8,
pub max_prefixes: u32,
pub max_rules_per_prefix: u8,
pub max_queues: u32,
}
impl Default for CaptureConfig {
fn default() -> CaptureConfig {
CaptureConfig {
match_field: MatchField::Dst,
arp: true,
neighbor_discovery: true,
default_action: Action::PASS,
min_prefix_v4: 1,
min_prefix_v6: 1,
max_prefixes: 1024,
max_rules_per_prefix: 8,
max_queues: 64,
}
}
}
impl CaptureConfig {
pub fn validate(&self) -> Result<()> {
if self.min_prefix_v4 == 0 || self.min_prefix_v6 == 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"xdp: min_prefix_v4/min_prefix_v6 must be at least 1; a /0 \
matches every packet on the interface",
));
}
if self.min_prefix_v4 > 32 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("xdp: min_prefix_v4 is /{}, max is /32", self.min_prefix_v4),
));
}
if self.min_prefix_v6 > 128 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("xdp: min_prefix_v6 is /{}, max is /128", self.min_prefix_v6),
));
}
if self.default_action != Action::PASS && self.default_action != Action::DROP {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"xdp: default_action must be PASS or DROP, got {:?}",
self.default_action
),
));
}
if self.max_prefixes == 0 || self.max_queues == 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"xdp: max_prefixes and max_queues must be non-zero",
));
}
if self.max_rules_per_prefix == 0 || self.max_rules_per_prefix > MAX_RULES_PER_PREFIX {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"xdp: max_rules_per_prefix is {}, must be 1-{MAX_RULES_PER_PREFIX}",
self.max_rules_per_prefix
),
));
}
Ok(())
}
pub fn check_prefix(&self, prefix: IpPrefix) -> Result<()> {
let (min, family) = if prefix.is_v4() {
(self.min_prefix_v4.max(1), "IPv4")
} else {
(self.min_prefix_v6.max(1), "IPv6")
};
if prefix.bits() >= min {
return Ok(());
}
let why = if prefix.bits() == 0 {
" — a /0 matches every packet on the interface".to_string()
} else {
format!(" — the {family} floor is /{min}")
};
Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("xdp: refusing to capture {prefix}{why}"),
))
}
}
fn coverage(prefixes: &[IpPrefix], v4: bool) -> u128 {
let width: u32 = if v4 { 32 } else { 128 };
prefixes
.iter()
.filter(|p| p.is_v4() == v4)
.fold(0u128, |acc, p| {
let host_bits = width - u32::from(p.bits()).min(width);
let n = 1u128.checked_shl(host_bits).unwrap_or(u128::MAX);
acc.saturating_add(n)
})
}
fn family_total(v4: bool) -> u128 {
if v4 { 1u128 << 32 } else { u128::MAX }
}
fn check_coverage(held: &[IpPrefix], new: IpPrefix) -> Result<()> {
let v4 = new.is_v4();
let mut combined = held.to_vec();
combined.push(new);
if coverage(&combined, v4) >= family_total(v4) {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"xdp: refusing to capture {new}: it would leave the capture set \
covering every {} address on the interface",
if v4 { "IPv4" } else { "IPv6" }
),
));
}
Ok(())
}
#[derive(Debug)]
pub struct CaptureMaps {
pub xskmap: Map,
pub v4: Map,
pub v6: Map,
}
impl CaptureMaps {
pub fn create(cfg: &CaptureConfig) -> Result<CaptureMaps> {
let value = value_size(cfg.max_rules_per_prefix);
Ok(CaptureMaps {
xskmap: Map::xskmap(cfg.max_queues)?,
v4: Map::lpm_trie(4, value, cfg.max_prefixes)?,
v6: Map::lpm_trie(16, value, cfg.max_prefixes)?,
})
}
}
fn stage_v4(asm: &mut Asm, slot: i16, pkt_off: i16) {
asm.emit(Insn::mov64_imm(R1, 32));
asm.emit(Insn::stx(Size::W, R10, slot, R1));
asm.emit(Insn::ldx(Size::W, R1, R7, pkt_off));
asm.emit(Insn::stx(Size::W, R10, slot + 4, R1));
}
fn stage_v6(asm: &mut Asm, slot: i16, pkt_off: i16) {
asm.emit(Insn::mov64_imm(R1, 128));
asm.emit(Insn::stx(Size::W, R10, slot, R1));
for w in 0..4i16 {
asm.emit(Insn::ldx(Size::W, R1, R7, pkt_off + w * 4));
asm.emit(Insn::stx(Size::W, R10, slot + 4 + w * 4, R1));
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Family {
V4,
V6,
}
fn load_l4(asm: &mut Asm, family: Family, port_off: i16) {
let l_done = asm.label();
asm.emit(Insn::mov64_imm(RULE_PORT, NO_PORT));
match family {
Family::V4 => {
asm.emit(Insn::ldx(Size::B, RULE_PROTO, R7, IPV4_PROTO));
asm.emit(Insn::ldx(Size::H, R1, R7, IPV4_FRAG));
asm.jump(
Insn::jmp_imm(Jmp::JSET, R1, host_be16(IPV4_FRAG_OFF_MASK), 0),
l_done,
);
asm.emit(Insn::ldx(Size::B, R2, R7, ETH_HLEN as i16));
asm.emit(Insn::and64_imm(R2, 0x0f));
asm.emit(Insn::lsh64_imm(R2, 2));
asm.jump(Insn::jmp_imm(Jmp::JLT, R2, IPV4_MIN - ETH_HLEN, 0), l_done);
asm.emit(Insn::mov64_reg(R1, R7));
asm.emit(Insn::add64_imm(R1, ETH_HLEN));
asm.emit(Insn::add64_reg(R1, R2));
}
Family::V6 => {
asm.emit(Insn::ldx(Size::B, RULE_PROTO, R7, IPV6_NEXT));
asm.emit(Insn::mov64_reg(R1, R7));
asm.emit(Insn::add64_imm(R1, IPV6_MIN));
}
}
asm.emit(Insn::mov64_reg(R5, R1));
asm.emit(Insn::add64_imm(R5, 4));
asm.jump(Insn::jmp_reg(Jmp::JGT, R5, R8, 0), l_done);
asm.emit(Insn::ldx(Size::H, RULE_PORT, R1, port_off));
asm.place(l_done);
}
fn lookup(asm: &mut Asm, map_fd: i32, slot: i16, miss: Label) {
asm.emit_all(&ld_map_fd(R1, map_fd));
asm.emit(Insn::mov64_reg(R2, R10));
asm.emit(Insn::add64_imm(R2, slot as i32));
asm.emit(Insn::call(BPF_FUNC_MAP_LOOKUP_ELEM));
asm.jump(Insn::jmp_imm(Jmp::JEQ, R0, 0, 0), miss);
}
fn match_rules(asm: &mut Asm, max_rules: u8, l4: Option<(Family, i16)>, hit: Label, miss: Label) {
for i in 0..max_rules as i16 {
let next = asm.label();
let off = i * RULE_SIZE as i16;
asm.emit(Insn::ldx(Size::B, R1, R0, off));
asm.jump(Insn::jmp_imm(Jmp::JEQ, R1, KIND_END as i32, 0), miss);
asm.jump(Insn::jmp_imm(Jmp::JEQ, R1, KIND_ANY as i32, 0), hit);
if let Some((family, port_off)) = l4 {
if i == 0 {
load_l4(asm, family, port_off);
asm.emit(Insn::ldx(Size::B, R1, R0, off));
}
asm.emit(Insn::ldx(Size::B, R2, R0, off + 1));
asm.jump(Insn::jmp_reg(Jmp::JNE, R2, RULE_PROTO, 0), next);
asm.jump(Insn::jmp_imm(Jmp::JEQ, R1, KIND_PROTO as i32, 0), hit);
asm.emit(Insn::ldx(Size::H, R2, R0, off + 2));
asm.jump(Insn::jmp_reg(Jmp::JEQ, R2, RULE_PORT, 0), hit);
}
asm.place(next);
}
asm.jump(Insn::ja(0), miss);
}
fn lookup_and_match(
asm: &mut Asm,
cfg: &CaptureConfig,
map_fd: i32,
slot: i16,
l4: Option<(Family, i16)>,
hit: Label,
) {
let miss = asm.label();
lookup(asm, map_fd, slot, miss);
match_rules(asm, cfg.max_rules_per_prefix, l4, hit, miss);
asm.place(miss);
}
fn need_bytes(asm: &mut Asm, n: i32, miss: Label) {
asm.emit(Insn::mov64_reg(R1, R7));
asm.emit(Insn::add64_imm(R1, n));
asm.jump(Insn::jmp_reg(Jmp::JGT, R1, R8, 0), miss);
}
pub fn build_program(cfg: &CaptureConfig, maps: &CaptureMaps) -> Result<Vec<Insn>> {
build_program_with_fds(
cfg,
maps.xskmap.as_raw_fd(),
maps.v4.as_raw_fd(),
maps.v6.as_raw_fd(),
)
}
fn build_program_with_fds(
cfg: &CaptureConfig,
xskmap_fd: i32,
v4_fd: i32,
v6_fd: i32,
) -> Result<Vec<Insn>> {
let mut asm = Asm::new();
let l_v4 = asm.label();
let l_v6 = asm.label();
let l_arp = asm.label();
let l_redirect = asm.label();
let l_default = asm.label();
asm.emit(Insn::mov64_reg(R6, R1));
asm.emit(Insn::ldx(Size::W, R7, R6, 0));
asm.emit(Insn::ldx(Size::W, R8, R6, 4));
need_bytes(&mut asm, ETH_HLEN, l_default);
asm.emit(Insn::ldx(Size::H, R2, R7, ETH_TYPE));
asm.jump(
Insn::jmp_imm(Jmp::JEQ, R2, host_be16(EtherType::IPV4.0), 0),
l_v4,
);
asm.jump(
Insn::jmp_imm(Jmp::JEQ, R2, host_be16(EtherType::IPV6.0), 0),
l_v6,
);
if cfg.arp {
asm.jump(
Insn::jmp_imm(Jmp::JEQ, R2, host_be16(EtherType::ARP.0), 0),
l_arp,
);
}
asm.jump(Insn::ja(0), l_default);
asm.place(l_v4);
need_bytes(&mut asm, IPV4_MIN, l_default);
if cfg.match_field.wants_dst() {
stage_v4(&mut asm, V4_DST_KEY, IPV4_DST);
}
if cfg.match_field.wants_src() {
stage_v4(&mut asm, V4_SRC_KEY, IPV4_SRC);
}
if cfg.match_field.wants_dst() {
let l4 = Some((Family::V4, L4_DPORT));
lookup_and_match(&mut asm, cfg, v4_fd, V4_DST_KEY, l4, l_redirect);
}
if cfg.match_field.wants_src() {
let l4 = Some((Family::V4, L4_SPORT));
lookup_and_match(&mut asm, cfg, v4_fd, V4_SRC_KEY, l4, l_redirect);
}
asm.jump(Insn::ja(0), l_default);
asm.place(l_v6);
need_bytes(&mut asm, IPV6_MIN, l_default);
if cfg.match_field.wants_dst() {
stage_v6(&mut asm, V6_DST_KEY, IPV6_DST);
}
if cfg.match_field.wants_src() {
stage_v6(&mut asm, V6_SRC_KEY, IPV6_SRC);
}
if cfg.match_field.wants_dst() {
let l4 = Some((Family::V6, L4_DPORT));
lookup_and_match(&mut asm, cfg, v6_fd, V6_DST_KEY, l4, l_redirect);
}
if cfg.match_field.wants_src() {
let l4 = Some((Family::V6, L4_SPORT));
lookup_and_match(&mut asm, cfg, v6_fd, V6_SRC_KEY, l4, l_redirect);
}
asm.jump(Insn::ja(0), l_default);
if cfg.arp {
asm.place(l_arp);
need_bytes(&mut asm, ARP_MIN, l_default);
asm.emit(Insn::ldx(Size::H, R2, R7, ARP_PTYPE));
asm.jump(
Insn::jmp_imm(Jmp::JNE, R2, host_be16(EtherType::IPV4.0), 0),
l_default,
);
if cfg.match_field.wants_dst() {
stage_v4(&mut asm, V4_DST_KEY, ARP_TPA);
}
if cfg.match_field.wants_src() {
stage_v4(&mut asm, V4_SRC_KEY, ARP_SPA);
}
if cfg.match_field.wants_dst() {
lookup_and_match(&mut asm, cfg, v4_fd, V4_DST_KEY, None, l_redirect);
}
if cfg.match_field.wants_src() {
lookup_and_match(&mut asm, cfg, v4_fd, V4_SRC_KEY, None, l_redirect);
}
asm.jump(Insn::ja(0), l_default);
}
asm.place(l_redirect);
asm.emit_all(&ld_map_fd(R1, xskmap_fd));
asm.emit(Insn::ldx(Size::W, R2, R6, XDP_MD_RX_QUEUE_INDEX));
asm.emit(Insn::mov64_imm(R3, Action::PASS.0 as i32));
asm.emit(Insn::call(BPF_FUNC_REDIRECT_MAP));
asm.emit(Insn::exit());
asm.place(l_default);
asm.emit(Insn::mov64_imm(R0, cfg.default_action.0 as i32));
asm.emit(Insn::exit());
asm.build()
}
#[derive(Debug, Clone)]
struct Entry {
prefix: IpPrefix,
rules: Vec<Rule>,
}
#[derive(Debug)]
pub struct Capture {
maps: CaptureMaps,
prog: Program,
link: Link,
cfg: CaptureConfig,
entries: Mutex<Vec<Entry>>,
}
impl Capture {
pub fn attach(ifindex: u32, cfg: CaptureConfig, mode: Mode) -> Result<Capture> {
cfg.validate()?;
let maps = CaptureMaps::create(&cfg)?;
let insns = build_program(&cfg, &maps)?;
let prog = Program::load(&insns, "pktkit_cap")?;
let link = prog.attach(ifindex, mode)?;
Ok(Capture {
maps,
prog,
link,
cfg,
entries: Mutex::new(Vec::new()),
})
}
#[inline]
pub fn mode(&self) -> Mode {
self.link.mode()
}
#[inline]
pub fn xskmap(&self) -> &Map {
&self.maps.xskmap
}
pub fn test_run(&self, frame: &[u8], repeat: u32) -> Result<TestRun> {
self.prog.test_run(frame, repeat)
}
pub fn add(&self, prefix: IpPrefix) -> Result<()> {
self.add_rule(prefix, Rule::Any)
}
pub fn add_rule(&self, prefix: IpPrefix, rule: Rule) -> Result<()> {
let prefix = prefix.masked();
self.cfg.check_prefix(prefix)?;
rule.validate()?;
let mut held = self.entries.lock().unwrap();
match held.iter().position(|e| e.prefix == prefix) {
Some(i) => {
if held[i].rules.contains(&rule) {
return Ok(());
}
if held[i].rules.len() >= usize::from(self.cfg.max_rules_per_prefix) {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"xdp: {prefix} already holds {} rules, the configured maximum",
held[i].rules.len()
),
));
}
let mut rules = held[i].rules.clone();
rules.push(rule);
self.write(prefix, &rules)?;
held[i].rules = rules;
}
None => {
let prefixes: Vec<IpPrefix> = held.iter().map(|e| e.prefix).collect();
check_coverage(&prefixes, prefix)?;
self.write(prefix, &[rule])?;
held.push(Entry {
prefix,
rules: vec![rule],
});
}
}
if rule == Rule::Any
&& let Some(sn) = self.solicited_node(prefix)
{
self.sync_solicited_node(&held, sn)?;
}
Ok(())
}
pub fn remove(&self, prefix: IpPrefix) -> Result<bool> {
let prefix = prefix.masked();
let mut held = self.entries.lock().unwrap();
let had = match held.iter().position(|e| e.prefix == prefix) {
Some(i) => {
held.remove(i);
true
}
None => false,
};
let removed = self.map_for(prefix).delete(lpm_key(prefix).as_bytes())?;
self.after_removal(&held, prefix)?;
Ok(had || removed)
}
pub fn remove_rule(&self, prefix: IpPrefix, rule: Rule) -> Result<bool> {
let prefix = prefix.masked();
let mut held = self.entries.lock().unwrap();
let Some(i) = held.iter().position(|e| e.prefix == prefix) else {
return Ok(false);
};
let Some(r) = held[i].rules.iter().position(|r| *r == rule) else {
return Ok(false);
};
let mut rules = held[i].rules.clone();
rules.remove(r);
if rules.is_empty() {
self.map_for(prefix).delete(lpm_key(prefix).as_bytes())?;
held.remove(i);
} else {
self.write(prefix, &rules)?;
held[i].rules = rules;
}
if rule == Rule::Any {
self.after_removal(&held, prefix)?;
}
Ok(true)
}
pub fn contains(&self, addr: IpAddr) -> Result<bool> {
Ok(!self.rules_for(addr)?.is_empty())
}
pub fn rules_for(&self, addr: IpAddr) -> Result<Vec<Rule>> {
let full = IpPrefix::new(addr, if addr.is_ipv4() { 32 } else { 128 });
let mut out = vec![0u8; value_size(self.cfg.max_rules_per_prefix) as usize];
if self
.map_for(full)
.lookup(lpm_key(full).as_bytes(), &mut out)?
{
Ok(decode_rules(&out))
} else {
Ok(Vec::new())
}
}
pub fn prefixes(&self) -> Vec<IpPrefix> {
self.entries
.lock()
.unwrap()
.iter()
.map(|e| e.prefix)
.collect()
}
pub fn rules(&self, prefix: IpPrefix) -> Vec<Rule> {
let prefix = prefix.masked();
self.entries
.lock()
.unwrap()
.iter()
.find(|e| e.prefix == prefix)
.map(|e| e.rules.clone())
.unwrap_or_default()
}
fn write(&self, prefix: IpPrefix, rules: &[Rule]) -> Result<()> {
self.map_for(prefix).update(
lpm_key(prefix).as_bytes(),
&encode_rules(rules, self.cfg.max_rules_per_prefix),
UpdateFlags::ANY,
)
}
#[inline]
fn map_for(&self, prefix: IpPrefix) -> &Map {
if prefix.is_v4() {
&self.maps.v4
} else {
&self.maps.v6
}
}
fn after_removal(&self, held: &[Entry], prefix: IpPrefix) -> Result<()> {
if let Some(sn) = self.solicited_node(prefix) {
self.sync_solicited_node(held, sn)?;
}
if is_solicited_node_group(prefix) {
self.sync_solicited_node(held, prefix)?;
}
Ok(())
}
fn sync_solicited_node(&self, held: &[Entry], sn: IpPrefix) -> Result<()> {
if held.iter().any(|e| e.prefix == sn) {
return Ok(());
}
let needed = held
.iter()
.any(|e| e.rules.contains(&Rule::Any) && self.solicited_node(e.prefix) == Some(sn));
if needed {
self.write(sn, &[Rule::Any])
} else {
self.map_for(sn).delete(lpm_key(sn).as_bytes()).map(|_| ())
}
}
fn solicited_node(&self, prefix: IpPrefix) -> Option<IpPrefix> {
if !self.cfg.neighbor_discovery || prefix.bits() != 128 {
return None;
}
match prefix.addr() {
IpAddr::V6(a) => Some(IpPrefix::new(solicited_node_multicast(a).into(), 128)),
IpAddr::V4(_) => None,
}
}
}
pub fn solicited_node_multicast(addr: Ipv6Addr) -> Ipv6Addr {
let o = addr.octets();
let mut sn = [0u8; 16];
sn[0] = 0xff;
sn[1] = 0x02;
sn[11] = 0x01;
sn[12] = 0xff;
sn[13..16].copy_from_slice(&o[13..16]);
Ipv6Addr::from(sn)
}
fn is_solicited_node_group(prefix: IpPrefix) -> bool {
match prefix.addr() {
IpAddr::V6(a) if prefix.bits() == 128 => {
let o = a.octets();
o[..13] == [0xff, 0x02, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0x01, 0xff]
}
_ => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::xdp::insn::{BPF_ADD, BPF_ALU64, BPF_JMP, BPF_K, BPF_STX};
use std::net::Ipv4Addr;
const UDP: Protocol = Protocol::UDP;
const TCP: Protocol = Protocol::TCP;
fn v4(a: [u8; 4], bits: u8) -> IpPrefix {
IpPrefix::new(Ipv4Addr::from(a).into(), bits)
}
fn program(cfg: &CaptureConfig) -> Vec<Insn> {
build_program_with_fds(cfg, 10, 11, 12).unwrap()
}
fn jumps(p: &[Insn]) -> Vec<usize> {
p.iter()
.enumerate()
.filter(|(_, i)| i.code & 0x07 == BPF_JMP)
.map(|(n, _)| n)
.collect()
}
#[test]
fn every_jump_lands_inside_the_program() {
for cfg in [
CaptureConfig::default(),
CaptureConfig {
match_field: MatchField::Either,
..Default::default()
},
CaptureConfig {
arp: false,
match_field: MatchField::Src,
..Default::default()
},
CaptureConfig {
max_rules_per_prefix: 1,
..Default::default()
},
CaptureConfig {
max_rules_per_prefix: MAX_RULES_PER_PREFIX,
match_field: MatchField::Either,
..Default::default()
},
] {
let p = program(&cfg);
for n in jumps(&p) {
let i = p[n];
if i.code == (BPF_JMP | 0x80) || i.code == (BPF_JMP | 0x90) {
continue;
}
let target = n as isize + 1 + i.off as isize;
assert!(
target >= 0 && target < p.len() as isize,
"jump at {n} targets {target}, program is {} insns",
p.len()
);
}
}
}
#[test]
fn program_ends_with_the_default_verdict() {
let cfg = CaptureConfig::default();
let p = program(&cfg);
let n = p.len();
assert_eq!(p[n - 1], Insn::exit());
assert_eq!(p[n - 2], Insn::mov64_imm(R0, Action::PASS.0 as i32));
}
#[test]
fn drop_default_is_honoured() {
let cfg = CaptureConfig {
default_action: Action::DROP,
..Default::default()
};
let p = program(&cfg);
assert_eq!(p[p.len() - 2], Insn::mov64_imm(R0, Action::DROP.0 as i32));
}
#[test]
fn redirect_falls_back_to_pass_on_an_unbound_queue() {
let p = program(&CaptureConfig::default());
let call = p
.iter()
.position(|i| *i == Insn::call(BPF_FUNC_REDIRECT_MAP))
.expect("redirect call present");
assert_eq!(p[call - 1], Insn::mov64_imm(R3, Action::PASS.0 as i32));
}
#[test]
fn dst_only_does_one_lookup_per_family() {
let p = program(&CaptureConfig {
match_field: MatchField::Dst,
arp: false,
..Default::default()
});
let n = p
.iter()
.filter(|i| **i == Insn::call(BPF_FUNC_MAP_LOOKUP_ELEM))
.count();
assert_eq!(n, 2, "one v4 + one v6 lookup");
}
#[test]
fn either_doubles_the_lookups() {
let p = program(&CaptureConfig {
match_field: MatchField::Either,
arp: false,
..Default::default()
});
let n = p
.iter()
.filter(|i| **i == Insn::call(BPF_FUNC_MAP_LOOKUP_ELEM))
.count();
assert_eq!(n, 4);
}
#[test]
fn arp_adds_a_third_family_branch() {
let with = program(&CaptureConfig::default());
let without = program(&CaptureConfig {
arp: false,
..Default::default()
});
assert!(with.len() > without.len());
let n = with
.iter()
.filter(|i| **i == Insn::call(BPF_FUNC_MAP_LOOKUP_ELEM))
.count();
assert_eq!(n, 3, "v4 + v6 + arp");
}
#[test]
fn packet_reads_never_follow_a_helper_call() {
let p = program(&CaptureConfig {
match_field: MatchField::Either,
..Default::default()
});
let mut seen_call = false;
for i in &p {
if *i == Insn::call(BPF_FUNC_MAP_LOOKUP_ELEM) {
seen_call = true;
}
if i.code & 0x07 == 0x01 && (i.regs >> 4) == R7 {
assert!(!seen_call, "packet read after a helper call");
}
if i.code & 0x07 == BPF_JMP && i.code != (BPF_JMP | 0x80) && i.off != 0 {
seen_call = false;
}
}
}
#[test]
fn every_staged_key_is_written_before_it_is_read() {
let p = program(&CaptureConfig {
match_field: MatchField::Either,
..Default::default()
});
let add64_imm = BPF_ALU64 | BPF_K | BPF_ADD;
let mut written: Vec<i16> = Vec::new();
for i in &p {
if i.code & 0x07 == BPF_STX && (i.regs & 0x0f) == R10 {
written.push(i.off);
}
if i.code == add64_imm && (i.regs & 0x0f) == R2 && i.imm < 0 {
let slot = i.imm as i16;
assert!(written.contains(&slot), "lookup key at {slot} never staged");
}
}
}
#[test]
fn a_default_route_is_never_capturable() {
let cfg = CaptureConfig::default();
for p in [
IpPrefix::new(Ipv4Addr::UNSPECIFIED.into(), 0),
IpPrefix::new(Ipv6Addr::UNSPECIFIED.into(), 0),
] {
let e = cfg.check_prefix(p).unwrap_err();
assert_eq!(e.kind(), io::ErrorKind::InvalidInput);
assert!(
e.to_string().contains("every packet"),
"error should say why: {e}"
);
}
}
#[test]
fn a_zero_floor_cannot_be_configured() {
for cfg in [
CaptureConfig {
min_prefix_v4: 0,
..Default::default()
},
CaptureConfig {
min_prefix_v6: 0,
..Default::default()
},
] {
assert!(cfg.validate().is_err());
}
let smuggled = CaptureConfig {
min_prefix_v4: 0,
..Default::default()
};
assert!(
smuggled
.check_prefix(IpPrefix::new(Ipv4Addr::UNSPECIFIED.into(), 0))
.is_err()
);
}
#[test]
fn ordinary_prefixes_are_accepted() {
let cfg = CaptureConfig::default();
cfg.check_prefix(v4([10, 0, 0, 7], 32)).unwrap();
cfg.check_prefix(v4([10, 0, 0, 0], 24)).unwrap();
cfg.check_prefix(v4([10, 0, 0, 0], 8)).unwrap();
cfg.check_prefix(v4([0, 0, 0, 0], 1)).unwrap();
}
#[test]
fn a_tighter_floor_is_enforced_per_family() {
let cfg = CaptureConfig {
min_prefix_v4: 24,
min_prefix_v6: 64,
..Default::default()
};
cfg.validate().unwrap();
cfg.check_prefix(v4([10, 0, 0, 0], 24)).unwrap();
assert!(cfg.check_prefix(v4([10, 0, 0, 0], 16)).is_err());
let net: Ipv6Addr = "2001:db8::".parse().unwrap();
cfg.check_prefix(IpPrefix::new(net.into(), 64)).unwrap();
assert!(cfg.check_prefix(IpPrefix::new(net.into(), 48)).is_err());
}
#[test]
fn a_floor_wider_than_the_family_is_rejected() {
assert!(
CaptureConfig {
min_prefix_v4: 33,
..Default::default()
}
.validate()
.is_err()
);
assert!(
CaptureConfig {
min_prefix_v6: 129,
..Default::default()
}
.validate()
.is_err()
);
}
#[test]
fn two_halves_cannot_add_up_to_the_whole_interface() {
let low = v4([0, 0, 0, 0], 1);
let high = v4([128, 0, 0, 0], 1);
check_coverage(&[], low).unwrap();
let e = check_coverage(&[low], high).unwrap_err();
assert!(e.to_string().contains("every IPv4 address"), "{e}");
}
#[test]
fn four_quarters_cannot_either() {
let quarters: Vec<IpPrefix> = [0u8, 64, 128, 192]
.iter()
.map(|&a| v4([a, 0, 0, 0], 2))
.collect();
for i in 0..3 {
check_coverage(&quarters[..i], quarters[i]).unwrap();
}
assert!(check_coverage(&quarters[..3], quarters[3]).is_err());
}
#[test]
fn ipv6_halves_are_caught_without_overflowing() {
let low = IpPrefix::new("::".parse::<Ipv6Addr>().unwrap().into(), 1);
let high = IpPrefix::new("8000::".parse::<Ipv6Addr>().unwrap().into(), 1);
check_coverage(&[], low).unwrap();
assert!(check_coverage(&[low], high).is_err());
}
#[test]
fn coverage_is_counted_per_family() {
let v6_low = IpPrefix::new("::".parse::<Ipv6Addr>().unwrap().into(), 1);
let v6_high = IpPrefix::new("8000::".parse::<Ipv6Addr>().unwrap().into(), 1);
check_coverage(&[v6_low, v6_high], v4([10, 0, 0, 1], 32)).unwrap();
}
#[test]
fn realistic_sets_stay_far_from_the_limit() {
let mut held: Vec<IpPrefix> = (0..1000)
.map(|i| v4([10, (i / 256) as u8, (i % 256) as u8, 1], 32))
.collect();
held.push(v4([192, 168, 0, 0], 16));
held.push(v4([172, 16, 0, 0], 12));
check_coverage(&held, v4([10, 0, 0, 0], 8)).unwrap();
}
#[test]
fn coverage_totals_are_exact_at_the_boundary() {
assert_eq!(coverage(&[v4([10, 0, 0, 1], 32)], true), 1);
assert_eq!(coverage(&[v4([10, 0, 0, 0], 24)], true), 256);
assert_eq!(coverage(&[v4([0, 0, 0, 0], 1)], true), 1 << 31);
assert_eq!(family_total(true), 1u128 << 32);
assert_eq!(coverage(&[], false), 0);
}
#[test]
fn default_action_must_be_a_terminal_verdict() {
for a in [Action::PASS, Action::DROP] {
CaptureConfig {
default_action: a,
..Default::default()
}
.validate()
.unwrap();
}
for a in [Action::REDIRECT, Action::TX, Action::ABORTED] {
assert!(
CaptureConfig {
default_action: a,
..Default::default()
}
.validate()
.is_err()
);
}
}
#[test]
fn the_default_configuration_validates() {
CaptureConfig::default().validate().unwrap();
}
#[test]
fn solicited_node_follows_rfc4291() {
let a: Ipv6Addr = "2001:db8::dead:beef".parse().unwrap();
let sn = solicited_node_multicast(a);
assert_eq!(sn, "ff02::1:ffad:beef".parse::<Ipv6Addr>().unwrap());
}
#[test]
fn solicited_node_only_depends_on_the_low_24_bits() {
let a: Ipv6Addr = "2001:db8::1:2:3".parse().unwrap();
let b: Ipv6Addr = "fe80::ffff:1:2:3".parse().unwrap();
assert_eq!(solicited_node_multicast(a), solicited_node_multicast(b));
}
#[test]
fn ethertype_constants_are_compared_in_wire_order() {
let p = program(&CaptureConfig::default());
let load = p
.iter()
.position(|i| *i == Insn::ldx(Size::H, R2, R7, ETH_TYPE))
.unwrap();
assert_eq!(p[load + 1].imm, host_be16(0x0800));
assert_eq!(p[load + 2].imm, host_be16(0x86DD));
assert_eq!(p[load + 3].imm, host_be16(0x0806));
}
#[test]
fn v4_prefix_round_trips_through_a_key() {
let p = IpPrefix::new(Ipv4Addr::new(198, 51, 100, 7).into(), 32);
assert_eq!(lpm_key(p).as_bytes()[4..], [198, 51, 100, 7]);
}
#[test]
fn a_port_rule_needs_tcp_or_udp() {
Rule::Port(TCP, 443).validate().unwrap();
Rule::Port(UDP, 53).validate().unwrap();
Rule::Proto(Protocol::ICMP).validate().unwrap();
Rule::Any.validate().unwrap();
for p in [Protocol::ICMP, Protocol::GRE, Protocol(132)] {
let e = Rule::Port(p, 80).validate().unwrap_err();
assert_eq!(e.kind(), io::ErrorKind::InvalidInput);
}
}
#[test]
fn rules_round_trip_through_the_value() {
let rules = [Rule::Any, Rule::Proto(Protocol::GRE), Rule::Port(TCP, 443)];
let v = encode_rules(&rules, 8);
assert_eq!(v.len(), 32);
assert_eq!(decode_rules(&v), rules);
assert_eq!(v[12], KIND_END);
}
#[test]
fn a_port_is_stored_in_wire_order() {
assert_eq!(
Rule::Port(UDP, 0x1234).encode(),
[KIND_PORT, 17, 0x12, 0x34]
);
}
#[test]
fn an_empty_list_decodes_to_nothing() {
assert!(decode_rules(&encode_rules(&[], 4)).is_empty());
}
#[test]
fn value_size_follows_the_rule_cap() {
assert_eq!(value_size(1), 4);
assert_eq!(value_size(8), 32);
assert_eq!(value_size(MAX_RULES_PER_PREFIX), 256);
}
#[test]
fn the_rule_cap_is_bounded() {
for n in [0, MAX_RULES_PER_PREFIX + 1] {
assert!(
CaptureConfig {
max_rules_per_prefix: n,
..Default::default()
}
.validate()
.is_err()
);
}
CaptureConfig {
max_rules_per_prefix: MAX_RULES_PER_PREFIX,
..Default::default()
}
.validate()
.unwrap();
}
#[test]
fn solicited_node_groups_are_recognised() {
let sn: Ipv6Addr = "ff02::1:ffad:beef".parse().unwrap();
assert!(is_solicited_node_group(IpPrefix::new(sn.into(), 128)));
assert!(!is_solicited_node_group(IpPrefix::new(sn.into(), 104)));
let other: Ipv6Addr = "ff02::16".parse().unwrap();
assert!(!is_solicited_node_group(IpPrefix::new(other.into(), 128)));
assert!(!is_solicited_node_group(v4([224, 0, 0, 1], 32)));
}
#[test]
fn rule_reads_stay_inside_the_value() {
for n in [1u8, 3, 8, MAX_RULES_PER_PREFIX] {
let cfg = CaptureConfig {
max_rules_per_prefix: n,
match_field: MatchField::Either,
..Default::default()
};
let size = value_size(n) as i16;
for i in program(&cfg) {
if i.code & 0x07 == 0x01 && (i.regs >> 4) == R0 {
let width = match Size(i.code & 0x18) {
Size::B => 1,
Size::H => 2,
Size::W => 4,
_ => 8,
};
assert!(
i.off >= 0 && i.off + width <= size,
"read at {} past {size}",
i.off
);
}
}
}
}
#[test]
fn the_arp_branch_never_compares_a_protocol() {
let p = program(&CaptureConfig {
match_field: MatchField::Dst,
..Default::default()
});
let arp = p
.iter()
.position(|i| *i == Insn::ldx(Size::H, R2, R7, ARP_PTYPE))
.expect("arp branch present");
let redirect = p
.iter()
.position(|i| *i == Insn::call(BPF_FUNC_REDIRECT_MAP))
.unwrap();
for i in &p[arp..redirect] {
assert_ne!(*i, Insn::mov64_imm(RULE_PORT, NO_PORT));
assert_ne!(*i, Insn::ldx(Size::B, R2, R0, 1));
}
}
#[test]
fn the_transport_header_is_not_parsed_before_a_lookup() {
let p = program(&CaptureConfig::default());
let first_call = p
.iter()
.position(|i| *i == Insn::call(BPF_FUNC_MAP_LOOKUP_ELEM))
.unwrap();
let head = &p[..first_call];
assert!(!head.contains(&Insn::mov64_imm(RULE_PORT, NO_PORT)));
assert!(!head.contains(&Insn::ldx(Size::B, RULE_PROTO, R7, IPV4_PROTO)));
let parses = p
.iter()
.filter(|i| **i == Insn::mov64_imm(RULE_PORT, NO_PORT))
.count();
assert_eq!(parses, 2, "one v4 site and one v6 site under Dst");
}
#[test]
fn any_is_encoded_first() {
let rules = [Rule::Port(TCP, 443), Rule::Any, Rule::Proto(Protocol::GRE)];
assert_eq!(
decode_rules(&encode_rules(&rules, 8)),
[Rule::Any, Rule::Port(TCP, 443), Rule::Proto(Protocol::GRE)]
);
}
mod vm {
use super::*;
const PKT: u64 = 0x1000_0000;
const STACK_TOP: u64 = 0x2000_0200;
const VALUE: u64 = 0x3000_0000;
const CTX: u64 = 0x4000_0000;
const MAP: u64 = 0x5000_0000;
pub const XSK_FD: i32 = 10;
pub const V4_FD: i32 = 11;
pub const V6_FD: i32 = 12;
pub struct Trie {
pub addr_len: usize,
pub entries: Vec<(u32, Vec<u8>, Vec<u8>)>,
}
impl Trie {
fn lookup(&self, key: &[u8]) -> Option<&[u8]> {
assert_eq!(key.len(), 4 + self.addr_len, "key does not fit this trie");
let bits = u32::from_ne_bytes(key[..4].try_into().unwrap());
let addr = &key[4..];
self.entries
.iter()
.filter(|(plen, paddr, _)| *plen <= bits && prefix_eq(paddr, addr, *plen))
.max_by_key(|(plen, _, _)| *plen)
.map(|(_, _, v)| v.as_slice())
}
}
fn prefix_eq(a: &[u8], b: &[u8], bits: u32) -> bool {
let full = (bits / 8) as usize;
if a[..full] != b[..full] {
return false;
}
let rem = bits % 8;
rem == 0 || {
let mask = 0xffu8 << (8 - rem);
a[full] & mask == b[full] & mask
}
}
pub struct Vm<'a> {
pub v4: Trie,
pub v6: Trie,
pub xsk_queues: Vec<u32>,
pub rx_queue: u32,
pub steps: usize,
pkt: &'a [u8],
stack: [u8; 512],
value: Vec<u8>,
ctx: [u8; 20],
regs: [u64; 11],
}
impl<'a> Vm<'a> {
pub fn new(pkt: &'a [u8], v4: Trie, v6: Trie) -> Vm<'a> {
Vm {
v4,
v6,
xsk_queues: vec![0],
rx_queue: 0,
steps: 0,
pkt,
stack: [0; 512],
value: Vec::new(),
ctx: [0; 20],
regs: [0; 11],
}
}
fn mem(&mut self, addr: u64, len: usize) -> &mut [u8] {
let (base, buf): (u64, &mut [u8]) = if (STACK_TOP - 512..STACK_TOP).contains(&addr)
{
(STACK_TOP - 512, &mut self.stack[..])
} else if (VALUE..VALUE + self.value.len() as u64).contains(&addr) {
(VALUE, &mut self.value[..])
} else if (CTX..CTX + 20).contains(&addr) {
(CTX, &mut self.ctx[..])
} else {
panic!("access to unmapped address {addr:#x}")
};
let off = (addr - base) as usize;
assert!(off + len <= buf.len(), "access past the end of a region");
&mut buf[off..off + len]
}
fn load(&mut self, addr: u64, len: usize) -> u64 {
let mut b = [0u8; 8];
if (PKT..PKT + self.pkt.len() as u64).contains(&addr) {
let off = (addr - PKT) as usize;
assert!(
off + len <= self.pkt.len(),
"packet read at {off}+{len} past data_end — the verifier would reject this"
);
b[..len].copy_from_slice(&self.pkt[off..off + len]);
} else {
b[..len].copy_from_slice(self.mem(addr, len));
}
u64::from_ne_bytes(b)
}
fn store(&mut self, addr: u64, len: usize, v: u64) {
let b = v.to_ne_bytes();
self.mem(addr, len).copy_from_slice(&b[..len]);
}
fn call(&mut self, func: i32) {
match func {
BPF_FUNC_MAP_LOOKUP_ELEM => {
let fd = (self.regs[1] - MAP) as i32;
let addr_len = match fd {
V4_FD => 4,
V6_FD => 16,
_ => panic!("lookup on non-trie fd {fd}"),
};
let key = self.load_bytes(self.regs[2], 4 + addr_len);
let trie = if fd == V4_FD { &self.v4 } else { &self.v6 };
match trie.lookup(&key).map(<[u8]>::to_vec) {
Some(v) => {
self.value = v;
self.regs[0] = VALUE;
}
None => self.regs[0] = 0,
}
}
BPF_FUNC_REDIRECT_MAP => {
assert_eq!(self.regs[1], MAP + XSK_FD as u64);
let q = self.regs[2] as u32;
self.regs[0] = if self.xsk_queues.contains(&q) {
Action::REDIRECT.0 as u64
} else {
self.regs[3] & 0xf
};
}
_ => panic!("unknown helper {func}"),
}
for r in 1..=5 {
self.regs[r] = 0xdead_beef_dead_beef;
}
}
fn load_bytes(&mut self, addr: u64, len: usize) -> Vec<u8> {
(0..len)
.map(|i| self.load(addr + i as u64, 1) as u8)
.collect()
}
pub fn run(&mut self, prog: &[Insn]) -> u32 {
let end = PKT + self.pkt.len() as u64;
self.ctx[..4].copy_from_slice(&(PKT as u32).to_ne_bytes());
self.ctx[4..8].copy_from_slice(&(end as u32).to_ne_bytes());
self.ctx[16..20].copy_from_slice(&self.rx_queue.to_ne_bytes());
self.regs[1] = CTX;
self.regs[10] = STACK_TOP;
let mut pc = 0usize;
self.steps = 0;
loop {
self.steps += 1;
assert!(self.steps < 10_000, "program does not terminate");
let i = prog[pc];
let dst = (i.regs & 0x0f) as usize;
let src = (i.regs >> 4) as usize;
let class = i.code & 0x07;
pc += 1;
match class {
0x00 => {
assert_eq!(i.code, 0x18);
self.regs[dst] = MAP + i.imm as u64;
pc += 1;
}
0x01 => {
let len = width(i.code);
let addr = self.regs[src].wrapping_add(i.off as i64 as u64);
self.regs[dst] = self.load(addr, len);
}
0x02 | 0x03 => {
let len = width(i.code);
let addr = self.regs[dst].wrapping_add(i.off as i64 as u64);
let v = if class == 0x02 {
i.imm as u64
} else {
self.regs[src]
};
self.store(addr, len, v);
}
0x07 => {
let operand = if i.code & 0x08 != 0 {
self.regs[src]
} else {
i.imm as i64 as u64
};
match i.code & 0xf0 {
0xb0 => self.regs[dst] = operand,
0x00 => self.regs[dst] = self.regs[dst].wrapping_add(operand),
0x50 => self.regs[dst] &= operand,
0x60 => self.regs[dst] <<= operand,
op => panic!("unsupported alu op {op:#x}"),
}
}
0x05 => {
let op = i.code & 0xf0;
if op == 0x80 {
self.call(i.imm);
continue;
}
if op == 0x90 {
return self.regs[0] as u32;
}
let a = self.regs[dst];
let b = if i.code & 0x08 != 0 {
self.regs[src]
} else {
i.imm as i64 as u64
};
let taken = match op {
0x00 => true,
0x10 => a == b,
0x20 => a > b,
0x30 => a >= b,
0x40 => a & b != 0,
0x50 => a != b,
0xa0 => a < b,
op => panic!("unsupported jump op {op:#x}"),
};
if taken {
pc = (pc as isize + i.off as isize) as usize;
}
}
c => panic!("unsupported class {c:#x}"),
}
}
}
}
fn width(code: u8) -> usize {
match Size(code & 0x18) {
Size::B => 1,
Size::H => 2,
Size::W => 4,
_ => 8,
}
}
}
use vm::{Trie, Vm};
fn tries(cfg: &CaptureConfig, set: &[(IpPrefix, &[Rule])]) -> (Trie, Trie) {
let mut v4 = Trie {
addr_len: 4,
entries: Vec::new(),
};
let mut v6 = Trie {
addr_len: 16,
entries: Vec::new(),
};
for (prefix, rules) in set {
let key = lpm_key(*prefix);
let entry = (
u32::from(prefix.bits()),
key.as_bytes()[4..].to_vec(),
encode_rules(rules, cfg.max_rules_per_prefix),
);
if prefix.is_v4() {
v4.entries.push(entry);
} else {
v6.entries.push(entry);
}
}
(v4, v6)
}
fn verdict(cfg: &CaptureConfig, set: &[(IpPrefix, &[Rule])], pkt: &[u8]) -> u32 {
let prog = build_program_with_fds(cfg, vm::XSK_FD, vm::V4_FD, vm::V6_FD).unwrap();
let (v4, v6) = tries(cfg, set);
Vm::new(pkt, v4, v6).run(&prog)
}
fn steps(cfg: &CaptureConfig, set: &[(IpPrefix, &[Rule])], pkt: &[u8]) -> usize {
let prog = build_program_with_fds(cfg, vm::XSK_FD, vm::V4_FD, vm::V6_FD).unwrap();
let (v4, v6) = tries(cfg, set);
let mut vm = Vm::new(pkt, v4, v6);
vm.run(&prog);
vm.steps
}
const REDIRECT: u32 = Action::REDIRECT.0;
const PASS: u32 = Action::PASS.0;
fn eth(ethertype: u16, payload: &[u8]) -> Vec<u8> {
let mut f = vec![0x02, 0, 0, 0, 0, 1, 0x02, 0, 0, 0, 0, 2];
f.extend_from_slice(ðertype.to_be_bytes());
f.extend_from_slice(payload);
f
}
fn ipv4_with(
proto: Protocol,
src: [u8; 4],
dst: [u8; 4],
ihl: u8,
frag: u16,
l4: &[u8],
) -> Vec<u8> {
let mut h = vec![0u8; usize::from(ihl) * 4];
h[0] = 0x40 | ihl;
let total = (h.len() + l4.len()) as u16;
h[2..4].copy_from_slice(&total.to_be_bytes());
h[6..8].copy_from_slice(&frag.to_be_bytes());
h[8] = 64;
h[9] = proto.as_u8();
h[12..16].copy_from_slice(&src);
h[16..20].copy_from_slice(&dst);
h.extend_from_slice(l4);
eth(EtherType::IPV4.0, &h)
}
fn ipv4(proto: Protocol, src: [u8; 4], dst: [u8; 4], l4: &[u8]) -> Vec<u8> {
ipv4_with(proto, src, dst, 5, 0, l4)
}
fn ipv6(next: Protocol, src: &str, dst: &str, l4: &[u8]) -> Vec<u8> {
let mut h = vec![0u8; 40];
h[0] = 0x60;
h[4..6].copy_from_slice(&(l4.len() as u16).to_be_bytes());
h[6] = next.as_u8();
h[7] = 64;
h[8..24].copy_from_slice(&src.parse::<Ipv6Addr>().unwrap().octets());
h[24..40].copy_from_slice(&dst.parse::<Ipv6Addr>().unwrap().octets());
h.extend_from_slice(l4);
eth(EtherType::IPV6.0, &h)
}
fn ports(sport: u16, dport: u16) -> Vec<u8> {
let mut l4 = Vec::new();
l4.extend_from_slice(&sport.to_be_bytes());
l4.extend_from_slice(&dport.to_be_bytes());
l4.extend_from_slice(&[0u8; 16]);
l4
}
fn arp(spa: [u8; 4], tpa: [u8; 4]) -> Vec<u8> {
let mut a = vec![0u8; 28];
a[..2].copy_from_slice(&1u16.to_be_bytes());
a[2..4].copy_from_slice(&EtherType::IPV4.0.to_be_bytes());
a[4] = 6;
a[5] = 4;
a[6..8].copy_from_slice(&1u16.to_be_bytes());
a[14..18].copy_from_slice(&spa);
a[24..28].copy_from_slice(&tpa);
eth(EtherType::ARP.0, &a)
}
const HOST: [u8; 4] = [10, 0, 0, 7];
const PEER: [u8; 4] = [10, 0, 0, 9];
#[test]
fn any_rule_takes_every_protocol_on_the_address() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Any])];
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(1, 2))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &ports(1, 2))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(Protocol::ICMP, PEER, HOST, &[8, 0, 0, 0])),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, HOST, PEER, &ports(1, 2))),
PASS
);
}
#[test]
fn proto_rule_takes_one_protocol_only() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Proto(UDP)])];
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(1000, 53))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &ports(1000, 53))),
PASS
);
assert_eq!(
verdict(&cfg, set, &ipv4(Protocol::ICMP, PEER, HOST, &[8, 0, 0, 0])),
PASS
);
}
#[test]
fn port_rule_takes_one_port_of_one_protocol() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Port(UDP, 51820)])];
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(4000, 51820))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &ports(4000, 51820))),
PASS
);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(4000, 53))),
PASS
);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(51820, 4000))),
PASS
);
}
#[test]
fn port_is_found_behind_ipv4_options() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Port(TCP, 443)])];
for ihl in [5u8, 6, 8, 15] {
let f = ipv4_with(TCP, PEER, HOST, ihl, 0, &ports(4000, 443));
assert_eq!(verdict(&cfg, set, &f), REDIRECT, "ihl={ihl}");
let f = ipv4_with(TCP, PEER, HOST, ihl, 0, &ports(4000, 80));
assert_eq!(verdict(&cfg, set, &f), PASS, "ihl={ihl}");
}
}
#[test]
fn a_bogus_ihl_cannot_match_a_port() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Port(TCP, 0x0a00)])];
let mut f = ipv4_with(TCP, PEER, HOST, 5, 0, &ports(1, 2));
f[14] = 0x44;
assert_eq!(verdict(&cfg, set, &f), PASS);
}
#[test]
fn only_the_first_fragment_carries_a_port() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Port(UDP, 53)])];
let f = ipv4_with(UDP, PEER, HOST, 5, 0x2000, &ports(4000, 53));
assert_eq!(verdict(&cfg, set, &f), REDIRECT);
let f = ipv4_with(UDP, PEER, HOST, 5, 0x2000 | 185, &ports(4000, 53));
assert_eq!(verdict(&cfg, set, &f), PASS);
let f = ipv4_with(UDP, PEER, HOST, 5, 185, &ports(4000, 53));
assert_eq!(verdict(&cfg, set, &f), PASS);
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Proto(UDP)])];
assert_eq!(verdict(&cfg, set, &f), REDIRECT);
}
#[test]
fn a_truncated_transport_header_cannot_match_a_port() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] =
&[(v4(HOST, 32), &[Rule::Port(UDP, 53), Rule::Proto(TCP)])];
assert_eq!(verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &[])), PASS);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &[0x00, 0x35, 0x00])),
PASS
);
assert_eq!(verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &[])), REDIRECT);
}
#[test]
fn the_rule_list_is_walked_in_full() {
let cfg = CaptureConfig::default();
let rules: &[Rule] = &[
Rule::Proto(Protocol::ICMP),
Rule::Port(TCP, 443),
Rule::Port(UDP, 53),
Rule::Port(TCP, 22),
];
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), rules)];
assert_eq!(
verdict(&cfg, set, &ipv4(Protocol::ICMP, PEER, HOST, &[8, 0, 0, 0])),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &ports(1, 443))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(1, 53))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &ports(1, 22))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &ports(1, 80))),
PASS
);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(1, 443))),
PASS
);
assert_eq!(
verdict(&cfg, set, &ipv4(Protocol::GRE, PEER, HOST, &[0; 4])),
PASS
);
}
#[test]
fn a_full_rule_list_has_no_terminator_and_still_stops() {
let cfg = CaptureConfig {
max_rules_per_prefix: 2,
..Default::default()
};
let rules: &[Rule] = &[Rule::Port(TCP, 1), Rule::Port(TCP, 2)];
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), rules)];
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &ports(9, 2))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &ports(9, 3))),
PASS
);
}
#[test]
fn any_anywhere_in_the_list_wins() {
let cfg = CaptureConfig::default();
let rules: &[Rule] = &[Rule::Port(TCP, 1), Rule::Any];
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), rules)];
assert_eq!(
verdict(&cfg, set, &ipv4(Protocol::GRE, PEER, HOST, &[0; 4])),
REDIRECT
);
}
#[test]
fn any_is_honoured_wherever_it_sits_in_the_value() {
let cfg = CaptureConfig::default();
let mut value = encode_rules(&[Rule::Port(TCP, 1)], cfg.max_rules_per_prefix);
value[RULE_SIZE..2 * RULE_SIZE].copy_from_slice(&Rule::Any.encode());
let (mut v4t, v6t) = tries(&cfg, &[]);
v4t.entries.push((32, HOST.to_vec(), value));
let prog = build_program_with_fds(&cfg, vm::XSK_FD, vm::V4_FD, vm::V6_FD).unwrap();
let f = ipv4(Protocol::GRE, PEER, HOST, &[0; 4]);
assert_eq!(Vm::new(&f, v4t, v6t).run(&prog), REDIRECT);
}
#[test]
fn a_miss_is_the_cheapest_path_through_the_program() {
let cfg = CaptureConfig::default();
let any: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Any])];
let port: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Port(UDP, 53)])];
let hit = ipv4(UDP, PEER, HOST, &ports(1, 53));
let miss = ipv4(UDP, HOST, PEER, &ports(1, 53));
assert_eq!(steps(&cfg, port, &miss), 23);
assert_eq!(steps(&cfg, any, &miss), steps(&cfg, port, &miss));
assert!(steps(&cfg, any, &hit) < steps(&cfg, port, &hit));
assert!(steps(&cfg, port, &hit) > steps(&cfg, port, &miss));
}
#[test]
fn src_match_judges_the_source_port() {
let cfg = CaptureConfig {
match_field: MatchField::Src,
..Default::default()
};
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Port(UDP, 51820)])];
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, HOST, PEER, &ports(51820, 4000))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, HOST, PEER, &ports(4000, 51820))),
PASS
);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(4000, 51820))),
PASS
);
}
#[test]
fn either_falls_through_to_the_source_lookup_after_a_port_mismatch() {
let cfg = CaptureConfig {
match_field: MatchField::Either,
..Default::default()
};
let set: &[(IpPrefix, &[Rule])] = &[
(v4(HOST, 32), &[Rule::Port(UDP, 51820)]),
(v4(PEER, 32), &[Rule::Port(UDP, 4000)]),
];
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(4000, 53))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(5000, 53))),
PASS
);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, PEER, HOST, &ports(9, 51820))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(UDP, HOST, PEER, &ports(51820, 9))),
REDIRECT
);
}
#[test]
fn a_subnet_rule_applies_to_every_address_in_it() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] = &[(v4([10, 0, 0, 0], 24), &[Rule::Port(TCP, 80)])];
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, [10, 0, 0, 200], &ports(1, 80))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, [10, 0, 1, 200], &ports(1, 80))),
PASS
);
let set: &[(IpPrefix, &[Rule])] = &[
(v4([10, 0, 0, 0], 24), &[Rule::Port(TCP, 80)]),
(v4(HOST, 32), &[Rule::Port(TCP, 22)]),
];
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &ports(1, 22))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv4(TCP, PEER, HOST, &ports(1, 80))),
PASS
);
}
#[test]
fn arp_is_captured_only_under_any() {
let cfg = CaptureConfig::default();
let who_has = arp(PEER, HOST);
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Any])];
assert_eq!(verdict(&cfg, set, &who_has), REDIRECT);
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Port(UDP, 51820)])];
assert_eq!(verdict(&cfg, set, &who_has), PASS);
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Proto(UDP)])];
assert_eq!(verdict(&cfg, set, &who_has), PASS);
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Proto(UDP), Rule::Any])];
assert_eq!(verdict(&cfg, set, &who_has), REDIRECT);
assert_eq!(verdict(&cfg, set, &arp(HOST, PEER)), PASS);
}
#[test]
fn arp_follows_the_sender_under_src() {
let cfg = CaptureConfig {
match_field: MatchField::Src,
..Default::default()
};
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Any])];
assert_eq!(verdict(&cfg, set, &arp(HOST, PEER)), REDIRECT);
assert_eq!(verdict(&cfg, set, &arp(PEER, HOST)), PASS);
}
#[test]
fn arp_is_off_when_disabled() {
let cfg = CaptureConfig {
arp: false,
..Default::default()
};
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Any])];
assert_eq!(verdict(&cfg, set, &arp(PEER, HOST)), PASS);
}
const HOST6: &str = "2001:db8::7";
const PEER6: &str = "2001:db8::9";
fn v6(addr: &str, bits: u8) -> IpPrefix {
IpPrefix::new(addr.parse::<Ipv6Addr>().unwrap().into(), bits)
}
#[test]
fn ipv6_rules_behave_like_ipv4_ones() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] = &[(v6(HOST6, 128), &[Rule::Port(UDP, 51820)])];
assert_eq!(
verdict(&cfg, set, &ipv6(UDP, PEER6, HOST6, &ports(4000, 51820))),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv6(UDP, PEER6, HOST6, &ports(4000, 53))),
PASS
);
assert_eq!(
verdict(&cfg, set, &ipv6(TCP, PEER6, HOST6, &ports(4000, 51820))),
PASS
);
assert_eq!(
verdict(&cfg, set, &ipv6(UDP, HOST6, PEER6, &ports(51820, 4000))),
PASS
);
let set: &[(IpPrefix, &[Rule])] =
&[(v6("2001:db8::", 64), &[Rule::Proto(Protocol::ICMPV6)])];
assert_eq!(
verdict(
&cfg,
set,
&ipv6(Protocol::ICMPV6, PEER6, HOST6, &[128, 0, 0, 0])
),
REDIRECT
);
assert_eq!(
verdict(&cfg, set, &ipv6(UDP, PEER6, HOST6, &ports(1, 2))),
PASS
);
let set: &[(IpPrefix, &[Rule])] = &[(v6(HOST6, 128), &[Rule::Any])];
assert_eq!(
verdict(&cfg, set, &ipv6(Protocol::GRE, PEER6, HOST6, &[0; 4])),
REDIRECT
);
}
#[test]
fn ipv6_extension_headers_are_not_walked() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] = &[(v6(HOST6, 128), &[Rule::Port(UDP, 53)])];
let mut ext = vec![UDP.as_u8(), 0, 0, 53, 0, 0, 0, 1];
ext.extend_from_slice(&ports(4000, 53));
assert_eq!(
verdict(&cfg, set, &ipv6(Protocol(44), PEER6, HOST6, &ext)),
PASS
);
let mut ext = vec![UDP.as_u8(), 0];
ext.extend_from_slice(&53u16.to_be_bytes());
ext.extend_from_slice(&[0; 20]);
assert_eq!(
verdict(&cfg, set, &ipv6(Protocol(44), PEER6, HOST6, &ext)),
PASS
);
let set: &[(IpPrefix, &[Rule])] = &[(v6(HOST6, 128), &[Rule::Proto(Protocol(44))])];
assert_eq!(
verdict(&cfg, set, &ipv6(Protocol(44), PEER6, HOST6, &ext)),
REDIRECT
);
}
#[test]
fn frames_too_short_for_their_header_take_the_default() {
let cfg = CaptureConfig::default();
let set: &[(IpPrefix, &[Rule])] =
&[(v4(HOST, 32), &[Rule::Any]), (v6(HOST6, 128), &[Rule::Any])];
let f = ipv4(UDP, PEER, HOST, &ports(1, 2));
assert_eq!(verdict(&cfg, set, &f[..30]), PASS);
let f = ipv6(UDP, PEER6, HOST6, &ports(1, 2));
assert_eq!(verdict(&cfg, set, &f[..50]), PASS);
assert_eq!(verdict(&cfg, set, &f[..10]), PASS);
let f = arp(PEER, HOST);
assert_eq!(verdict(&cfg, set, &f[..40]), PASS);
let f = ipv4(UDP, PEER, HOST, &[]);
assert_eq!(verdict(&cfg, set, &f), REDIRECT);
let f = ipv6(UDP, PEER6, HOST6, &[]);
assert_eq!(verdict(&cfg, set, &f), REDIRECT);
}
#[test]
fn unmatched_traffic_takes_the_configured_default() {
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Any])];
let f = ipv4(UDP, HOST, PEER, &ports(1, 2));
let drop = CaptureConfig {
default_action: Action::DROP,
..Default::default()
};
assert_eq!(verdict(&drop, set, &f), Action::DROP.0);
assert_eq!(
verdict(&drop, set, &ipv4(UDP, PEER, HOST, &ports(1, 2))),
REDIRECT
);
assert_eq!(verdict(&CaptureConfig::default(), set, &f), PASS);
assert_eq!(verdict(&drop, set, ð(0x88cc, &[0; 40])), Action::DROP.0);
}
#[test]
fn a_queue_with_no_socket_passes_to_the_host() {
let cfg = CaptureConfig::default();
let prog = build_program_with_fds(&cfg, vm::XSK_FD, vm::V4_FD, vm::V6_FD).unwrap();
let set: &[(IpPrefix, &[Rule])] = &[(v4(HOST, 32), &[Rule::Any])];
let (v4t, v6t) = tries(&cfg, set);
let f = ipv4(UDP, PEER, HOST, &ports(1, 2));
let mut vm = Vm::new(&f, v4t, v6t);
vm.xsk_queues = vec![0, 1];
vm.rx_queue = 3;
assert_eq!(vm.run(&prog), PASS);
vm.rx_queue = 1;
assert_eq!(vm.run(&prog), REDIRECT);
}
}