use std::collections::HashMap;
use std::sync::Mutex;
use std::time::{Duration, Instant};
use axum::http::HeaderMap;
const MAX_KEYS: usize = 10_000;
pub fn client_ip(headers: &HeaderMap) -> String {
if let Some(v) = headers.get("x-forwarded-for")
&& let Ok(s) = v.to_str()
{
if let Some(first) = s.split(',').next() {
let trimmed = first.trim();
if !trimmed.is_empty() {
return trimmed.to_string();
}
}
}
if let Some(v) = headers.get("x-real-ip")
&& let Ok(s) = v.to_str()
{
let trimmed = s.trim();
if !trimmed.is_empty() {
return trimmed.to_string();
}
}
"unknown".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 client_ip_prefers_x_forwarded_for_first_hop() {
let mut h = HeaderMap::new();
h.insert("x-forwarded-for", "203.0.113.7, 10.0.0.1".parse().unwrap());
assert_eq!(client_ip(&h), "203.0.113.7");
}
#[test]
fn client_ip_falls_back_to_x_real_ip() {
let mut h = HeaderMap::new();
h.insert("x-real-ip", "198.51.100.4".parse().unwrap());
assert_eq!(client_ip(&h), "198.51.100.4");
}
#[test]
fn client_ip_unknown_when_no_proxy_headers() {
let h = HeaderMap::new();
assert_eq!(client_ip(&h), "unknown");
}
}