use std::net::IpAddr;
use flowscope::L4Proto;
use crate::config::ipnet::IpNet;
#[derive(Debug, Clone, PartialEq)]
pub enum Predicate {
Always,
Atom(Atom),
And(Box<Predicate>, Box<Predicate>),
Or(Box<Predicate>, Box<Predicate>),
Not(Box<Predicate>),
}
impl Predicate {
pub fn and(self, other: Predicate) -> Predicate {
match (self, other) {
(Predicate::Always, p) | (p, Predicate::Always) => p,
(a, b) => Predicate::And(Box::new(a), Box::new(b)),
}
}
pub fn or(self, other: Predicate) -> Predicate {
match (self, other) {
(Predicate::Always, _) | (_, Predicate::Always) => Predicate::Always,
(a, b) => Predicate::Or(Box::new(a), Box::new(b)),
}
}
pub fn negate(self) -> Predicate {
Predicate::Not(Box::new(self))
}
pub fn eval(&self, src: &dyn FieldSource) -> bool {
match self {
Predicate::Always => true,
Predicate::Atom(a) => a.eval(src),
Predicate::And(l, r) => l.eval(src) && r.eval(src),
Predicate::Or(l, r) => l.eval(src) || r.eval(src),
Predicate::Not(p) => !p.eval(src),
}
}
pub fn is_fully_kernel_pushable(&self) -> bool {
match self {
Predicate::Always => true,
Predicate::Atom(a) => a.is_kernel_pushable(),
Predicate::And(l, r) | Predicate::Or(l, r) => {
l.is_fully_kernel_pushable() && r.is_fully_kernel_pushable()
}
Predicate::Not(p) => p.is_fully_kernel_pushable(),
}
}
pub fn kernel_approx(&self) -> Predicate {
match self {
Predicate::Always => Predicate::Always,
Predicate::Atom(a) if a.is_kernel_pushable() => Predicate::Atom(a.clone()),
Predicate::Atom(_) => Predicate::Always,
Predicate::And(l, r) => l.kernel_approx().and(r.kernel_approx()),
Predicate::Or(l, r) => l.kernel_approx().or(r.kernel_approx()),
Predicate::Not(p) if p.is_fully_kernel_pushable() => {
Predicate::Not(Box::new(p.kernel_approx()))
}
Predicate::Not(_) => Predicate::Always,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum Atom {
Proto(L4Proto),
SrcPort(u16),
DstPort(u16),
AnyPort(u16),
SrcHost(IpAddr),
DstHost(IpAddr),
AnyHost(IpAddr),
SrcNet(IpNet),
DstNet(IpNet),
AnyNet(IpNet),
VlanId(u16),
EtherType(u16),
SniGlob(Glob),
HttpHostGlob(Glob),
DnsQnameGlob(Glob),
BytesOver(u64),
PacketsOver(u64),
}
impl Atom {
pub fn is_kernel_pushable(&self) -> bool {
matches!(
self,
Atom::Proto(_)
| Atom::SrcPort(_)
| Atom::DstPort(_)
| Atom::AnyPort(_)
| Atom::SrcHost(_)
| Atom::DstHost(_)
| Atom::AnyHost(_)
| Atom::SrcNet(_)
| Atom::DstNet(_)
| Atom::AnyNet(_)
| Atom::VlanId(_)
| Atom::EtherType(_)
)
}
fn eval(&self, src: &dyn FieldSource) -> bool {
match self {
Atom::Proto(p) => src.l4proto() == Some(*p),
Atom::SrcPort(p) => src.src_port() == Some(*p),
Atom::DstPort(p) => src.dst_port() == Some(*p),
Atom::AnyPort(p) => src.src_port() == Some(*p) || src.dst_port() == Some(*p),
Atom::SrcHost(h) => src.src_ip() == Some(*h),
Atom::DstHost(h) => src.dst_ip() == Some(*h),
Atom::AnyHost(h) => src.src_ip() == Some(*h) || src.dst_ip() == Some(*h),
Atom::SrcNet(n) => src.src_ip().is_some_and(|ip| n.contains(&ip)),
Atom::DstNet(n) => src.dst_ip().is_some_and(|ip| n.contains(&ip)),
Atom::AnyNet(n) => {
src.src_ip().is_some_and(|ip| n.contains(&ip))
|| src.dst_ip().is_some_and(|ip| n.contains(&ip))
}
Atom::VlanId(v) => src.vlan_id() == Some(*v),
Atom::EtherType(t) => src.ethertype() == Some(*t),
Atom::SniGlob(g) => src.sni().is_some_and(|s| g.matches(s)),
Atom::HttpHostGlob(g) => src.http_host().is_some_and(|s| g.matches(s)),
Atom::DnsQnameGlob(g) => src.dns_qname().is_some_and(|s| g.matches(s)),
Atom::BytesOver(n) => src.total_bytes().is_some_and(|b| b > *n),
Atom::PacketsOver(n) => src.total_packets().is_some_and(|p| p > *n),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Glob {
pattern: String,
}
impl Glob {
pub fn new(pattern: impl Into<String>) -> Self {
Self {
pattern: pattern.into().to_ascii_lowercase(),
}
}
pub fn matches(&self, input: &str) -> bool {
glob_match(&self.pattern, &input.to_ascii_lowercase())
}
}
fn glob_match(pat: &str, text: &str) -> bool {
let p: Vec<u8> = pat.bytes().collect();
let t: Vec<u8> = text.bytes().collect();
let (mut pi, mut ti) = (0usize, 0usize);
let (mut star, mut mark) = (usize::MAX, 0usize);
while ti < t.len() {
if pi < p.len() && p[pi] == b'*' {
star = pi;
mark = ti;
pi += 1;
} else if pi < p.len() && p[pi] == t[ti] {
pi += 1;
ti += 1;
} else if star != usize::MAX {
pi = star + 1;
mark += 1;
ti = mark;
} else {
return false;
}
}
while pi < p.len() && p[pi] == b'*' {
pi += 1;
}
pi == p.len()
}
#[allow(unused_variables)]
pub trait FieldSource {
fn l4proto(&self) -> Option<L4Proto> {
None
}
fn src_port(&self) -> Option<u16> {
None
}
fn dst_port(&self) -> Option<u16> {
None
}
fn src_ip(&self) -> Option<IpAddr> {
None
}
fn dst_ip(&self) -> Option<IpAddr> {
None
}
fn vlan_id(&self) -> Option<u16> {
None
}
fn ethertype(&self) -> Option<u16> {
match self.src_ip()? {
IpAddr::V4(_) => Some(0x0800),
IpAddr::V6(_) => Some(0x86dd),
}
}
fn sni(&self) -> Option<&str> {
None
}
fn http_host(&self) -> Option<&str> {
None
}
fn dns_qname(&self) -> Option<&str> {
None
}
fn total_bytes(&self) -> Option<u64> {
None
}
fn total_packets(&self) -> Option<u64> {
None
}
}
#[cfg(test)]
mod tests {
use std::net::{IpAddr, Ipv4Addr};
use super::*;
#[derive(Default)]
struct Fields {
proto: Option<L4Proto>,
src_port: Option<u16>,
dst_port: Option<u16>,
src_ip: Option<IpAddr>,
dst_ip: Option<IpAddr>,
sni: Option<String>,
bytes: Option<u64>,
}
impl FieldSource for Fields {
fn l4proto(&self) -> Option<L4Proto> {
self.proto
}
fn src_port(&self) -> Option<u16> {
self.src_port
}
fn dst_port(&self) -> Option<u16> {
self.dst_port
}
fn src_ip(&self) -> Option<IpAddr> {
self.src_ip
}
fn dst_ip(&self) -> Option<IpAddr> {
self.dst_ip
}
fn sni(&self) -> Option<&str> {
self.sni.as_deref()
}
fn total_bytes(&self) -> Option<u64> {
self.bytes
}
}
fn tcp_443() -> Fields {
Fields {
proto: Some(L4Proto::Tcp),
src_port: Some(54321),
dst_port: Some(443),
src_ip: Some(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1))),
dst_ip: Some(IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34))),
..Default::default()
}
}
#[test]
fn always_matches_everything() {
assert!(Predicate::Always.eval(&Fields::default()));
}
#[test]
fn ethertype_derives_from_ip_version() {
assert!(Predicate::Atom(Atom::EtherType(0x0800)).eval(&tcp_443()));
assert!(!Predicate::Atom(Atom::EtherType(0x86dd)).eval(&tcp_443()));
assert!(!Predicate::Atom(Atom::EtherType(0x0806)).eval(&Fields::default()));
assert!(Atom::EtherType(0x0806).is_kernel_pushable());
}
#[test]
fn and_or_not_boolean_semantics() {
let f = tcp_443();
let tcp = Predicate::Atom(Atom::Proto(L4Proto::Tcp));
let p443 = Predicate::Atom(Atom::DstPort(443));
let p80 = Predicate::Atom(Atom::DstPort(80));
assert!(tcp.clone().and(p443.clone()).eval(&f));
assert!(!tcp.clone().and(p80.clone()).eval(&f));
assert!(p443.clone().or(p80.clone()).eval(&f));
assert!(p80.clone().negate().eval(&f));
assert!(!p443.clone().negate().eval(&f));
}
#[test]
fn always_is_and_identity_and_or_absorbing() {
let tcp = Predicate::Atom(Atom::Proto(L4Proto::Tcp));
assert_eq!(Predicate::Always.and(tcp.clone()), tcp);
assert_eq!(tcp.clone().and(Predicate::Always), tcp);
assert_eq!(Predicate::Always.or(tcp.clone()), Predicate::Always);
assert_eq!(tcp.or(Predicate::Always), Predicate::Always);
}
#[test]
fn absent_field_atom_does_not_match() {
let g = Predicate::Atom(Atom::SniGlob(Glob::new("*.bank")));
assert!(!g.eval(&tcp_443()));
let b = Predicate::Atom(Atom::BytesOver(1000));
assert!(!b.eval(&tcp_443()));
}
#[test]
fn net_and_count_atoms() {
let f = Fields {
src_ip: Some(IpAddr::V4(Ipv4Addr::new(10, 1, 2, 3))),
bytes: Some(2000),
..Default::default()
};
let net: IpNet = "10.1.0.0/16".parse().unwrap();
assert!(Predicate::Atom(Atom::SrcNet(net)).eval(&f));
let other: IpNet = "192.168.0.0/16".parse().unwrap();
assert!(!Predicate::Atom(Atom::SrcNet(other)).eval(&f));
assert!(Predicate::Atom(Atom::BytesOver(1999)).eval(&f));
assert!(!Predicate::Atom(Atom::BytesOver(2000)).eval(&f)); }
#[test]
fn sni_glob_matches() {
let f = Fields {
sni: Some("login.bank.example".into()),
..Default::default()
};
assert!(Predicate::Atom(Atom::SniGlob(Glob::new("*.bank.example"))).eval(&f));
assert!(Predicate::Atom(Atom::SniGlob(Glob::new("*.BANK.*"))).eval(&f)); assert!(!Predicate::Atom(Atom::SniGlob(Glob::new("*.gov"))).eval(&f));
}
#[test]
fn glob_edge_cases() {
assert!(Glob::new("*").matches("anything"));
assert!(Glob::new("*").matches(""));
assert!(Glob::new("abc").matches("abc"));
assert!(!Glob::new("abc").matches("abcd"));
assert!(Glob::new("a*c").matches("axxxc"));
assert!(Glob::new("a*c").matches("ac"));
assert!(!Glob::new("a*c").matches("ab"));
assert!(Glob::new("*.bank").matches("x.bank"));
assert!(!Glob::new("*.bank").matches("bank"));
assert!(Glob::new("api.*").matches("api.example.com"));
}
#[test]
fn kernel_approx_drops_userspace_atoms_to_always() {
let p = Predicate::Atom(Atom::Proto(L4Proto::Tcp))
.and(Predicate::Atom(Atom::DstPort(443)))
.and(Predicate::Atom(Atom::SniGlob(Glob::new("*.bank"))));
let k = p.kernel_approx();
let expected =
Predicate::Atom(Atom::Proto(L4Proto::Tcp)).and(Predicate::Atom(Atom::DstPort(443)));
assert_eq!(k, expected);
assert!(!p.is_fully_kernel_pushable());
assert!(k.is_fully_kernel_pushable());
}
#[test]
fn kernel_approx_or_with_userspace_branch_is_always() {
let p = Predicate::Atom(Atom::DstPort(443)).or(Predicate::Atom(Atom::BytesOver(1 << 20)));
assert_eq!(p.kernel_approx(), Predicate::Always);
}
#[test]
fn kernel_approx_not_only_pushed_when_fully_kernel() {
let p = Predicate::Atom(Atom::Proto(L4Proto::Tcp)).negate();
assert_eq!(
p.kernel_approx(),
Predicate::Not(Box::new(Predicate::Atom(Atom::Proto(L4Proto::Tcp))))
);
let q = Predicate::Atom(Atom::SniGlob(Glob::new("*.bank"))).negate();
assert_eq!(q.kernel_approx(), Predicate::Always);
}
#[test]
fn kernel_approx_is_a_conservative_superset() {
let preds = [
Predicate::Atom(Atom::Proto(L4Proto::Tcp)).and(Predicate::Atom(Atom::DstPort(443))),
Predicate::Atom(Atom::DstPort(443)).or(Predicate::Atom(Atom::BytesOver(10))),
Predicate::Atom(Atom::Proto(L4Proto::Udp))
.and(Predicate::Atom(Atom::SniGlob(Glob::new("*.x"))).negate()),
Predicate::Atom(Atom::Proto(L4Proto::Tcp)).negate(),
];
let sources = [
Fields {
proto: Some(L4Proto::Tcp),
dst_port: Some(443),
sni: Some("a.bank".into()),
bytes: Some(5),
..Default::default()
},
Fields {
proto: Some(L4Proto::Udp),
dst_port: Some(53),
bytes: Some(100),
..Default::default()
},
Fields {
proto: Some(L4Proto::Tcp),
dst_port: Some(80),
..Default::default()
},
Fields::default(),
];
for p in &preds {
let k = p.kernel_approx();
for f in &sources {
if p.eval(f) {
assert!(
k.eval(f),
"superset violated: p matched but kernel_approx didn't\n p={p:?}\n k={k:?}"
);
}
}
}
}
#[test]
fn kernel_pushability_classification() {
assert!(Atom::Proto(L4Proto::Tcp).is_kernel_pushable());
assert!(Atom::DstPort(443).is_kernel_pushable());
assert!(Atom::AnyNet("10.0.0.0/8".parse().unwrap()).is_kernel_pushable());
assert!(Atom::VlanId(100).is_kernel_pushable());
assert!(!Atom::SniGlob(Glob::new("*.bank")).is_kernel_pushable());
assert!(!Atom::BytesOver(1).is_kernel_pushable());
assert!(!Atom::PacketsOver(1).is_kernel_pushable());
}
}