use std::fmt;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Cidr {
bits: u128,
prefix: u8,
is_v6: bool,
}
impl Cidr {
pub fn new(addr: IpAddr, prefix: u8) -> Option<Self> {
match addr {
IpAddr::V4(v4) => {
if prefix > 32 {
return None;
}
Some(Cidr {
bits: u32::from(v4) as u128,
prefix,
is_v6: false,
})
}
IpAddr::V6(v6) => {
if prefix > 128 {
return None;
}
Some(Cidr {
bits: u128::from(v6),
prefix,
is_v6: true,
})
}
}
}
pub fn contains(&self, addr: IpAddr) -> bool {
let (value, is_v6) = match addr {
IpAddr::V4(v4) => (u32::from(v4) as u128, false),
IpAddr::V6(v6) => (u128::from(v6), true),
};
if is_v6 != self.is_v6 {
return false;
}
let total_bits = if self.is_v6 { 128 } else { 32 };
if self.prefix == 0 {
return true;
}
let shift = total_bits - self.prefix as u32;
(value >> shift) == (self.bits >> shift)
}
pub fn prefix(&self) -> u8 {
self.prefix
}
pub fn is_v6(&self) -> bool {
self.is_v6
}
}
impl std::str::FromStr for Cidr {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let s = s.trim();
if let Some((addr_s, prefix_s)) = s.split_once('/') {
let addr: IpAddr = addr_s
.trim()
.parse()
.map_err(|_| format!("invalid IP address '{addr_s}'"))?;
let prefix: u8 = prefix_s
.trim()
.parse()
.map_err(|_| format!("invalid prefix '{prefix_s}'"))?;
Cidr::new(addr, prefix).ok_or_else(|| format!("prefix /{prefix} too large for address"))
} else {
let addr: IpAddr = s.parse().map_err(|_| format!("invalid IP address '{s}'"))?;
let prefix = match addr {
IpAddr::V4(_) => 32,
IpAddr::V6(_) => 128,
};
Ok(Cidr::new(addr, prefix).expect("host prefix is always valid"))
}
}
}
impl fmt::Display for Cidr {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let addr = if self.is_v6 {
IpAddr::V6(Ipv6Addr::from(self.bits))
} else {
IpAddr::V4(Ipv4Addr::from(self.bits as u32))
};
write!(f, "{addr}/{}", self.prefix)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub struct PortRange {
pub min: u16,
pub max: u16,
}
impl PortRange {
pub const ANY: PortRange = PortRange {
min: 0,
max: u16::MAX,
};
pub fn contains(&self, port: u16) -> bool {
port >= self.min && port <= self.max
}
}
impl std::str::FromStr for PortRange {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let s = s.trim();
if let Some((lo, hi)) = s.split_once('-') {
let min: u16 = lo
.trim()
.parse()
.map_err(|_| format!("invalid port '{lo}'"))?;
let max: u16 = hi
.trim()
.parse()
.map_err(|_| format!("invalid port '{hi}'"))?;
if min > max {
return Err(format!("port range min ({min}) > max ({max})"));
}
Ok(PortRange { min, max })
} else {
let p: u16 = s.parse().map_err(|_| format!("invalid port '{s}'"))?;
Ok(PortRange { min: p, max: p })
}
}
}
pub(crate) fn validate_hostname(name: &str) -> Result<(), String> {
if name.chars().any(|c| c.is_control() || c.is_whitespace()) {
return Err("contains a control or whitespace character".into());
}
let stem = name.strip_suffix('.').unwrap_or(name);
if stem.is_empty() {
return Err("has no labels".into());
}
if stem.len() > 253 {
return Err("exceeds the maximum DNS name length of 253 bytes".into());
}
for label in stem.split('.') {
if label.is_empty() {
return Err("has an empty label".into());
}
if label.len() > 63 {
return Err("has a label longer than 63 bytes".into());
}
}
Ok(())
}
pub(crate) fn validate_acme_domain(name: &str) -> Result<(), String> {
validate_hostname(name)?;
let stem = name.strip_suffix('.').unwrap_or(name);
if stem.parse::<IpAddr>().is_ok() {
return Err("must be a DNS name, not an IP address".into());
}
if !stem.contains('.') {
return Err(
"must be a multi-label DNS name (a public CA cannot issue for a single-label or \
local name)"
.into(),
);
}
for label in stem.split('.') {
if !label
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'-')
{
return Err(
"has a non-LDH label (no wildcards or underscores; use punycode 'xn--' for IDN)"
.into(),
);
}
if label.starts_with('-') || label.ends_with('-') {
return Err("has a label with a leading or trailing hyphen".into());
}
}
const SPECIAL_USE_TLDS: &[&str] = &[
"local",
"localhost",
"test",
"invalid",
"example",
"internal",
"arpa",
"onion",
"alt",
];
let tld = stem.rsplit('.').next().unwrap_or(stem);
if SPECIAL_USE_TLDS.iter().any(|s| tld.eq_ignore_ascii_case(s)) {
return Err(format!(
"ends in the special-use TLD '.{tld}', which no public CA issues certificates for"
));
}
Ok(())
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum HostPattern {
Exact(String),
Suffix(String),
}
impl HostPattern {
pub fn parse(token: &str) -> Result<Self, String> {
let lower = token.trim().to_ascii_lowercase();
let (mut domain, suffix) = match lower.strip_prefix('.') {
Some(rest) => (rest.to_string(), true),
None => (lower, false),
};
if domain.ends_with('.') && !domain.ends_with("..") {
domain.pop();
}
let labels_valid = domain.split('.').all(|label| {
(1..=63).contains(&label.len())
&& !label.starts_with('-')
&& !label.ends_with('-')
&& label
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'-')
});
if !labels_valid {
return Err(format!("'{token}' is not a valid hostname pattern"));
}
Ok(if suffix {
HostPattern::Suffix(domain)
} else {
HostPattern::Exact(domain)
})
}
pub fn matches(&self, host: &str) -> bool {
let host = host.strip_suffix('.').unwrap_or(host);
match self {
HostPattern::Exact(h) => host.eq_ignore_ascii_case(h),
HostPattern::Suffix(domain) => {
host.eq_ignore_ascii_case(domain)
|| (host.len() > domain.len()
&& host.as_bytes()[host.len() - domain.len() - 1] == b'.'
&& host[host.len() - domain.len()..].eq_ignore_ascii_case(domain))
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct AddrSpec {
pub cidrs: Vec<Cidr>,
pub hosts: Vec<HostPattern>,
pub ports: Option<PortRange>,
}
impl AddrSpec {
pub fn any() -> Self {
AddrSpec {
cidrs: vec![
"0.0.0.0/0".parse().expect("constant CIDR is valid"),
"::/0".parse().expect("constant CIDR is valid"),
],
hosts: Vec::new(),
ports: None,
}
}
pub fn new(cidr: Cidr, ports: Option<PortRange>) -> Self {
AddrSpec {
cidrs: vec![cidr],
hosts: Vec::new(),
ports,
}
}
pub fn host(pattern: HostPattern, ports: Option<PortRange>) -> Self {
AddrSpec {
cidrs: Vec::new(),
hosts: vec![pattern],
ports,
}
}
pub fn matches(&self, ip: IpAddr, port: u16) -> bool {
self.cidrs.iter().any(|cidr| cidr.contains(ip)) && self.port_matches(port)
}
pub fn matches_dest(&self, host: Option<&str>, ip: IpAddr, port: u16) -> bool {
self.port_matches(port)
&& (self.cidrs.iter().any(|cidr| cidr.contains(ip))
|| host.is_some_and(|h| self.hosts.iter().any(|p| p.matches(h))))
}
pub fn matches_all(&self) -> bool {
self.ports.is_none()
&& self.cidrs.iter().any(|c| c.prefix() == 0 && !c.is_v6())
&& self.cidrs.iter().any(|c| c.prefix() == 0 && c.is_v6())
}
fn port_matches(&self, port: u16) -> bool {
match self.ports {
Some(range) => range.contains(port),
None => true,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn validate_hostname_rejects_control_characters() {
for bad in [
"ex\rample.com",
"ex\nample.com",
"ex\0ample.com",
"ex\tample.com",
] {
assert!(validate_hostname(bad).is_err(), "{bad:?} must be rejected");
}
}
#[test]
fn validate_hostname_rejects_unicode_controls_and_whitespace() {
for bad in [
"ex\u{85}ample.com",
"ex\u{2028}ample.com",
"ex\u{2029}ample.com",
"ex ample.com",
"ex\u{a0}ample.com",
] {
assert!(validate_hostname(bad).is_err(), "{bad:?} must be rejected");
}
}
#[test]
fn validate_hostname_rejects_malformed_labels() {
assert!(validate_hostname("foo..bar").is_err()); assert!(validate_hostname(".foo").is_err()); assert!(validate_hostname(".").is_err()); assert!(validate_hostname("foo..").is_err()); assert!(validate_hostname("foo.").is_ok()); let oversize = format!("{}.com", "a".repeat(64)); assert!(validate_hostname(&oversize).is_err());
let too_long = vec!["a".repeat(63); 4].join("."); assert!(validate_hostname(&too_long).is_err());
}
#[test]
fn validate_hostname_accepts_normal_and_absolute_names() {
for ok in [
"example.com",
"a.b.c.example.org",
"under_score.example", "xn--bcher-kva.example", "host.example.com.", ] {
assert!(validate_hostname(ok).is_ok(), "{ok:?} must be accepted");
}
}
#[test]
fn validate_acme_domain_rejects_non_issuable_names() {
for bad in [
"*.example.com", "_svc.example.com", "é.example.com", "-bad.example.com", "bad-.example.com", "foo..bar", "localhost", "127.0.0.1", "127.0.0.1.", "::1", "router.local", "app.test", "service.invalid", "host.home.arpa", "x.internal", "h.onion", "site.alt", "DEV.LOCAL", ] {
assert!(
validate_acme_domain(bad).is_err(),
"{bad:?} must be rejected"
);
}
}
#[test]
fn validate_acme_domain_accepts_issuable_names() {
for ok in [
"example.com",
"a.b.example.org",
"host-1.example.com", "xn--bcher-kva.example.com", "example.com.", ] {
assert!(validate_acme_domain(ok).is_ok(), "{ok:?} must be accepted");
}
}
#[test]
fn cidr_v4_contains() {
let net: Cidr = "10.0.0.0/8".parse().unwrap();
assert!(net.contains("10.1.2.3".parse().unwrap()));
assert!(net.contains("10.255.255.255".parse().unwrap()));
assert!(!net.contains("11.0.0.0".parse().unwrap()));
}
#[test]
fn cidr_any_v4() {
let net: Cidr = "0.0.0.0/0".parse().unwrap();
assert!(net.contains("8.8.8.8".parse().unwrap()));
assert!(net.contains("192.168.1.1".parse().unwrap()));
assert!(!net.contains("::1".parse().unwrap()));
}
#[test]
fn cidr_host_route() {
let net: Cidr = "127.0.0.1".parse().unwrap();
assert_eq!(net.prefix(), 32);
assert!(net.contains("127.0.0.1".parse().unwrap()));
assert!(!net.contains("127.0.0.2".parse().unwrap()));
}
#[test]
fn cidr_v6_contains() {
let net: Cidr = "fd00::/8".parse().unwrap();
assert!(net.is_v6());
assert!(net.contains("fd12:3456::1".parse().unwrap()));
assert!(!net.contains("fe80::1".parse().unwrap()));
assert!(!net.contains("10.0.0.1".parse().unwrap()));
}
#[test]
fn cidr_partial_byte_prefix() {
let net: Cidr = "192.168.1.0/25".parse().unwrap();
assert!(net.contains("192.168.1.0".parse().unwrap()));
assert!(net.contains("192.168.1.127".parse().unwrap()));
assert!(!net.contains("192.168.1.128".parse().unwrap()));
}
#[test]
fn cidr_rejects_oversized_prefix() {
assert!("10.0.0.0/33".parse::<Cidr>().is_err());
assert!("::/129".parse::<Cidr>().is_err());
}
#[test]
fn cidr_display_roundtrip() {
let net: Cidr = "172.16.0.0/12".parse().unwrap();
assert_eq!(net.to_string(), "172.16.0.0/12");
}
#[test]
fn port_range_single() {
let r: PortRange = "443".parse().unwrap();
assert_eq!(r, PortRange { min: 443, max: 443 });
assert!(r.contains(443));
assert!(!r.contains(444));
}
#[test]
fn port_range_span() {
let r: PortRange = "1000 - 2000".parse().unwrap();
assert!(r.contains(1000));
assert!(r.contains(1500));
assert!(r.contains(2000));
assert!(!r.contains(999));
assert!(!r.contains(2001));
}
#[test]
fn port_range_inverted_rejected() {
assert!("2000-1000".parse::<PortRange>().is_err());
}
#[test]
fn addr_spec_matches_ip_and_port() {
let spec = AddrSpec {
cidrs: vec!["192.168.0.0/16".parse().unwrap()],
hosts: Vec::new(),
ports: Some(PortRange { min: 80, max: 80 }),
};
assert!(spec.matches("192.168.5.5".parse().unwrap(), 80));
assert!(!spec.matches("192.168.5.5".parse().unwrap(), 81));
assert!(!spec.matches("10.0.0.1".parse().unwrap(), 80));
}
#[test]
fn addr_spec_any_port() {
let spec = AddrSpec {
cidrs: vec!["0.0.0.0/0".parse().unwrap()],
hosts: Vec::new(),
ports: None,
};
assert!(spec.matches("1.2.3.4".parse().unwrap(), 1));
assert!(spec.matches("1.2.3.4".parse().unwrap(), 65535));
}
#[test]
fn addr_spec_any_matches_v4_and_v6() {
let spec = AddrSpec::any();
assert!(spec.matches("8.8.8.8".parse().unwrap(), 53));
assert!(spec.matches("2001:db8::1".parse().unwrap(), 443));
}
#[test]
fn host_pattern_exact_and_suffix_match() {
let exact = HostPattern::parse("Example.com").unwrap();
assert_eq!(exact, HostPattern::Exact("example.com".into()));
assert!(exact.matches("example.com"));
assert!(exact.matches("EXAMPLE.COM"));
assert!(exact.matches("example.com.")); assert!(!exact.matches("example.com..")); assert!(!exact.matches("a.example.com"));
assert!(!exact.matches("notexample.com"));
let suffix = HostPattern::parse(".example.com").unwrap();
assert_eq!(suffix, HostPattern::Suffix("example.com".into()));
assert!(suffix.matches("example.com")); assert!(suffix.matches("a.example.com"));
assert!(suffix.matches("a.b.example.com"));
assert!(!suffix.matches("example.com.evil.com"));
assert!(!suffix.matches("notexample.com")); assert!(!suffix.matches("fooexample.com"));
}
#[test]
fn host_pattern_rejects_invalid() {
for bad in [
".",
"",
"..",
"exam ple.com",
"a..b.com",
"ex@mple.com",
"-example.com", "example-.com", ] {
assert!(HostPattern::parse(bad).is_err(), "should reject {bad:?}");
}
}
#[test]
fn host_pattern_tolerates_trailing_dot() {
assert_eq!(
HostPattern::parse("example.com.").unwrap(),
HostPattern::Exact("example.com".into())
);
assert_eq!(
HostPattern::parse(".example.com.").unwrap(),
HostPattern::Suffix("example.com".into())
);
assert!(HostPattern::parse("example.com..").is_err());
}
#[test]
fn addr_spec_matches_dest_by_host_or_ip() {
let host = AddrSpec::host(HostPattern::Suffix("example.com".into()), None);
assert!(host.matches_dest(Some("api.example.com"), "203.0.113.1".parse().unwrap(), 443));
assert!(!host.matches_dest(None, "203.0.113.1".parse().unwrap(), 443));
let cidr = AddrSpec::new("10.0.0.0/8".parse().unwrap(), None);
assert!(cidr.matches_dest(Some("anything"), "10.1.2.3".parse().unwrap(), 80));
assert!(!cidr.matches_dest(Some("anything"), "192.0.2.1".parse().unwrap(), 80));
}
#[test]
fn matches_all_requires_both_families_and_any_port() {
assert!(AddrSpec::any().matches_all());
assert!(!AddrSpec::new("0.0.0.0/0".parse().unwrap(), None).matches_all());
assert!(!AddrSpec::new("::/0".parse().unwrap(), None).matches_all());
assert!(!AddrSpec::new("10.0.0.0/8".parse().unwrap(), None).matches_all());
let both_families_one_port = AddrSpec {
ports: Some(PortRange { min: 80, max: 80 }),
..AddrSpec::any()
};
assert!(!both_families_one_port.matches_all());
}
}