use std::collections::BTreeMap;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use std::sync::OnceLock;
use argon2::password_hash::{PasswordHash, SaltString};
use argon2::{Argon2, PasswordHasher, PasswordVerifier};
use ipnet::IpNet;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct AccessConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub basic_auth: Option<BasicAuth>,
pub ip: IpRules,
#[serde(skip_serializing_if = "Option::is_none")]
pub rate_limit: Option<RateLimit>,
pub trusted_proxies: Vec<String>,
#[serde(default, skip_serializing_if = "waf_disabled")]
pub waf: crate::waf::WafConfig,
}
fn waf_disabled(waf: &crate::waf::WafConfig) -> bool {
!waf.is_enabled()
}
impl AccessConfig {
pub fn is_enforced(&self) -> bool {
self.basic_auth.is_some()
|| self.rate_limit.is_some()
|| !self.ip.allow.is_empty()
|| !self.ip.deny.is_empty()
|| self.waf.is_enabled()
}
pub fn is_trusted_proxy(&self, ip: IpAddr) -> bool {
ip_in_any(ip, &self.trusted_proxies)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct BasicAuth {
#[serde(default = "default_realm")]
pub realm: String,
#[serde(default)]
pub users: BTreeMap<String, String>,
}
fn default_realm() -> String {
"Restricted".to_string()
}
impl BasicAuth {
pub fn verify(&self, username: &str, password: &str) -> bool {
match self.users.get(username) {
Some(hash) => verify_password(password, hash),
None => {
let _ = verify_password(password, dummy_hash());
false
}
}
}
}
pub fn hash_password(password: &str) -> String {
let mut salt_bytes = [0u8; 16];
getrandom::getrandom(&mut salt_bytes).expect("system RNG");
let salt = SaltString::encode_b64(&salt_bytes).expect("salt encode");
Argon2::default()
.hash_password(password.as_bytes(), &salt)
.expect("argon2 hashing")
.to_string()
}
fn verify_password(password: &str, hash: &str) -> bool {
PasswordHash::new(hash)
.map(|parsed| {
Argon2::default()
.verify_password(password.as_bytes(), &parsed)
.is_ok()
})
.unwrap_or(false)
}
fn dummy_hash() -> &'static str {
static DUMMY: OnceLock<String> = OnceLock::new();
DUMMY.get_or_init(|| hash_password("dummy"))
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct IpRules {
pub allow: Vec<String>,
pub deny: Vec<String>,
}
impl IpRules {
pub fn allows(&self, ip: IpAddr) -> bool {
if ip_in_any(ip, &self.deny) {
return false;
}
if self.allow.is_empty() {
return true;
}
ip_in_any(ip, &self.allow)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct RateLimit {
pub rps: u32,
#[serde(default)]
pub burst: u32,
}
impl RateLimit {
pub fn burst_capacity(&self) -> u32 {
if self.burst == 0 {
self.rps.max(1)
} else {
self.burst
}
}
}
fn parse_net(spec: &str) -> Option<IpNet> {
let spec = spec.trim();
if let Ok(net) = spec.parse::<IpNet>() {
return Some(net);
}
if let Ok(ip) = spec.parse::<IpAddr>() {
let prefix = if ip.is_ipv4() { 32 } else { 128 };
return IpNet::new(ip, prefix).ok();
}
None
}
fn ip_in_any(ip: IpAddr, specs: &[String]) -> bool {
specs
.iter()
.filter_map(|spec| parse_net(spec))
.any(|net| net.contains(&ip))
}
pub fn resolve_client_ip(
peer: IpAddr,
forwarded_for: Option<&str>,
trusted_proxies: &[String],
) -> IpAddr {
if trusted_proxies.is_empty() || !ip_in_any(peer, trusted_proxies) {
return peer;
}
if let Some(chain) = forwarded_for {
for hop in chain.split(',').rev() {
if let Ok(ip) = hop.trim().parse::<IpAddr>() {
if !ip_in_any(ip, trusted_proxies) {
return ip;
}
}
}
}
peer
}
pub fn is_global_ip(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => ipv4_is_global(v4),
IpAddr::V6(v6) => ipv6_is_global(v6),
}
}
fn ipv4_is_global(a: Ipv4Addr) -> bool {
let [o0, o1, ..] = a.octets();
!(a.is_private()
|| a.is_loopback()
|| a.is_link_local() || a.is_unspecified()
|| a.is_broadcast()
|| a.is_documentation()
|| o0 == 0
|| (o0 == 100 && (64..128).contains(&o1)) || o0 >= 240) }
fn ipv6_is_global(a: Ipv6Addr) -> bool {
if a.is_loopback() || a.is_unspecified() || a.is_multicast() {
return false;
}
if let Some(v4) = a.to_ipv4_mapped() {
return ipv4_is_global(v4);
}
let first = a.segments()[0];
if (first & 0xfe00) == 0xfc00 {
return false; }
if (first & 0xffc0) == 0xfe80 {
return false; }
true
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn global_ip_blocks_internal_ranges() {
let public = ["8.8.8.8", "1.1.1.1", "2606:4700:4700::1111"];
let internal = [
"127.0.0.1",
"10.1.2.3",
"172.16.5.5",
"192.168.1.1",
"169.254.169.254", "100.64.0.1", "0.0.0.0",
"::1",
"fc00::1",
"fe80::1",
"::ffff:127.0.0.1", ];
for ip in public {
assert!(is_global_ip(ip.parse().unwrap()), "{ip} should be global");
}
for ip in internal {
assert!(
!is_global_ip(ip.parse().unwrap()),
"{ip} should be internal"
);
}
}
#[test]
fn ip_rules_deny_wins_and_allow_list_gates() {
let rules = IpRules {
allow: vec!["10.0.0.0/8".into()],
deny: vec!["10.1.2.3".into()],
};
assert!(rules.allows("10.5.5.5".parse().unwrap())); assert!(!rules.allows("10.1.2.3".parse().unwrap())); assert!(!rules.allows("192.168.0.1".parse().unwrap())); }
#[test]
fn ip_rules_default_allows_all() {
assert!(IpRules::default().allows("8.8.8.8".parse().unwrap()));
}
#[test]
fn ipv6_cidr_matches() {
let rules = IpRules {
allow: vec!["2001:db8::/32".into()],
deny: vec![],
};
assert!(rules.allows("2001:db8::1".parse().unwrap()));
assert!(!rules.allows("2001:dead::1".parse().unwrap()));
}
#[test]
fn xff_only_trusted_from_a_trusted_peer() {
let trusted = vec!["10.0.0.0/8".into()];
assert_eq!(
resolve_client_ip(
"10.0.0.5".parse().unwrap(),
Some("203.0.113.7, 10.0.0.9"),
&trusted,
),
"203.0.113.7".parse::<IpAddr>().unwrap()
);
assert_eq!(
resolve_client_ip(
"198.51.100.2".parse().unwrap(),
Some("203.0.113.7"),
&trusted,
),
"198.51.100.2".parse::<IpAddr>().unwrap()
);
}
#[test]
fn basic_auth_verifies_argon2() {
let hash = hash_password("hunter2");
assert!(
hash.starts_with("$argon2"),
"stored as an argon2 PHC string"
);
let auth = BasicAuth {
realm: "x".into(),
users: BTreeMap::from([("alice".to_string(), hash)]),
};
assert!(auth.verify("alice", "hunter2"));
assert!(!auth.verify("alice", "wrong"));
assert!(!auth.verify("bob", "hunter2")); }
}