use std::net::IpAddr;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct Rule {
network: u128,
prefix: u8,
}
impl Rule {
fn contains(&self, probe: u128) -> bool {
if self.prefix == 0 {
return true;
}
(probe ^ self.network) >> (128 - self.prefix) == 0
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct InvalidIpEntry(pub String);
impl std::fmt::Display for InvalidIpEntry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "not an IP address or CIDR range: {}", self.0)
}
}
impl std::error::Error for InvalidIpEntry {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct IpAllowlist {
rules: Vec<Rule>,
allow_all_v4: bool,
allow_all_v6: bool,
}
const ALLOW_ALL_V6: [&str; 3] = ["::/0", "::", "::0"];
const ALLOW_ALL_V4: [&str; 2] = ["0.0.0.0/0", "0.0.0.0"];
impl Default for IpAllowlist {
fn default() -> Self {
Self {
rules: vec![
rule_of(IpAddr::from([127, 0, 0, 1]), None),
rule_of(IpAddr::from([0, 0, 0, 0, 0, 0, 0, 1]), None),
],
allow_all_v4: false,
allow_all_v6: false,
}
}
}
impl IpAllowlist {
pub fn parse<I, S>(entries: I) -> Result<Self, InvalidIpEntry>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let mut out = Self::deny_all();
for entry in entries {
let entry = entry.as_ref();
if ALLOW_ALL_V6.contains(&entry) {
out.allow_all_v6 = true;
} else if ALLOW_ALL_V4.contains(&entry) {
out.allow_all_v4 = true;
} else {
out.rules.push(parse_entry(entry)?);
}
}
Ok(out)
}
pub fn parse_env(value: &str) -> Result<Self, InvalidIpEntry> {
Self::parse(value.split(','))
}
pub fn deny_all() -> Self {
Self {
rules: Vec::new(),
allow_all_v4: false,
allow_all_v6: false,
}
}
pub fn allows(&self, peer: IpAddr) -> bool {
let peer_is_v4 = matches!(peer, IpAddr::V4(_));
if peer_is_v4 && self.allow_all_v4 {
return true;
}
if !peer_is_v4 && self.allow_all_v6 {
return true;
}
let probe = to_v6(peer);
self.rules.iter().any(|r| r.contains(probe))
}
pub fn is_deny_all(&self) -> bool {
self.rules.is_empty() && !self.allow_all_v4 && !self.allow_all_v6
}
}
fn parse_entry(entry: &str) -> Result<Rule, InvalidIpEntry> {
let invalid = || InvalidIpEntry(entry.to_string());
let (address, mask) = match entry.split_once('/') {
Some((a, m)) => (a, Some(m.parse::<u8>().map_err(|_| invalid())?)),
None => (entry, None),
};
let address: IpAddr = address.parse().map_err(|_| invalid())?;
let width = if address.is_ipv4() { 32 } else { 128 };
if mask.is_some_and(|m| m > width) {
return Err(invalid());
}
Ok(rule_of(address, mask))
}
fn rule_of(address: IpAddr, mask: Option<u8>) -> Rule {
let prefix = match address {
IpAddr::V4(_) => 96 + mask.unwrap_or(32),
IpAddr::V6(_) => mask.unwrap_or(128),
};
let network = to_v6(address);
Rule {
network: mask_to(network, prefix),
prefix,
}
}
fn mask_to(value: u128, prefix: u8) -> u128 {
if prefix == 0 {
0
} else {
value & (u128::MAX << (128 - prefix))
}
}
fn to_v6(address: IpAddr) -> u128 {
match address {
IpAddr::V4(v4) => u128::from(v4.to_ipv6_mapped()),
IpAddr::V6(v6) => u128::from(v6),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn ip(s: &str) -> IpAddr {
s.parse().expect("test address")
}
fn list(entries: &[&str]) -> IpAllowlist {
IpAllowlist::parse(entries).expect("test entries")
}
#[test]
fn the_default_is_loopback_only() {
let allow = IpAllowlist::default();
for peer in ["127.0.0.1", "::1", "::ffff:127.0.0.1"] {
assert!(allow.allows(ip(peer)), "{peer} must be allowed");
}
for peer in [
"127.0.0.2",
"::ffff:127.0.0.2",
"10.0.0.1",
"192.168.1.10",
"::2",
"2001:db8::1",
] {
assert!(!allow.allows(ip(peer)), "{peer} must be refused");
}
}
const ORACLE: &[(&str, &[(&str, bool)])] = &[
(
"::/0",
&[
("127.0.0.1", false),
("::1", true),
("::ffff:127.0.0.1", true),
("127.0.0.2", false),
("10.1.2.3", false),
("::ffff:10.1.2.3", true),
("2001:db8::1", true),
],
),
(
"::",
&[
("127.0.0.1", false),
("::1", true),
("::ffff:127.0.0.1", true),
("2001:db8::1", true),
],
),
(
"::0",
&[
("127.0.0.1", false),
("::1", true),
("::ffff:127.0.0.1", true),
("2001:db8::1", true),
],
),
(
"0.0.0.0/0",
&[
("127.0.0.1", true),
("::1", false),
("::ffff:127.0.0.1", false),
("127.0.0.2", true),
("10.1.2.3", true),
("::ffff:10.1.2.3", false),
("2001:db8::1", false),
],
),
(
"0.0.0.0",
&[
("127.0.0.1", true),
("::1", false),
("::ffff:127.0.0.1", false),
("10.1.2.3", true),
],
),
(
"127.0.0.1",
&[
("127.0.0.1", true),
("::1", false),
("::ffff:127.0.0.1", true),
("127.0.0.2", false),
("10.1.2.3", false),
],
),
(
"::1",
&[
("127.0.0.1", false),
("::1", true),
("::ffff:127.0.0.1", false),
],
),
(
"::/64",
&[
("127.0.0.1", true),
("::1", true),
("::ffff:127.0.0.1", true),
("10.1.2.3", true),
("2001:db8::1", false),
],
),
(
"10.0.0.0/8",
&[
("127.0.0.1", false),
("::ffff:127.0.0.1", false),
("10.1.2.3", true),
("::ffff:10.1.2.3", true),
("2001:db8::1", false),
],
),
(
"2000::/3",
&[
("127.0.0.1", false),
("::ffff:127.0.0.1", false),
("2001:db8::1", true),
],
),
(
"::ffff:0.0.0.0/96",
&[
("127.0.0.1", true),
("::1", false),
("::ffff:127.0.0.1", true),
("10.1.2.3", true),
],
),
];
#[test]
fn every_oracle_row_matches() {
for (rule, peers) in ORACLE {
let allow = list(&[rule]);
for (peer, expected) in *peers {
assert_eq!(allow.allows(ip(peer)), *expected, "[{rule}] against {peer}");
}
}
}
#[test]
fn the_allow_all_literals_are_scoped_to_one_family() {
assert!(
!list(&["::/0"]).allows(ip("127.0.0.1")),
"::/0 is allowAllIpv6 and an IPv4 peer is not IPv6"
);
assert!(
!list(&["0.0.0.0/0"]).allows(ip("::ffff:127.0.0.1")),
"isIPv4 is false for a mapped address, so allowAllIpv4 does not apply to it"
);
let both = list(&["0.0.0.0/0", "::0"]);
for peer in ["127.0.0.1", "::1", "::ffff:127.0.0.1", "2001:db8::1"] {
assert!(both.allows(ip(peer)), "{peer}");
}
}
#[test]
fn the_bare_spellings_are_allow_all_rather_than_one_address() {
assert!(list(&["::"]).allows(ip("2001:db8::1")));
assert!(list(&["0.0.0.0"]).allows(ip("203.0.113.9")));
}
#[test]
fn a_numerically_equivalent_spelling_is_not_a_special_case() {
assert!(!list(&["0.0.0.0/32"]).allows(ip("127.0.0.1")));
}
#[test]
fn the_two_families_stay_separate_for_genuine_addresses() {
assert!(!list(&["0.0.0.0/0"]).allows(ip("2001:db8::1")));
assert!(!list(&["2000::/3"]).allows(ip("32.0.0.1")));
}
#[test]
fn cidr_ranges_match_by_prefix() {
let allow = list(&["10.0.1.0/24", "2001:db8::/32"]);
assert!(allow.allows(ip("10.0.1.0")));
assert!(allow.allows(ip("10.0.1.255")));
assert!(!allow.allows(ip("10.0.2.0")));
assert!(allow.allows(ip("2001:db8:1234::9")));
assert!(!allow.allows(ip("2001:db9::1")));
}
#[test]
fn a_non_canonical_base_is_masked_rather_than_refused() {
let allow = list(&["10.0.1.5/24"]);
assert!(allow.allows(ip("10.0.1.7")));
assert!(allow.allows(ip("10.0.1.5")));
assert!(!allow.allows(ip("10.0.2.7")));
}
#[test]
fn an_empty_allowlist_denies_everything_including_loopback() {
let allow = IpAllowlist::deny_all();
assert!(allow.is_deny_all());
for peer in ["127.0.0.1", "::1", "10.0.0.1"] {
assert!(!allow.allows(ip(peer)));
}
assert_eq!(
IpAllowlist::parse::<[&str; 0], &str>([]).expect("empty"),
allow
);
}
#[test]
fn the_env_spelling_is_comma_separated() {
let allow = IpAllowlist::parse_env("127.0.0.1,10.0.1.0/24,::1").expect("parse");
assert!(allow.allows(ip("127.0.0.1")));
assert!(allow.allows(ip("10.0.1.9")));
assert!(allow.allows(ip("::1")));
assert!(!allow.allows(ip("10.0.2.9")));
}
#[test]
fn the_env_spelling_matches_upstreams_strictness() {
assert_eq!(
IpAllowlist::parse_env("127.0.0.1, ::1").expect_err("space"),
InvalidIpEntry(" ::1".to_string())
);
assert_eq!(
IpAllowlist::parse_env("").expect_err("empty"),
InvalidIpEntry(String::new())
);
assert_eq!(
IpAllowlist::parse_env(" ").expect_err("blank"),
InvalidIpEntry(" ".to_string())
);
}
#[test]
fn an_out_of_range_prefix_is_refused_at_boot_rather_than_on_the_first_request() {
assert_eq!(
IpAllowlist::parse(["127.0.0.1/999"]).expect_err("mask"),
InvalidIpEntry("127.0.0.1/999".to_string())
);
}
#[test]
fn a_malformed_entry_is_refused_by_name() {
for entry in [
"",
"localhost",
"127.0.0.1/33",
"::1/129",
"127.0.0.1/x",
"1.2.3",
] {
let e = IpAllowlist::parse([entry]).expect_err("must refuse");
assert_eq!(e, InvalidIpEntry(entry.to_string()));
}
assert!(IpAllowlist::parse(["fe80::1%lo0"]).is_err());
}
}