From 55cb583e618643fabf494cdce7911740e8aea3f3 Mon Sep 17 00:00:00 2001 From: Dirkjan Ochtman Date: Thu, 3 Oct 2024 13:02:18 +0200 Subject: [PATCH] spf: refactor EvalContext --- crates/kumo-spf/src/eval.rs | 235 +++++++++++++++++----------------- crates/kumo-spf/src/lib.rs | 16 +-- crates/kumo-spf/src/record.rs | 4 +- 3 files changed, 123 insertions(+), 132 deletions(-) diff --git a/crates/kumo-spf/src/eval.rs b/crates/kumo-spf/src/eval.rs index e73b4c03..42fee92e 100644 --- a/crates/kumo-spf/src/eval.rs +++ b/crates/kumo-spf/src/eval.rs @@ -1,139 +1,133 @@ use crate::record::{MacroElement, MacroName}; -use std::collections::HashMap; +use std::fmt::Write; use std::net::IpAddr; use std::time::SystemTime; -#[derive(Default)] -pub struct EvalContext { - vars: HashMap, +pub struct EvalContext<'a> { + sender: &'a str, + local_part: &'a str, + sender_domain: &'a str, + domain: &'a str, + client_ip: IpAddr, + now: SystemTime, } -impl EvalContext { - pub fn new() -> Self { - let mut ctx = Self::default(); +impl<'a> EvalContext<'a> { + pub fn new(sender: &'a str, domain: &'a str, client_ip: IpAddr) -> Result { + let Some((local_part, sender_domain)) = sender.split_once('@') else { + return Err(format!( + "invalid sender {sender} is missing @ sign to delimit local part and domain" + )); + }; - ctx.set_var( - MacroName::CurrentUnixTimeStamp, - SystemTime::now() - .duration_since(SystemTime::UNIX_EPOCH) - .map(|d| d.as_secs()) - .unwrap_or(0), - ); - - ctx - } - - pub fn set_sender(&mut self, sender: V) -> Result<(), String> { - let sender = sender.to_string(); - let (local, domain) = sender.split_once('@').ok_or_else(|| { - format!("invalid sender {sender} is missing @ sign to delimit local part and domain") - })?; - - self.set_var(MacroName::LocalPart, local); - self.set_var(MacroName::SenderDomain, domain); - self.set_var(MacroName::Sender, sender); - - Ok(()) - } - - pub fn set_ip>(&mut self, ip: IP) { - let ip: IpAddr = ip.into(); - self.set_var(MacroName::ClientIp, ip); - - match ip { - IpAddr::V4(v4) => { - self.set_var(MacroName::Ip, v4); - } - IpAddr::V6(v6) => { - // For IPv6 addresses, the "i" macro expands to a dot-format address; - // it is intended for use in %{ir}. - let mut ip = String::new(); - for segment in v6.segments() { - for b in format!("{segment:04x}").chars() { - if !ip.is_empty() { - ip.push('.'); - } - ip.push(b); - } - } - self.set_var(MacroName::Ip, ip); - } - } - self.set_var( - MacroName::ReverseDns, - if ip.is_ipv4() { "in-addr" } else { "ip6" }, - ); - } - - pub fn set_client_ip>(&mut self, ip: IP) { - let ip: IpAddr = ip.into(); - self.set_var(MacroName::ClientIp, ip); - } - - pub fn set_var(&mut self, name: MacroName, value: V) { - self.vars.insert(name, value.to_string()); + Ok(Self { + sender, + local_part, + sender_domain, + domain, + client_ip, + now: SystemTime::now(), + }) } pub fn evaluate(&self, elements: &[MacroElement]) -> Result { - let mut result = String::new(); + let (mut result, mut buf) = (String::new(), String::new()); for element in elements { - match element { + let m = match element { MacroElement::Literal(t) => { result.push_str(&t); + continue; } - MacroElement::Macro(m) => { - eprintln!("apply {m:?}"); - let value = self - .vars - .get(&m.name) - .ok_or_else(|| format!("{:?} has no been set in EvalContext", m.name))?; - let delimiters = if m.delimiters.is_empty() { - "." - } else { - &m.delimiters - }; - let mut tokens: Vec<&str> = value.split(|c| delimiters.contains(c)).collect(); + MacroElement::Macro(m) => m, + }; - if m.reverse { - tokens.reverse(); + buf.clear(); + match m.name { + MacroName::Sender => buf.push_str(self.sender), + MacroName::LocalPart => buf.push_str(&self.local_part), + MacroName::SenderDomain => buf.push_str(&self.sender_domain), + MacroName::Domain => buf.push_str(&self.domain), + MacroName::ReverseDns => buf.push_str(match self.client_ip.is_ipv4() { + true => "in-addr", + false => "ip6", + }), + MacroName::ClientIp => { + buf.write_fmt(format_args!("{}", self.client_ip)).unwrap(); + } + MacroName::Ip => match self.client_ip { + IpAddr::V4(v4) => { + buf.write_fmt(format_args!("{}", v4)).unwrap(); } - - if let Some(n) = m.transformer_digits { - let n = n as usize; - while tokens.len() > n { - tokens.remove(0); - } - } - - let output = tokens.join("."); - - if m.url_escape { - // https://datatracker.ietf.org/doc/html/rfc7208#section-7.3: - // Uppercase macros expand exactly as their lowercase - // equivalents, and are then URL escaped. URL escaping - // MUST be performed for characters not in the - // "unreserved" set. - // https://datatracker.ietf.org/doc/html/rfc3986#section-2.3: - // unreserved = ALPHA / DIGIT / "-" / "." / "_" / "~" - for c in output.chars() { - if c.is_ascii_alphanumeric() - || c == '-' - || c == '.' - || c == '_' - || c == '~' - { - result.push(c); - } else { - let mut bytes = [0u8; 4]; - for b in c.encode_utf8(&mut bytes).bytes() { - result.push_str(&format!("%{b:02x}")); + IpAddr::V6(v6) => { + // For IPv6 addresses, the "i" macro expands to a dot-format address; + // it is intended for use in %{ir}. + for segment in v6.segments() { + for b in format!("{segment:04x}").chars() { + if !buf.is_empty() { + buf.push('.'); } + buf.push(b); } } + } + }, + MacroName::CurrentUnixTimeStamp => buf + .write_fmt(format_args!( + "{}", + self.now + .duration_since(SystemTime::UNIX_EPOCH) + .map(|d| d.as_secs()) + .unwrap_or(0) + )) + .unwrap(), + MacroName::RelayingHostName + | MacroName::HeloDomain + | MacroName::ValidatedDomainName => { + return Err(format!("{:?} has not been implemented", m.name)) + } + }; + + let delimiters = if m.delimiters.is_empty() { + "." + } else { + &m.delimiters + }; + + let mut tokens: Vec<&str> = buf.split(|c| delimiters.contains(c)).collect(); + + if m.reverse { + tokens.reverse(); + } + + if let Some(n) = m.transformer_digits { + let n = n as usize; + while tokens.len() > n { + tokens.remove(0); + } + } + + let output = tokens.join("."); + + if m.url_escape { + // https://datatracker.ietf.org/doc/html/rfc7208#section-7.3: + // Uppercase macros expand exactly as their lowercase + // equivalents, and are then URL escaped. URL escaping + // MUST be performed for characters not in the + // "unreserved" set. + // https://datatracker.ietf.org/doc/html/rfc3986#section-2.3: + // unreserved = ALPHA / DIGIT / "-" / "." / "_" / "~" + for c in output.chars() { + if c.is_ascii_alphanumeric() || c == '-' || c == '.' || c == '_' || c == '~' { + result.push(c); } else { - result.push_str(&output); + let mut bytes = [0u8; 4]; + for b in c.encode_utf8(&mut bytes).bytes() { + result.push_str(&format!("%{b:02x}")); + } } } + } else { + result.push_str(&output); } } @@ -150,9 +144,12 @@ mod test { fn test_eval() { // - let mut ctx = EvalContext::new(); - ctx.set_sender("strong-bad@email.example.com").unwrap(); - ctx.set_var(MacroName::Domain, "email.example.com"); + let mut ctx = EvalContext::new( + "strong-bad@email.example.com", + "email.example.com", + IpAddr::from([192, 0, 2, 3]), + ) + .unwrap(); for (input, expect) in &[ ("%{s}", "strong-bad@email.example.com"), @@ -175,8 +172,6 @@ mod test { k9::assert_equal!(&output, expect, "{input}"); } - ctx.set_ip("192.0.2.3".parse::().unwrap()); - for (input, expect) in &[ ( "%{ir}.%{v}._spf.%{d2}", @@ -202,7 +197,7 @@ mod test { k9::assert_equal!(&output, expect, "{input}"); } - ctx.set_ip("2001:db8::cb01".parse::().unwrap()); + ctx.client_ip = IpAddr::from([0x2001, 0xdb8, 0, 0, 0, 0, 0, 0xcb01]); for (input, expect) in &[ ( "%{ir}.%{v}._spf.%{d2}", diff --git a/crates/kumo-spf/src/lib.rs b/crates/kumo-spf/src/lib.rs index fec3ce60..18ad86c1 100644 --- a/crates/kumo-spf/src/lib.rs +++ b/crates/kumo-spf/src/lib.rs @@ -6,7 +6,7 @@ use error::SpfError; pub mod eval; use eval::EvalContext; pub mod record; -use record::{MacroName, Qualifier, Record}; +use record::{Qualifier, Record}; #[cfg(test)] mod tests; @@ -105,19 +105,15 @@ impl CheckHostParams { } }; - let mut cx = EvalContext::new(); - cx.set_ip(self.client_ip); - if let Err(err) = cx.set_sender(&self.sender) { - return SpfResult { - disposition: SpfDisposition::TempError, + match EvalContext::new(&self.sender, &self.domain, self.client_ip) { + Ok(cx) => record.evaluate(&cx, resolver).await, + Err(err) => SpfResult { + disposition: SpfDisposition::PermError, context: format!( "input sender parameter '{}' is malformed: {err}", self.sender ), - }; + }, } - cx.set_var(MacroName::Domain, &self.domain); - - record.evaluate(&cx, resolver).await } } diff --git a/crates/kumo-spf/src/record.rs b/crates/kumo-spf/src/record.rs index d9ece56f..006db24d 100644 --- a/crates/kumo-spf/src/record.rs +++ b/crates/kumo-spf/src/record.rs @@ -33,7 +33,7 @@ impl Record { Ok(Self { terms }) } - pub async fn evaluate(&self, cx: &EvalContext, resolver: &dyn Lookup) -> SpfResult { + pub async fn evaluate(&self, cx: &EvalContext<'_>, resolver: &dyn Lookup) -> SpfResult { for term in &self.terms { match term { Term::Directive(d) => match d.evaluate(cx, resolver).await { @@ -85,7 +85,7 @@ impl Directive { pub async fn evaluate( &self, - _cx: &EvalContext, + _cx: &EvalContext<'_>, _resolver: &dyn Lookup, ) -> Result, SpfResult> { let matched = match &self.mechanism {