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;
const MAX_KEY_BYTES: usize = 1024;
pub(crate) fn retry_after_message(prefix: &str, retry_after: u64) -> String {
if retry_after == 0 {
format!("{prefix} — try again later")
} else {
format!("{prefix} — try again in {retry_after} seconds")
}
}
#[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_or(peer, normalize_ip)
};
normalize_ip(client).to_string()
}
#[derive(Debug)]
pub struct RateLimiter {
state: Mutex<RateLimitState>,
max_attempts: usize,
window: Duration,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AttemptId(u64);
#[derive(Debug)]
struct Attempt {
id: AttemptId,
at: Instant,
}
#[derive(Debug)]
struct RateLimitState {
attempts: HashMap<String, Vec<Attempt>>,
last_sweep: Instant,
next_id: u64,
}
impl RateLimitState {
fn allocate(&mut self) -> Option<AttemptId> {
let id = AttemptId(self.next_id);
self.next_id = self.next_id.checked_add(1)?;
Some(id)
}
}
impl RateLimiter {
pub fn new(max_attempts: usize, window: Duration) -> Self {
Self {
state: Mutex::new(RateLimitState {
attempts: HashMap::new(),
last_sweep: Instant::now(),
next_id: 0,
}),
max_attempts,
window,
}
}
fn sweep(map: &mut HashMap<String, Vec<Attempt>>, now: Instant, window: Duration) {
map.retain(|_, entries| {
entries.retain(|entry| now.saturating_duration_since(entry.at) < window);
!entries.is_empty()
});
}
fn make_room(state: &mut RateLimitState, now: Instant, window: Duration) -> bool {
if now.saturating_duration_since(state.last_sweep) >= window {
Self::sweep(&mut state.attempts, now, window);
state.last_sweep = now;
}
state.attempts.len() < MAX_KEYS
}
pub fn check(&self, key: &str) -> bool {
self.reserve(key).is_some()
}
pub fn reserve(&self, key: &str) -> Option<AttemptId> {
if key.len() > MAX_KEY_BYTES {
return None;
}
let now = Instant::now();
let mut state = self.state.lock().unwrap_or_else(|e| e.into_inner());
if !state.attempts.contains_key(key) && !Self::make_room(&mut state, now, self.window) {
return None;
}
let id = state.allocate()?;
let entry = state.attempts.entry(key.to_string()).or_default();
entry.retain(|attempt| now.saturating_duration_since(attempt.at) < self.window);
if entry.len() >= self.max_attempts {
return None;
}
entry.push(Attempt { id, at: now });
Some(id)
}
pub fn refund(&self, key: &str, id: AttemptId) {
if key.len() > MAX_KEY_BYTES {
return;
}
let now = Instant::now();
let mut state = self.state.lock().unwrap_or_else(|e| e.into_inner());
let Some(entry) = state.attempts.get_mut(key) else {
return;
};
entry.retain(|attempt| {
attempt.id != id && now.saturating_duration_since(attempt.at) < self.window
});
if entry.is_empty() {
state.attempts.remove(key);
}
}
pub fn retry_after(&self, key: &str) -> u64 {
let state = self.state.lock().unwrap_or_else(|e| e.into_inner());
let now = Instant::now();
match state.attempts.get(key) {
Some(entries) if !entries.is_empty() => {
let oldest = entries[0].at;
let elapsed = now.saturating_duration_since(oldest);
self.window
.checked_sub(elapsed)
.filter(|remaining| !remaining.is_zero())
.map_or(0, |remaining| remaining.as_secs() + 1)
}
_ => 0,
}
}
#[cfg(test)]
pub(crate) fn contains_key(&self, key: &str) -> bool {
self.state
.lock()
.unwrap_or_else(|e| e.into_inner())
.attempts
.contains_key(key)
}
}
pub struct Reservation<'a> {
limiter: &'a RateLimiter,
first: (String, AttemptId),
second: (String, AttemptId),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ReservationRejection {
First,
Second,
}
impl<'a> Reservation<'a> {
pub fn acquire(
limiter: &'a RateLimiter,
first: &str,
second: &str,
) -> Result<Self, ReservationRejection> {
let first_id = limiter.reserve(first).ok_or(ReservationRejection::First)?;
let Some(second_id) = limiter.reserve(second) else {
limiter.refund(first, first_id);
return Err(ReservationRejection::Second);
};
Ok(Self {
limiter,
first: (first.to_string(), first_id),
second: (second.to_string(), second_id),
})
}
pub fn refund(self) {
self.limiter.refund(&self.first.0, self.first.1);
self.limiter.refund(&self.second.0, self.second.1);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_refund_gives_back_exactly_the_slot_its_reservation_took() {
let rl = RateLimiter::new(2, Duration::from_secs(60));
let first = rl.reserve("k").expect("free");
assert!(rl.reserve("k").is_some());
assert!(rl.reserve("k").is_none(), "the budget is spent");
rl.refund("k", first);
assert!(rl.reserve("k").is_some(), "the refunded slot came back");
assert!(rl.reserve("k").is_none(), "and only that one slot");
}
#[test]
fn a_refund_removes_only_its_own_attempt_not_the_newest() {
let rl = RateLimiter::new(3, Duration::from_secs(60));
let slow_success = rl.reserve("shared").expect("free");
let later_failure = rl.reserve("shared").expect("free");
assert_ne!(slow_success, later_failure);
rl.refund("shared", slow_success);
assert!(rl.retry_after("shared") > 0, "the failure is still held");
assert!(rl.reserve("shared").is_some());
assert!(rl.reserve("shared").is_some());
assert!(
rl.reserve("shared").is_none(),
"the later failure still costs its slot"
);
rl.refund("shared", later_failure);
assert!(rl.reserve("shared").is_some());
}
#[test]
fn refunding_the_last_attempt_frees_the_key() {
let rl = RateLimiter::new(2, Duration::from_secs(60));
let id = rl.reserve("k").expect("free");
assert!(rl.contains_key("k"));
rl.refund("k", id);
assert!(
!rl.contains_key("k"),
"an emptied key must not hold a slot in the bounded table"
);
}
#[test]
fn refunding_an_unknown_key_or_a_spent_id_is_a_no_op() {
let rl = RateLimiter::new(1, Duration::from_secs(60));
rl.refund("never-seen", AttemptId(999));
assert!(!rl.contains_key("never-seen"));
let id = rl.reserve("k").expect("free");
rl.refund("k", id);
rl.refund("k", id);
assert!(rl.reserve("k").is_some());
assert!(rl.reserve("k").is_none(), "still capped at one");
}
#[test]
fn refunding_an_expired_attempt_is_a_no_op() {
let rl = RateLimiter::new(1, Duration::from_secs(0));
let id = rl.reserve("k").expect("free");
rl.refund("k", id);
assert!(rl.reserve("k").is_some());
}
#[test]
fn an_oversized_key_is_refused_and_refunding_it_is_harmless() {
let rl = RateLimiter::new(1, Duration::from_secs(60));
let huge = "x".repeat(MAX_KEY_BYTES + 1);
assert!(rl.reserve(&huge).is_none());
rl.refund(&huge, AttemptId(0));
assert!(!rl.contains_key(&huge));
}
#[test]
fn retry_after_reflects_only_the_attempts_still_held() {
let rl = RateLimiter::new(2, Duration::from_secs(60));
let refunded = rl.reserve("k").expect("free");
let kept = rl.reserve("k").expect("free");
rl.refund("k", refunded);
assert!(rl.retry_after("k") > 0, "the kept failure still counts");
rl.refund("k", kept);
assert_eq!(rl.retry_after("k"), 0, "nothing is held any more");
}
#[test]
fn concurrent_checks_admit_at_most_max_attempts() {
const MAX: usize = 5;
const RACERS: usize = 64;
let rl = std::sync::Arc::new(RateLimiter::new(MAX, Duration::from_secs(60)));
let admitted = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
let barrier = std::sync::Arc::new(std::sync::Barrier::new(RACERS));
let mut handles = Vec::new();
for _ in 0..RACERS {
let (rl, admitted, barrier) = (rl.clone(), admitted.clone(), barrier.clone());
handles.push(std::thread::spawn(move || {
barrier.wait();
if rl.check("contended") {
admitted.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}
}));
}
for handle in handles {
handle.join().unwrap();
}
assert_eq!(
admitted.load(std::sync::atomic::Ordering::SeqCst),
MAX,
"exactly the budget, no more and no fewer"
);
}
#[test]
fn refunds_are_safe_under_contention() {
let rl = std::sync::Arc::new(RateLimiter::new(4, Duration::from_secs(60)));
let barrier = std::sync::Arc::new(std::sync::Barrier::new(16));
let mut handles = Vec::new();
for _ in 0..16 {
let (rl, barrier) = (rl.clone(), barrier.clone());
handles.push(std::thread::spawn(move || {
barrier.wait();
if let Some(id) = rl.reserve("busy") {
rl.refund("busy", id);
}
}));
}
for handle in handles {
handle.join().unwrap();
}
for _ in 0..4 {
assert!(rl.check("busy"), "every slot survived the churn");
}
assert!(!rl.check("busy"));
}
#[test]
fn retry_after_message_omits_zero_delay() {
assert_eq!(
retry_after_message("too many requests", 0),
"too many requests — try again later"
);
assert_eq!(
retry_after_message("too many requests", 3),
"too many requests — try again in 3 seconds"
);
}
#[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<Attempt>> = HashMap::new();
let window = Duration::from_millis(1);
map.insert(
"old".into(),
vec![Attempt {
id: AttemptId(0),
at: Instant::now(),
}],
);
std::thread::sleep(Duration::from_millis(5));
RateLimiter::sweep(&mut map, Instant::now(), window);
assert!(map.is_empty(), "expired keys should be evicted");
}
#[test]
fn full_table_rejects_new_identity_without_evicting_existing_state() {
let rl = RateLimiter::new(1, Duration::from_secs(60));
for index in 0..MAX_KEYS {
assert!(rl.check(&format!("identity-{index}")));
}
assert!(!rl.check("new-identity"));
assert!(!rl.check("identity-0"));
}
#[test]
fn expired_keys_are_swept_before_admitting_new_identity() {
let rl = RateLimiter::new(1, Duration::ZERO);
assert!(rl.check("expired"));
assert!(rl.check("new"));
}
#[test]
fn a_full_table_refuses_new_identities_and_keeps_existing_ones() {
let rl = RateLimiter::new(5, Duration::from_secs(3600));
for index in 0..MAX_KEYS {
assert!(rl.check(&format!("identity-{index}")));
}
assert!(
!rl.check("new-identity"),
"a saturated table must not admit an untracked identity"
);
assert!(!rl.contains_key("new-identity"));
assert!(rl.check("identity-0"));
}
#[test]
fn a_swept_table_admits_new_identities_again() {
let rl = RateLimiter::new(5, Duration::from_secs(0));
for index in 0..MAX_KEYS {
assert!(rl.check(&format!("identity-{index}")));
}
assert!(rl.check("new-identity"));
}
#[test]
fn oversized_keys_are_rejected_without_allocation() {
let rl = RateLimiter::new(1, Duration::from_secs(60));
let oversized = "x".repeat(MAX_KEY_BYTES + 1);
assert!(!rl.check(&oversized));
rl.refund(&oversized, AttemptId(0));
assert_eq!(rl.retry_after(&oversized), 0);
}
#[test]
fn default_proxy_configuration_does_not_trust_loopback_headers() {
let peer = "127.0.0.1".parse().unwrap();
let mut headers = HeaderMap::new();
headers.insert("x-forwarded-for", "198.51.100.9".parse().unwrap());
assert_eq!(client_ip(peer, &headers, &[]), "127.0.0.1");
}
#[test]
fn a_failed_attempt_costs_exactly_one_slot() {
let rl = RateLimiter::new(5, Duration::from_secs(60));
for i in 0..5 {
assert!(rl.check("u"), "attempt {i} should be allowed");
}
assert!(!rl.check("u"), "6th attempt should be blocked");
}
#[test]
fn a_reservation_takes_both_keys_and_gives_both_back() {
let rl = RateLimiter::new(1, Duration::from_secs(60));
let held = Reservation::acquire(&rl, "ip", "id").expect("both keys are free");
assert!(!rl.check("ip"), "the address slot is held");
assert!(!rl.check("id"), "the account slot is held");
held.refund();
assert!(rl.check("ip"));
assert!(rl.check("id"));
}
#[test]
fn a_refused_second_key_releases_the_first() {
let rl = RateLimiter::new(1, Duration::from_secs(60));
assert!(rl.check("id"), "spend the account's only slot");
assert!(matches!(
Reservation::acquire(&rl, "ip", "id"),
Err(ReservationRejection::Second)
));
assert!(rl.check("ip"), "the address slot must have been given back");
}
#[test]
fn a_refused_first_key_is_reported_exactly() {
let rl = RateLimiter::new(1, Duration::from_secs(60));
assert!(rl.check("ip"), "spend the address's only slot");
assert!(matches!(
Reservation::acquire(&rl, "ip", "id"),
Err(ReservationRejection::First)
));
assert!(rl.check("id"), "the untouched account key stays free");
}
#[test]
fn a_reservation_that_is_never_refunded_stays_spent() {
let rl = RateLimiter::new(2, Duration::from_secs(60));
let _spent = Reservation::acquire(&rl, "ip", "id").expect("free");
drop(_spent);
assert!(rl.check("ip"));
assert!(!rl.check("ip"));
}
#[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]"));
}
}