use std::collections::HashMap;
use std::net::IpAddr;
use std::sync::Mutex;
use std::time::{Duration, Instant};
use axum::http::HeaderMap;
const MAX_KEYS: usize = 10_000;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct IpNetwork {
address: IpAddr,
prefix_len: u8,
}
impl IpNetwork {
pub fn parse(input: &str) -> Result<Self, String> {
let input = input.trim();
let (address_text, prefix_text) = input
.split_once('/')
.map_or((input, None), |(address, prefix)| (address, Some(prefix)));
let address = address_text
.parse::<IpAddr>()
.map_err(|_| "must be an IP address or CIDR range".to_string())?;
let max_prefix = match address {
IpAddr::V4(_) => 32,
IpAddr::V6(_) => 128,
};
let prefix_len = match prefix_text {
Some(prefix) => prefix
.parse::<u8>()
.map_err(|_| format!("invalid prefix length {prefix:?}"))?,
None => max_prefix,
};
if prefix_len > max_prefix {
return Err(format!(
"prefix length {prefix_len} exceeds the {max_prefix}-bit address family"
));
}
Ok(Self {
address,
prefix_len,
})
}
pub fn contains(&self, ip: IpAddr) -> bool {
match (self.address, normalize_ip(ip)) {
(IpAddr::V4(network), IpAddr::V4(ip)) => {
prefix_matches(&network.octets(), &ip.octets(), self.prefix_len)
}
(IpAddr::V6(network), IpAddr::V6(ip)) => {
prefix_matches(&network.octets(), &ip.octets(), self.prefix_len)
}
_ => false,
}
}
}
pub fn parse_trusted_proxies(entries: &[String]) -> Result<Vec<IpNetwork>, String> {
entries
.iter()
.enumerate()
.map(|(index, entry)| {
IpNetwork::parse(entry)
.map_err(|error| format!("trusted_proxies[{index}] ({entry:?}): {error}"))
})
.collect()
}
pub fn normalize_ip(ip: IpAddr) -> IpAddr {
let IpAddr::V6(ipv6) = ip else {
return ip;
};
let octets = ipv6.octets();
if octets[..10] == [0; 10] && octets[10] == 0xff && octets[11] == 0xff {
IpAddr::V4(std::net::Ipv4Addr::new(
octets[12], octets[13], octets[14], octets[15],
))
} else {
IpAddr::V6(ipv6)
}
}
fn prefix_matches(network: &[u8], ip: &[u8], prefix_len: u8) -> bool {
let whole_bytes = usize::from(prefix_len / 8);
let remaining_bits = prefix_len % 8;
network[..whole_bytes] == ip[..whole_bytes]
&& (remaining_bits == 0
|| (network[whole_bytes] & (!0u8 << (8 - remaining_bits)))
== (ip[whole_bytes] & (!0u8 << (8 - remaining_bits))))
}
fn header_ip(headers: &HeaderMap, name: &str) -> Option<IpAddr> {
headers
.get(name)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.trim().parse::<IpAddr>().ok())
}
fn x_forwarded_for_client_ip(headers: &HeaderMap, trusted_proxies: &[IpNetwork]) -> Option<IpAddr> {
let mut entries = Vec::new();
for value in headers.get_all("x-forwarded-for").iter() {
let line = value.to_str().ok()?;
entries.extend(line.split(','));
}
for entry in entries.into_iter().rev() {
let ip = normalize_ip(entry.trim().parse::<IpAddr>().ok()?);
if !trusted_proxies.iter().any(|range| range.contains(ip)) {
return Some(ip);
}
}
None
}
pub fn client_ip(peer: IpAddr, headers: &HeaderMap, trusted_proxies: &[IpNetwork]) -> String {
let peer = normalize_ip(peer);
if !trusted_proxies.iter().any(|range| range.contains(peer)) {
return peer.to_string();
}
let client = if headers.contains_key("x-forwarded-for") {
x_forwarded_for_client_ip(headers, trusted_proxies).unwrap_or(peer)
} else {
header_ip(headers, "x-real-ip")
.map(normalize_ip)
.unwrap_or(peer)
};
normalize_ip(client).to_string()
}
#[derive(Debug)]
pub struct RateLimiter {
attempts: Mutex<HashMap<String, Vec<Instant>>>,
max_attempts: usize,
window: Duration,
}
impl RateLimiter {
pub fn new(max_attempts: usize, window: Duration) -> Self {
Self {
attempts: Mutex::new(HashMap::new()),
max_attempts,
window,
}
}
fn sweep(map: &mut HashMap<String, Vec<Instant>>, window: Duration) {
let now = Instant::now();
map.retain(|_, entries| {
entries.retain(|t| now.duration_since(*t) < window);
!entries.is_empty()
});
}
pub fn check(&self, key: &str) -> bool {
let now = Instant::now();
let mut map = self.attempts.lock().unwrap_or_else(|e| e.into_inner());
if map.len() > MAX_KEYS {
Self::sweep(&mut map, self.window);
}
let entry = map.entry(key.to_string()).or_default();
entry.retain(|t| now.duration_since(*t) < self.window);
if entry.len() >= self.max_attempts {
return false;
}
entry.push(now);
true
}
pub fn peek(&self, key: &str) -> bool {
let now = Instant::now();
let mut map = self.attempts.lock().unwrap_or_else(|e| e.into_inner());
match map.get_mut(key) {
Some(entry) => {
entry.retain(|t| now.duration_since(*t) < self.window);
entry.len() < self.max_attempts
}
None => true,
}
}
pub fn record_failure(&self, key: &str) {
let now = Instant::now();
let mut map = self.attempts.lock().unwrap_or_else(|e| e.into_inner());
if map.len() > MAX_KEYS {
Self::sweep(&mut map, self.window);
}
let entry = map.entry(key.to_string()).or_default();
entry.retain(|t| now.duration_since(*t) < self.window);
entry.push(now);
}
pub fn retry_after(&self, key: &str) -> u64 {
let now = Instant::now();
let map = self.attempts.lock().unwrap_or_else(|e| e.into_inner());
match map.get(key) {
Some(entries) if !entries.is_empty() => {
let oldest = entries[0];
let elapsed = now.duration_since(oldest);
if elapsed < self.window {
(self.window - elapsed).as_secs() + 1
} else {
0
}
}
_ => 0,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn allows_under_limit() {
let rl = RateLimiter::new(3, Duration::from_secs(60));
assert!(rl.check("user1"));
assert!(rl.check("user1"));
assert!(rl.check("user1"));
}
#[test]
fn blocks_over_limit() {
let rl = RateLimiter::new(2, Duration::from_secs(60));
assert!(rl.check("user1"));
assert!(rl.check("user1"));
assert!(!rl.check("user1")); }
#[test]
fn different_keys_independent() {
let rl = RateLimiter::new(1, Duration::from_secs(60));
assert!(rl.check("user1"));
assert!(rl.check("user2")); assert!(!rl.check("user1")); }
#[test]
fn retry_after_nonzero_when_limited() {
let rl = RateLimiter::new(1, Duration::from_secs(60));
rl.check("user1");
assert!(rl.retry_after("user1") > 0);
}
#[test]
fn sweep_removes_expired_keys() {
let mut map: HashMap<String, Vec<Instant>> = HashMap::new();
let window = Duration::from_millis(1);
map.insert("old".into(), vec![Instant::now()]);
std::thread::sleep(Duration::from_millis(5));
RateLimiter::sweep(&mut map, window);
assert!(map.is_empty(), "expired keys should be evicted");
}
#[test]
fn peek_does_not_record() {
let rl = RateLimiter::new(1, Duration::from_secs(60));
assert!(rl.peek("user1"));
assert!(rl.peek("user1"));
assert!(rl.peek("user1"));
rl.record_failure("user1");
assert!(
!rl.peek("user1"),
"should be at limit after exactly one recorded failure"
);
}
#[test]
fn peek_reflects_recorded_failures_one_to_one() {
let rl = RateLimiter::new(5, Duration::from_secs(60));
for i in 0..5 {
assert!(rl.peek("u"), "attempt {i} should be allowed");
rl.record_failure("u");
}
assert!(!rl.peek("u"), "6th attempt should be blocked");
}
#[test]
fn untrusted_peer_ignores_spoofed_forwarded_headers() {
let trusted = parse_trusted_proxies(&["127.0.0.0/8".into()]).unwrap();
let peer = "203.0.113.5".parse().unwrap();
let mut h = HeaderMap::new();
h.insert("x-forwarded-for", "203.0.113.7, 10.0.0.1".parse().unwrap());
h.insert("x-real-ip", "198.51.100.4".parse().unwrap());
assert_eq!(client_ip(peer, &h, &trusted), "203.0.113.5");
}
#[test]
fn repeated_forwarded_for_lines_use_the_proxy_appended_line() {
let trusted = parse_trusted_proxies(&["127.0.0.0/8".into()]).unwrap();
let peer = "127.0.0.1".parse().unwrap();
let mut h = HeaderMap::new();
h.append("x-forwarded-for", "198.51.100.10".parse().unwrap());
h.append("x-forwarded-for", "203.0.113.9".parse().unwrap());
assert_eq!(client_ip(peer, &h, &trusted), "203.0.113.9");
}
#[test]
fn trusted_proxy_chain_skips_trusted_intermediate_hops() {
let trusted = parse_trusted_proxies(&["127.0.0.0/8".into(), "10.0.0.0/8".into()])
.unwrap();
let peer = "127.0.0.1".parse().unwrap();
let mut h = HeaderMap::new();
h.insert("x-forwarded-for", "203.0.113.9, 10.0.0.2".parse().unwrap());
assert_eq!(client_ip(peer, &h, &trusted), "203.0.113.9");
}
#[test]
fn all_trusted_forwarded_for_hops_fall_back_to_peer() {
let trusted = parse_trusted_proxies(&["127.0.0.0/8".into(), "10.0.0.0/8".into()])
.unwrap();
let peer = "127.0.0.1".parse().unwrap();
let mut h = HeaderMap::new();
h.insert("x-forwarded-for", "10.0.0.2, 127.0.0.2".parse().unwrap());
assert_eq!(client_ip(peer, &h, &trusted), "127.0.0.1");
}
#[test]
fn trusted_peer_falls_back_to_x_real_ip() {
let trusted = parse_trusted_proxies(&["127.0.0.0/8".into()]).unwrap();
let peer = "127.0.0.1".parse().unwrap();
let mut h = HeaderMap::new();
h.insert("x-real-ip", "198.51.100.4".parse().unwrap());
assert_eq!(client_ip(peer, &h, &trusted), "198.51.100.4");
}
#[test]
fn malformed_forwarded_for_does_not_fall_through_to_x_real_ip() {
let trusted = parse_trusted_proxies(&["127.0.0.0/8".into()]).unwrap();
let peer = "127.0.0.1".parse().unwrap();
let mut h = HeaderMap::new();
h.insert("x-forwarded-for", "1.2.3.4:5678".parse().unwrap());
h.insert("x-real-ip", "198.51.100.4".parse().unwrap());
assert_eq!(client_ip(peer, &h, &trusted), "127.0.0.1");
}
#[test]
fn trusted_peer_without_valid_headers_falls_back_to_peer() {
let trusted = parse_trusted_proxies(&["127.0.0.0/8".into()]).unwrap();
let peer = "127.0.0.1".parse().unwrap();
let h = HeaderMap::new();
assert_eq!(client_ip(peer, &h, &trusted), "127.0.0.1");
}
#[test]
fn cidr_matcher_handles_v4_v6_and_ipv4_mapped_ipv6() {
let v4_exact = IpNetwork::parse("192.0.2.1").unwrap();
assert!(v4_exact.contains("192.0.2.1".parse().unwrap()));
assert!(!v4_exact.contains("192.0.2.2".parse().unwrap()));
let v4_everything = IpNetwork::parse("0.0.0.0/0").unwrap();
assert!(v4_everything.contains("0.0.0.0".parse().unwrap()));
assert!(v4_everything.contains("255.255.255.255".parse().unwrap()));
let v4_range = IpNetwork::parse("10.0.0.0/8").unwrap();
assert!(v4_range.contains("10.255.255.255".parse().unwrap()));
assert!(!v4_range.contains("11.0.0.0".parse().unwrap()));
assert!(v4_range.contains("::ffff:10.1.2.3".parse().unwrap()));
let v6_loopback = IpNetwork::parse("::1/128").unwrap();
assert!(v6_loopback.contains("::1".parse().unwrap()));
assert!(!v6_loopback.contains("::2".parse().unwrap()));
let v6_everything = IpNetwork::parse("::/0").unwrap();
assert!(v6_everything.contains("::1".parse().unwrap()));
assert!(v6_everything.contains("2001:db8::1".parse().unwrap()));
let v6_range = IpNetwork::parse("2001:db8:1234:5678::/61").unwrap();
assert!(v6_range.contains("2001:db8:1234:567f::1".parse().unwrap()));
assert!(!v6_range.contains("2001:db8:1234:5680::1".parse().unwrap()));
}
#[test]
fn ipv4_mapped_ipv6_is_normalized_for_bucket_keys() {
let peer = "::ffff:192.0.2.1".parse().unwrap();
assert_eq!(client_ip(peer, &HeaderMap::new(), &[]), "192.0.2.1");
}
#[test]
fn invalid_trusted_proxy_range_is_rejected() {
let error = parse_trusted_proxies(&["10.0.0.0/99".into()]).unwrap_err();
assert!(error.contains("trusted_proxies[0]"));
}
}