use crate::errors::{SecurityError, SecurityResult};
use governor::{
clock::DefaultClock,
state::{InMemoryState, NotKeyed},
Quota, RateLimiter as GovernorRateLimiter,
};
use std::collections::HashMap;
use std::net::IpAddr;
use std::num::NonZeroU32;
use std::sync::{Arc, RwLock};
use std::time::Duration;
#[derive(Debug, Clone)]
pub struct RateLimitConfig {
pub authenticated_rps: u32,
pub unauthenticated_rps: u32,
pub burst_size: u32,
pub window_seconds: u64,
pub ban_duration_seconds: u64,
pub ban_threshold: usize,
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
authenticated_rps: 100,
unauthenticated_rps: 10,
burst_size: 50,
window_seconds: 60,
ban_duration_seconds: 3600, ban_threshold: 10,
}
}
}
pub struct RateLimiter {
config: RateLimitConfig,
authenticated_limiter: Arc<GovernorRateLimiter<NotKeyed, InMemoryState, DefaultClock>>,
unauthenticated_limiter: Arc<GovernorRateLimiter<NotKeyed, InMemoryState, DefaultClock>>,
per_ip_limiters: Arc<RwLock<HashMap<IpAddr, IpLimiter>>>,
banned_ips: Arc<RwLock<HashMap<IpAddr, BanInfo>>>,
}
#[derive(Debug, Clone)]
struct IpLimiter {
limiter: Arc<GovernorRateLimiter<NotKeyed, InMemoryState, DefaultClock>>,
violations: usize,
last_violation: std::time::Instant,
}
#[derive(Debug, Clone)]
pub struct BanInfo {
pub banned_at: std::time::Instant,
pub reason: String,
pub violations: usize,
}
impl RateLimiter {
pub fn new(config: RateLimitConfig) -> Self {
let authenticated_quota = Quota::per_second(
NonZeroU32::new(config.authenticated_rps).unwrap_or(NonZeroU32::new(100).unwrap())
).allow_burst(
NonZeroU32::new(config.burst_size).unwrap_or(NonZeroU32::new(50).unwrap())
);
let unauthenticated_quota = Quota::per_second(
NonZeroU32::new(config.unauthenticated_rps).unwrap_or(NonZeroU32::new(10).unwrap())
).allow_burst(
NonZeroU32::new(config.burst_size / 5).unwrap_or(NonZeroU32::new(10).unwrap())
);
Self {
config,
authenticated_limiter: Arc::new(GovernorRateLimiter::direct(authenticated_quota)),
unauthenticated_limiter: Arc::new(GovernorRateLimiter::direct(
unauthenticated_quota,
)),
per_ip_limiters: Arc::new(RwLock::new(HashMap::new())),
banned_ips: Arc::new(RwLock::new(HashMap::new())),
}
}
pub fn check_request(
&self,
ip: IpAddr,
authenticated: bool,
) -> SecurityResult<()> {
if self.is_banned(ip) {
return Err(SecurityError::RateLimitExceeded(
"IP address is temporarily banned".to_string(),
));
}
let limiter = if authenticated {
&self.authenticated_limiter
} else {
&self.unauthenticated_limiter
};
if limiter.check().is_err() {
self.record_violation(ip, "Global rate limit exceeded");
return Err(SecurityError::RateLimitExceeded(
"Too many requests. Please try again later".to_string(),
));
}
let mut limiters = self.per_ip_limiters.write().unwrap();
let ip_limiter = limiters.entry(ip).or_insert_with(|| {
let quota = if authenticated {
Quota::per_second(
NonZeroU32::new(self.config.authenticated_rps / 10)
.unwrap_or(NonZeroU32::new(10).unwrap())
)
} else {
Quota::per_second(
NonZeroU32::new(self.config.unauthenticated_rps)
.unwrap_or(NonZeroU32::new(10).unwrap())
)
}
.allow_burst(NonZeroU32::new(10).unwrap());
IpLimiter {
limiter: Arc::new(GovernorRateLimiter::direct(quota)),
violations: 0,
last_violation: std::time::Instant::now(),
}
});
if ip_limiter.limiter.check().is_err() {
drop(limiters); self.record_violation(ip, "Per-IP rate limit exceeded");
return Err(SecurityError::RateLimitExceeded(format!(
"Too many requests from IP {}. Please try again later",
ip
)));
}
Ok(())
}
fn is_banned(&self, ip: IpAddr) -> bool {
let banned = self.banned_ips.read().unwrap();
if let Some(ban_info) = banned.get(&ip) {
let elapsed = ban_info.banned_at.elapsed();
let ban_duration = Duration::from_secs(self.config.ban_duration_seconds);
if elapsed < ban_duration {
return true;
}
}
false
}
fn record_violation(&self, ip: IpAddr, reason: &str) {
let mut limiters = self.per_ip_limiters.write().unwrap();
if let Some(ip_limiter) = limiters.get_mut(&ip) {
ip_limiter.violations += 1;
ip_limiter.last_violation = std::time::Instant::now();
if ip_limiter.violations >= self.config.ban_threshold {
let violations = ip_limiter.violations; drop(limiters); self.ban_ip(ip, reason.to_string(), violations);
}
}
}
fn ban_ip(&self, ip: IpAddr, reason: String, violations: usize) {
let mut banned = self.banned_ips.write().unwrap();
banned.insert(
ip,
BanInfo {
banned_at: std::time::Instant::now(),
reason,
violations,
},
);
tracing::warn!(
ip = %ip,
violations = violations,
"IP address banned due to rate limit violations"
);
}
pub fn ban(&self, ip: IpAddr, reason: String) {
self.ban_ip(ip, reason, 0);
}
pub fn unban(&self, ip: IpAddr) {
let mut banned = self.banned_ips.write().unwrap();
if banned.remove(&ip).is_some() {
tracing::info!(ip = %ip, "IP address unbanned");
}
}
pub fn get_banned_ips(&self) -> Vec<(IpAddr, BanInfo)> {
let banned = self.banned_ips.read().unwrap();
banned
.iter()
.map(|(ip, info)| (*ip, info.clone()))
.collect()
}
pub fn cleanup(&self) {
let mut banned = self.banned_ips.write().unwrap();
let ban_duration = Duration::from_secs(self.config.ban_duration_seconds);
banned.retain(|_, ban_info| ban_info.banned_at.elapsed() < ban_duration);
let mut limiters = self.per_ip_limiters.write().unwrap();
limiters.retain(|_, ip_limiter| {
ip_limiter.last_violation.elapsed() < Duration::from_secs(3600)
});
}
pub fn get_stats(&self) -> RateLimitStats {
let banned = self.banned_ips.read().unwrap();
let limiters = self.per_ip_limiters.read().unwrap();
RateLimitStats {
active_limiters: limiters.len(),
banned_ips: banned.len(),
total_violations: limiters.values().map(|l| l.violations).sum(),
}
}
}
#[derive(Debug, Clone)]
pub struct RateLimitStats {
pub active_limiters: usize,
pub banned_ips: usize,
pub total_violations: usize,
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::Ipv4Addr;
use std::thread;
use std::time::Duration;
#[test]
fn test_rate_limiter_basic() {
let config = RateLimitConfig {
authenticated_rps: 10,
unauthenticated_rps: 5,
burst_size: 10,
..Default::default()
};
let limiter = RateLimiter::new(config);
let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
assert!(limiter.check_request(ip, true).is_ok());
for _ in 0..9 {
assert!(limiter.check_request(ip, true).is_ok());
}
assert!(limiter.check_request(ip, true).is_err());
}
#[test]
fn test_per_ip_limiting() {
let config = RateLimitConfig {
authenticated_rps: 100,
unauthenticated_rps: 10,
burst_size: 20,
..Default::default()
};
let limiter = RateLimiter::new(config);
let ip1 = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
let ip2 = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 2));
for _ in 0..10 {
limiter.check_request(ip1, false).ok();
}
assert!(limiter.check_request(ip2, false).is_ok());
}
#[test]
fn test_banning() {
let config = RateLimitConfig {
authenticated_rps: 5,
unauthenticated_rps: 5,
burst_size: 10,
ban_threshold: 3,
ban_duration_seconds: 1,
..Default::default()
};
let limiter = RateLimiter::new(config);
let ip = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1));
for _ in 0..20 {
limiter.check_request(ip, false).ok();
}
assert!(limiter.is_banned(ip));
thread::sleep(Duration::from_secs(2));
limiter.cleanup();
assert!(!limiter.is_banned(ip));
}
#[test]
fn test_manual_ban() {
let limiter = RateLimiter::new(RateLimitConfig::default());
let ip = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1));
limiter.ban(ip, "Test ban".to_string());
assert!(limiter.is_banned(ip));
limiter.unban(ip);
assert!(!limiter.is_banned(ip));
}
#[test]
fn test_stats() {
let limiter = RateLimiter::new(RateLimitConfig::default());
let ip = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1));
limiter.check_request(ip, false).ok();
let stats = limiter.get_stats();
assert_eq!(stats.active_limiters, 1);
}
#[test]
fn test_cleanup() {
let limiter = RateLimiter::new(RateLimitConfig::default());
let ip = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1));
limiter.check_request(ip, false).ok();
assert_eq!(limiter.get_stats().active_limiters, 1);
limiter.cleanup();
assert_eq!(limiter.get_stats().active_limiters, 1);
}
}