use base64::Engine as _;
use rsa::pkcs8::DecodePublicKey;
use rsa::Pkcs1v15Sign;
use sha2::{Digest, Sha256};
use crate::canon::{canonicalize_body, canonicalize_header};
use crate::error::{DkimError, DkimResult};
use crate::header::{Algorithm, DkimHeader};
use crate::resolver::DkimResolver;
pub async fn verify<R: DkimResolver + ?Sized>(resolver: &R, raw_message: &[u8]) -> DkimResult {
match verify_inner(resolver, raw_message).await {
Ok(r) => r,
Err(e) => e.to_result(),
}
}
async fn verify_inner<R: DkimResolver + ?Sized>(
resolver: &R,
raw_message: &[u8],
) -> Result<DkimResult, DkimError> {
let (header_value, signed_headers_raw, body_offset) = extract_dkim_signature(raw_message)?;
let header = DkimHeader::parse(&header_value)?;
if let Some(x) = header.expiration {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
if now > x {
return Err(DkimError::Expired);
}
}
let body = &raw_message[body_offset..];
let canon_body_bytes = canonicalize_body(body, header.canon_body, header.body_length);
let mut body_hasher = Sha256::new();
body_hasher.update(&canon_body_bytes);
let actual_body_hash = body_hasher.finalize();
let expected_body_hash = base64::engine::general_purpose::STANDARD
.decode(&header.body_hash_b64)
.map_err(|_| DkimError::InvalidBase64("bh".into()))?;
if actual_body_hash.as_slice() != expected_body_hash.as_slice() {
return Err(DkimError::BodyHashMismatch);
}
let pubkey_domain = format!("{}._domainkey.{}", header.selector, header.domain);
let txts = resolver.lookup_txt(&pubkey_domain).await?;
if txts.is_empty() {
return Err(DkimError::DnsPermError(format!(
"no TXT at {pubkey_domain}"
)));
}
let key_txt = txts
.iter()
.find(|s| s.contains("p="))
.ok_or_else(|| DkimError::InvalidKey("no p= tag in TXT".into()))?;
let key_bytes = extract_public_key(key_txt)?;
let mut signed_block = Vec::new();
for name in &header.signed_headers {
if let Some(value) = find_header_value(signed_headers_raw, name) {
let canon = canonicalize_header(name, value, header.canon_header);
signed_block.extend_from_slice(&canon);
}
}
let dkim_sig_b_cleared = clear_b_value(&header_value);
let canon_dkim = canonicalize_header(
"DKIM-Signature",
&dkim_sig_b_cleared,
header.canon_header,
);
let canon_dkim_trimmed = if canon_dkim.ends_with(b"\r\n") {
&canon_dkim[..canon_dkim.len() - 2]
} else {
&canon_dkim
};
signed_block.extend_from_slice(canon_dkim_trimmed);
let signature_bytes = base64::engine::general_purpose::STANDARD
.decode(&header.signature_b64)
.map_err(|_| DkimError::InvalidBase64("b".into()))?;
match header.algorithm {
Algorithm::RsaSha256 => {
let public_key = rsa::RsaPublicKey::from_public_key_der(&key_bytes).map_err(|e| {
DkimError::InvalidKey(format!("RSA PKCS8 decode failed: {e}"))
})?;
let mut hasher = Sha256::new();
hasher.update(&signed_block);
let digest = hasher.finalize();
let scheme = Pkcs1v15Sign::new::<Sha256>();
public_key
.verify(scheme, &digest, &signature_bytes)
.map_err(|_| DkimError::SignatureMismatch)?;
}
Algorithm::Ed25519Sha256 => {
if key_bytes.len() != 32 {
return Err(DkimError::InvalidKey(format!(
"ed25519 key wrong length: {} (expected 32)",
key_bytes.len()
)));
}
let mut key_arr = [0u8; 32];
key_arr.copy_from_slice(&key_bytes);
let verifying_key = ed25519_dalek::VerifyingKey::from_bytes(&key_arr).map_err(
|e| DkimError::InvalidKey(format!("ed25519 key decode: {e}")),
)?;
let mut hasher = Sha256::new();
hasher.update(&signed_block);
let digest = hasher.finalize();
if signature_bytes.len() != 64 {
return Err(DkimError::InvalidBase64(format!(
"b= ed25519 sig wrong length: {}",
signature_bytes.len()
)));
}
let mut sig_arr = [0u8; 64];
sig_arr.copy_from_slice(&signature_bytes);
let signature = ed25519_dalek::Signature::from_bytes(&sig_arr);
use ed25519_dalek::Verifier as _;
verifying_key
.verify(&digest, &signature)
.map_err(|_| DkimError::SignatureMismatch)?;
}
}
Ok(DkimResult::Pass)
}
fn extract_dkim_signature(raw: &[u8]) -> Result<(String, &[u8], usize), DkimError> {
let body_offset = find_body_offset(raw).ok_or(DkimError::MissingHeader)?;
let headers_raw = &raw[..body_offset_minus_blank(body_offset, raw)];
let value = find_header_value_in_raw(headers_raw, b"DKIM-Signature")?;
Ok((value, headers_raw, body_offset))
}
fn body_offset_minus_blank(body_offset: usize, raw: &[u8]) -> usize {
if body_offset >= 2 && &raw[body_offset - 2..body_offset] == b"\r\n" {
body_offset - 2
} else if body_offset >= 1 && raw[body_offset - 1] == b'\n' {
body_offset - 1
} else {
body_offset
}
}
fn find_body_offset(raw: &[u8]) -> Option<usize> {
let mut i = 0;
while i < raw.len() {
if i + 3 < raw.len() && &raw[i..i + 4] == b"\r\n\r\n" {
return Some(i + 4);
}
if i + 1 < raw.len() && raw[i] == b'\n' && raw[i + 1] == b'\n' {
return Some(i + 2);
}
i += 1;
}
None
}
fn find_header_value_in_raw(headers: &[u8], name: &[u8]) -> Result<String, DkimError> {
let mut i = 0;
while i < headers.len() {
if i + name.len() < headers.len()
&& headers[i..i + name.len()].eq_ignore_ascii_case(name)
&& headers[i + name.len()] == b':'
{
let value_start = i + name.len() + 1;
let mut j = value_start;
while j < headers.len() {
if headers[j] == b'\n' {
let after = j + 1;
if after < headers.len() && matches!(headers[after], b' ' | b'\t') {
j += 1;
continue;
}
return Ok(String::from_utf8_lossy(&headers[value_start..j]).into_owned());
}
j += 1;
}
return Ok(String::from_utf8_lossy(&headers[value_start..j]).into_owned());
}
while i < headers.len() && headers[i] != b'\n' {
i += 1;
}
i += 1;
}
Err(DkimError::MissingHeader)
}
fn find_header_value<'a>(headers: &'a [u8], name: &str) -> Option<&'a str> {
let bytes = name.as_bytes();
let mut i = 0;
while i < headers.len() {
if i + bytes.len() < headers.len()
&& headers[i..i + bytes.len()].eq_ignore_ascii_case(bytes)
&& headers[i + bytes.len()] == b':'
{
let value_start = i + bytes.len() + 1;
let mut j = value_start;
while j < headers.len() {
if headers[j] == b'\n' {
let after = j + 1;
if after < headers.len() && matches!(headers[after], b' ' | b'\t') {
j += 1;
continue;
}
let end = if j > value_start && headers[j - 1] == b'\r' {
j - 1
} else {
j
};
return std::str::from_utf8(&headers[value_start..end]).ok();
}
j += 1;
}
return std::str::from_utf8(&headers[value_start..j]).ok();
}
while i < headers.len() && headers[i] != b'\n' {
i += 1;
}
i += 1;
}
None
}
fn clear_b_value(value: &str) -> String {
let bytes = value.as_bytes();
let mut out = Vec::with_capacity(bytes.len());
let mut i = 0;
while i < bytes.len() {
let is_b_start = if i + 1 < bytes.len() && bytes[i] == b'b' && bytes[i + 1] == b'=' {
let mut k = i;
while k > 0 {
k -= 1;
if !matches!(bytes[k], b' ' | b'\t' | b'\r' | b'\n') {
break;
}
}
k == 0 || bytes[k] == b';'
} else {
false
};
if is_b_start {
out.extend_from_slice(b"b=");
i += 2;
while i < bytes.len() && bytes[i] != b';' {
i += 1;
}
continue;
}
out.push(bytes[i]);
i += 1;
}
String::from_utf8_lossy(&out).into_owned()
}
fn extract_public_key(txt: &str) -> Result<Vec<u8>, DkimError> {
let p_value = txt
.split(';')
.find_map(|t| {
let t = t.trim();
t.strip_prefix("p=")
})
.ok_or_else(|| DkimError::InvalidKey("p= tag missing".into()))?;
let p_value = p_value
.chars()
.filter(|c| !matches!(c, ' ' | '\t' | '\r' | '\n'))
.collect::<String>();
if p_value.is_empty() {
return Err(DkimError::InvalidKey("p= empty (key revoked)".into()));
}
base64::engine::general_purpose::STANDARD
.decode(p_value.as_bytes())
.map_err(|e| DkimError::InvalidKey(format!("p= base64 decode: {e}")))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn body_offset_simple_crlf_crlf() {
let raw = b"From: a\r\nTo: b\r\n\r\nhello";
let off = find_body_offset(raw).unwrap();
assert_eq!(&raw[off..], b"hello");
}
#[test]
fn body_offset_lf_lf() {
let raw = b"From: a\nTo: b\n\nhello";
let off = find_body_offset(raw).unwrap();
assert_eq!(&raw[off..], b"hello");
}
#[test]
fn body_offset_no_blank_line() {
let raw = b"From: a\r\nTo: b\r\n";
assert!(find_body_offset(raw).is_none());
}
#[test]
fn find_header_extracts_value() {
let raw = b"From: alice@e.com\r\nTo: bob@e.com\r\n";
assert_eq!(find_header_value(raw, "From"), Some(" alice@e.com"));
assert_eq!(find_header_value(raw, "to"), Some(" bob@e.com"));
}
#[test]
fn find_header_handles_folded() {
let raw = b"X-Long: line1\r\n line2\r\nFrom: a\r\n";
let val = find_header_value(raw, "X-Long").unwrap();
assert!(val.contains("line1"));
assert!(val.contains("line2"));
}
#[test]
fn find_header_returns_none_if_absent() {
let raw = b"From: a\r\n";
assert!(find_header_value(raw, "Missing").is_none());
}
#[test]
fn clear_b_replaces_value_only() {
let v = " v=1; a=rsa-sha256; b=ABCDEFG; d=e.com";
let cleared = clear_b_value(v);
assert!(cleared.contains("b=;") || cleared.contains("b= "));
assert!(!cleared.contains("ABCDEFG"));
assert!(cleared.contains("v=1"));
assert!(cleared.contains("a=rsa-sha256"));
assert!(cleared.contains("d=e.com"));
}
#[test]
fn extract_pubkey_finds_p_tag() {
let txt = "v=DKIM1; k=rsa; p=AA==";
let der = extract_public_key(txt).unwrap();
assert!(!der.is_empty());
}
#[test]
fn extract_pubkey_rejects_missing_p() {
let txt = "v=DKIM1; k=rsa";
let r = extract_public_key(txt);
assert!(matches!(r, Err(DkimError::InvalidKey(_))));
}
#[test]
fn extract_pubkey_rejects_empty_p() {
let txt = "v=DKIM1; k=rsa; p=";
let r = extract_public_key(txt);
assert!(matches!(r, Err(DkimError::InvalidKey(_))));
}
#[test]
fn extract_pubkey_strips_wsp_in_p() {
let txt = "v=DKIM1; k=rsa; p=AA == ";
let der = extract_public_key(txt).unwrap();
assert!(!der.is_empty());
}
}