use std::net::{IpAddr, Ipv6Addr, SocketAddr};
use std::sync::{Arc, LazyLock, Mutex};
use std::time::Duration;
use axum::extract::Request;
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use http::Method;
use quick_cache::sync::Cache;
use web_time::Instant;
use super::AppState;
use super::client_ip::{self, ClientIp, ExtractClientIP};
use crate::cnf;
use crate::ntw::error::Error as NetError;
const WARN_SUPPRESS_WINDOW: Duration = Duration::from_secs(60);
enum Decision {
Allowed,
Limited {
retry_after_secs: u64,
warn: bool,
},
}
struct Bucket {
tokens: f64,
refilled: Instant,
last_warned: Option<Instant>,
}
struct AuthRateLimiter {
buckets: Cache<IpAddr, Arc<Mutex<Bucket>>>,
burst: f64,
per_sec: f64,
}
impl AuthRateLimiter {
fn new(burst: u32, per_minute: u32, max_tracked: usize) -> Self {
Self {
buckets: Cache::new(max_tracked),
burst: f64::from(burst),
per_sec: f64::from(per_minute) / 60.0,
}
}
fn check_at(&self, key: IpAddr, now: Instant) -> Decision {
let bucket = self
.buckets
.get_or_insert_with(&key, || {
Ok::<_, std::convert::Infallible>(Arc::new(Mutex::new(Bucket {
tokens: self.burst,
refilled: now,
last_warned: None,
})))
})
.expect("bucket construction is infallible");
let mut bucket = bucket.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
let elapsed = now.saturating_duration_since(bucket.refilled).as_secs_f64();
bucket.tokens = (bucket.tokens + elapsed * self.per_sec).min(self.burst);
bucket.refilled = now;
if bucket.tokens >= 1.0 {
bucket.tokens -= 1.0;
Decision::Allowed
} else {
let retry_after_secs = ((1.0 - bucket.tokens) / self.per_sec).ceil().max(1.0) as u64;
let warn = bucket
.last_warned
.is_none_or(|at| now.saturating_duration_since(at) >= WARN_SUPPRESS_WINDOW);
if warn {
bucket.last_warned = Some(now);
}
Decision::Limited {
retry_after_secs,
warn,
}
}
}
}
fn derive_key(raw: &str) -> Option<IpAddr> {
let candidate = raw.rsplit(',').next()?.trim();
let ip = candidate
.parse::<IpAddr>()
.or_else(|_| candidate.parse::<SocketAddr>().map(|addr| addr.ip()))
.ok()?;
Some(match ip {
IpAddr::V4(v4) => IpAddr::V4(v4),
IpAddr::V6(v6) => match v6.to_ipv4_mapped() {
Some(v4) => IpAddr::V4(v4),
None => {
let s = v6.segments();
IpAddr::V6(Ipv6Addr::new(s[0], s[1], s[2], s[3], 0, 0, 0, 0))
}
},
})
}
fn request_key(request: &Request) -> Option<IpAddr> {
let strategy = request.extensions().get::<AppState>().map(|state| state.client_ip);
if let Some(ClientIp::Forwarded) = strategy {
let raw = request.headers().get(http::header::FORWARDED)?.to_str().ok()?;
return client_ip::parse_forwarded_for_last(raw).as_deref().and_then(derive_key);
}
request
.extensions()
.get::<ExtractClientIP>()
.and_then(|ExtractClientIP(ip)| ip.as_deref())
.and_then(derive_key)
}
static LIMITER: LazyLock<Option<AuthRateLimiter>> = LazyLock::new(|| {
if !*cnf::HTTP_AUTH_RATE_LIMIT_ENABLED {
return None;
}
let burst = *cnf::HTTP_AUTH_RATE_LIMIT_BURST;
let per_minute = *cnf::HTTP_AUTH_RATE_LIMIT_PER_MINUTE;
let max_tracked = *cnf::HTTP_AUTH_RATE_LIMIT_MAX_TRACKED_CLIENTS;
if burst == 0 || per_minute == 0 || max_tracked == 0 {
return None;
}
Some(AuthRateLimiter::new(burst, per_minute, max_tracked))
});
fn check(limiter: &AuthRateLimiter, key: IpAddr, endpoint: &str) -> Result<(), NetError> {
match limiter.check_at(key, Instant::now()) {
Decision::Allowed => Ok(()),
Decision::Limited {
retry_after_secs,
warn,
} => {
if warn {
warn!("Rate limiting authentication attempts from '{key}' on /{endpoint}");
} else {
debug!("Rate limited authentication attempt from '{key}' on /{endpoint}");
}
Err(NetError::TooManyRequests(retry_after_secs))
}
}
}
pub(super) async fn auth_rate_limit_middleware(request: Request, next: Next) -> Response {
if let Err(err) = admit(&request) {
return err.into_response();
}
next.run(request).await
}
fn admit(request: &Request) -> Result<(), NetError> {
let Some(limiter) = LIMITER.as_ref() else {
return Ok(());
};
if request.method() != Method::POST {
return Ok(());
}
let endpoint = match request.uri().path() {
"/signin" => "signin",
"/signup" => "signup",
_ => return Ok(()),
};
let Some(key) = request_key(request) else {
return Ok(());
};
check(limiter, key, endpoint)
}
#[cfg(test)]
mod tests {
use super::*;
fn ip(s: &str) -> IpAddr {
s.parse().unwrap()
}
#[test]
fn allows_burst_then_limits_with_retry_after() {
let l = AuthRateLimiter::new(3, 60, 16);
let now = Instant::now();
for i in 0..3 {
assert!(matches!(l.check_at(ip("1.2.3.4"), now), Decision::Allowed), "attempt {i}");
}
match l.check_at(ip("1.2.3.4"), now) {
Decision::Limited {
retry_after_secs,
warn,
} => {
assert_eq!(retry_after_secs, 1);
assert!(warn);
}
Decision::Allowed => panic!("fourth attempt within the burst window must be limited"),
}
assert!(matches!(
l.check_at(ip("1.2.3.4"), now),
Decision::Limited {
warn: false,
..
}
));
}
#[test]
fn budget_refills_over_time() {
let l = AuthRateLimiter::new(2, 60, 16);
let now = Instant::now();
for _ in 0..2 {
assert!(matches!(l.check_at(ip("9.9.9.9"), now), Decision::Allowed));
}
assert!(matches!(l.check_at(ip("9.9.9.9"), now), Decision::Limited { .. }));
assert!(matches!(
l.check_at(ip("9.9.9.9"), now + Duration::from_secs(1)),
Decision::Allowed
));
}
#[test]
fn buckets_are_per_address() {
let l = AuthRateLimiter::new(1, 60, 16);
let now = Instant::now();
assert!(matches!(l.check_at(ip("1.1.1.1"), now), Decision::Allowed));
assert!(matches!(l.check_at(ip("1.1.1.1"), now), Decision::Limited { .. }));
assert!(matches!(l.check_at(ip("2.2.2.2"), now), Decision::Allowed));
}
#[test]
fn warns_at_most_once_per_suppression_window() {
let l = AuthRateLimiter::new(1, 60, 16);
let t0 = Instant::now();
assert!(matches!(l.check_at(ip("3.3.3.3"), t0), Decision::Allowed));
assert!(matches!(
l.check_at(ip("3.3.3.3"), t0),
Decision::Limited {
warn: true,
..
}
));
let t1 = t0 + Duration::from_secs(2);
assert!(matches!(l.check_at(ip("3.3.3.3"), t1), Decision::Allowed));
assert!(matches!(
l.check_at(ip("3.3.3.3"), t1),
Decision::Limited {
warn: false,
..
}
));
let t2 = t0 + Duration::from_secs(61);
assert!(matches!(l.check_at(ip("3.3.3.3"), t2), Decision::Allowed));
assert!(matches!(
l.check_at(ip("3.3.3.3"), t2),
Decision::Limited {
warn: true,
..
}
));
}
#[test]
fn ipv6_addresses_share_a_slash64_bucket() {
let a = derive_key("2001:db8:1:2:aaaa::1").unwrap();
let b = derive_key("2001:db8:1:2::beef").unwrap();
assert_eq!(a, b, "addresses inside one /64 must share a key");
let c = derive_key("2001:db8:1:3::1").unwrap();
assert_ne!(a, c, "addresses in different /64s must not share a key");
let l = AuthRateLimiter::new(1, 60, 16);
let now = Instant::now();
assert!(matches!(l.check_at(a, now), Decision::Allowed));
assert!(matches!(l.check_at(b, now), Decision::Limited { .. }));
}
#[test]
fn derive_key_uses_only_the_last_chain_element() {
assert_eq!(derive_key("6.6.6.6, 7.7.7.7"), Some(ip("7.7.7.7")));
assert_eq!(derive_key("1.1.1.1, 2.2.2.2, 7.7.7.7"), Some(ip("7.7.7.7")));
let l = AuthRateLimiter::new(1, 60, 16);
let now = Instant::now();
let first = derive_key("1.1.1.1, 9.9.9.9").unwrap();
let second = derive_key("1.1.1.2, 9.9.9.9").unwrap();
assert!(matches!(l.check_at(first, now), Decision::Allowed));
assert!(matches!(l.check_at(second, now), Decision::Limited { .. }));
}
#[test]
fn forwarded_key_comes_from_the_proxy_appended_element() {
let key = client_ip::parse_forwarded_for_last("for=1.1.1.1, for=9.9.9.9")
.as_deref()
.and_then(derive_key);
assert_eq!(key, Some(ip("9.9.9.9")));
let spoofed = client_ip::parse_forwarded_for_last("for=1.1.1.2, for=9.9.9.9")
.as_deref()
.and_then(derive_key);
assert_eq!(spoofed, key);
let v6 = client_ip::parse_forwarded_for_last(r#"for="[2001:db8:1:2::1]:443""#)
.as_deref()
.and_then(derive_key);
assert_eq!(v6, Some(ip("2001:db8:1:2::")));
}
#[test]
fn derive_key_parses_addresses_and_rejects_garbage() {
assert_eq!(derive_key("1.2.3.4"), Some(ip("1.2.3.4")));
assert_eq!(derive_key(" 1.2.3.4 "), Some(ip("1.2.3.4")));
assert_eq!(derive_key("1.2.3.4:5555"), Some(ip("1.2.3.4")));
assert_eq!(derive_key("[2001:db8::1]:443"), Some(ip("2001:db8::")));
assert_eq!(derive_key("::ffff:1.2.3.4"), Some(ip("1.2.3.4")));
assert_eq!(derive_key("not-an-ip"), None);
assert_eq!(derive_key(""), None);
assert_eq!(derive_key("unknown, also-unknown"), None);
}
#[test]
fn tracking_is_bounded_at_capacity() {
let l = AuthRateLimiter::new(1, 60, 4);
let now = Instant::now();
for i in 0..64u8 {
let key = IpAddr::V4(std::net::Ipv4Addr::new(10, 0, 0, i));
assert!(matches!(l.check_at(key, now), Decision::Allowed));
}
assert!(l.buckets.len() <= 4, "store grew past capacity: {}", l.buckets.len());
assert!(matches!(l.check_at(ip("10.0.0.63"), now), Decision::Limited { .. }));
}
}