use dashmap::DashMap;
use serde::Serialize;
use std::net::IpAddr;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct RateLimitConfig {
pub requests_per_minute: u32,
pub burst_size: u32,
pub refill_interval_ms: u64,
pub whitelist: Vec<IpAddr>,
pub enabled: bool,
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
requests_per_minute: 60,
burst_size: 10,
refill_interval_ms: 1000,
whitelist: vec![],
enabled: true,
}
}
}
#[derive(Debug)]
pub struct TokenBucket {
tokens: f64,
max_tokens: f64,
refill_rate: f64,
last_update: Instant,
}
impl TokenBucket {
pub fn new(max_tokens: f64, refill_rate: f64) -> Self {
Self {
tokens: max_tokens,
max_tokens,
refill_rate,
last_update: Instant::now(),
}
}
pub fn try_consume(&mut self) -> bool {
self.refill();
if self.tokens >= 1.0 {
self.tokens -= 1.0;
true
} else {
false
}
}
pub fn tokens_remaining(&mut self) -> f64 {
self.refill();
self.tokens
}
pub fn time_until_refill(&self) -> Duration {
if self.tokens >= 1.0 {
Duration::ZERO
} else {
let needed = 1.0 - self.tokens;
let seconds = needed / self.refill_rate;
Duration::from_secs_f64(seconds)
}
}
fn refill(&mut self) {
let now = Instant::now();
let elapsed = now.duration_since(self.last_update);
let new_tokens = elapsed.as_secs_f64() * self.refill_rate;
self.tokens = (self.tokens + new_tokens).min(self.max_tokens);
self.last_update = now;
}
pub fn last_update(&self) -> Instant {
self.last_update
}
}
#[derive(Debug, Clone)]
pub enum RateLimitResult {
Allowed {
remaining: u32,
reset_at: u64,
},
Limited {
retry_after: u64,
},
}
pub struct RateLimiter {
buckets: DashMap<IpAddr, TokenBucket>,
config: RateLimitConfig,
#[allow(dead_code)]
start_time: Instant,
total_requests: AtomicU64,
rejected_requests: AtomicU64,
}
impl RateLimiter {
pub fn new(config: RateLimitConfig) -> Self {
Self {
buckets: DashMap::new(),
config,
start_time: Instant::now(),
total_requests: AtomicU64::new(0),
rejected_requests: AtomicU64::new(0),
}
}
pub fn check(&self, ip: IpAddr) -> RateLimitResult {
self.total_requests.fetch_add(1, Ordering::Relaxed);
if !self.config.enabled {
return RateLimitResult::Allowed {
remaining: u32::MAX,
reset_at: 0,
};
}
if self.config.whitelist.contains(&ip) {
return RateLimitResult::Allowed {
remaining: u32::MAX,
reset_at: 0,
};
}
let mut bucket = self.buckets.entry(ip).or_insert_with(|| {
let refill_rate = self.config.requests_per_minute as f64 / 60.0;
TokenBucket::new(self.config.burst_size as f64, refill_rate)
});
if bucket.try_consume() {
let remaining = bucket.tokens_remaining() as u32;
let reset_at = self.calculate_reset_time();
RateLimitResult::Allowed { remaining, reset_at }
} else {
self.rejected_requests.fetch_add(1, Ordering::Relaxed);
let retry_after = bucket.time_until_refill().as_secs().max(1);
RateLimitResult::Limited { retry_after }
}
}
fn calculate_reset_time(&self) -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default();
let current_secs = now.as_secs();
(current_secs / 60 + 1) * 60
}
pub fn cleanup_expired(&self, max_age: Duration) {
let now = Instant::now();
self.buckets.retain(|_: &IpAddr, bucket: &mut TokenBucket| {
now.duration_since(bucket.last_update()) < max_age
});
}
pub fn tracked_ips(&self) -> usize {
self.buckets.len()
}
pub fn total_requests(&self) -> u64 {
self.total_requests.load(Ordering::Relaxed)
}
pub fn rejected_requests(&self) -> u64 {
self.rejected_requests.load(Ordering::Relaxed)
}
pub fn is_enabled(&self) -> bool {
self.config.enabled
}
pub fn requests_per_minute(&self) -> u32 {
self.config.requests_per_minute
}
pub fn burst_size(&self) -> u32 {
self.config.burst_size
}
}
#[derive(Debug, Clone, Serialize)]
pub struct RateLimitStatus {
pub enabled: bool,
pub requests_per_minute: u32,
pub burst_size: u32,
pub your_remaining: u32,
pub reset_at: u64,
}
#[derive(Debug, Clone, Serialize)]
pub struct RateLimitError {
pub error: String,
pub message: String,
pub retry_after: u64,
}
impl RateLimitError {
pub fn new(retry_after: u64) -> Self {
Self {
error: "rate_limit_exceeded".to_string(),
message: format!(
"Too many requests. Please retry after {} seconds.",
retry_after
),
retry_after,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::Ipv4Addr;
use std::thread;
#[test]
fn test_token_bucket_new() {
let bucket = TokenBucket::new(10.0, 1.0);
assert_eq!(bucket.max_tokens, 10.0);
assert_eq!(bucket.tokens, 10.0);
assert_eq!(bucket.refill_rate, 1.0);
}
#[test]
fn test_token_bucket_consume_success() {
let mut bucket = TokenBucket::new(10.0, 1.0);
assert!(bucket.try_consume());
assert!(bucket.tokens_remaining() < 10.0);
}
#[test]
fn test_token_bucket_exhausted() {
let mut bucket = TokenBucket::new(2.0, 0.1);
assert!(bucket.try_consume()); assert!(bucket.try_consume()); assert!(!bucket.try_consume()); }
#[test]
fn test_token_bucket_refill() {
let mut bucket = TokenBucket::new(10.0, 100.0);
for _ in 0..10 {
bucket.try_consume();
}
assert!(bucket.tokens_remaining() < 1.0);
thread::sleep(Duration::from_millis(50));
assert!(bucket.tokens_remaining() > 0.0);
}
#[test]
fn test_rate_limiter_new() {
let config = RateLimitConfig::default();
let limiter = RateLimiter::new(config);
assert!(limiter.is_enabled());
assert_eq!(limiter.requests_per_minute(), 60);
assert_eq!(limiter.burst_size(), 10);
}
#[test]
fn test_rate_limiter_per_ip() {
let config = RateLimitConfig {
burst_size: 2,
..Default::default()
};
let limiter = RateLimiter::new(config);
let ip1: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
let ip2: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 2));
assert!(matches!(limiter.check(ip1), RateLimitResult::Allowed { .. }));
assert!(matches!(limiter.check(ip1), RateLimitResult::Allowed { .. }));
assert!(matches!(limiter.check(ip1), RateLimitResult::Limited { .. }));
assert!(matches!(limiter.check(ip2), RateLimitResult::Allowed { .. }));
assert_eq!(limiter.tracked_ips(), 2);
}
#[test]
fn test_rate_limiter_whitelist() {
let whitelisted_ip: IpAddr = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1));
let config = RateLimitConfig {
burst_size: 1,
whitelist: vec![whitelisted_ip],
..Default::default()
};
let limiter = RateLimiter::new(config);
for _ in 0..100 {
assert!(matches!(
limiter.check(whitelisted_ip),
RateLimitResult::Allowed { .. }
));
}
}
#[test]
fn test_rate_limit_remaining() {
let config = RateLimitConfig {
burst_size: 5,
..Default::default()
};
let limiter = RateLimiter::new(config);
let ip: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
if let RateLimitResult::Allowed { remaining, .. } = limiter.check(ip) {
assert!(remaining <= 4); } else {
panic!("Should be allowed");
}
}
#[test]
fn test_rate_limit_exceeded() {
let config = RateLimitConfig {
burst_size: 1,
requests_per_minute: 1,
..Default::default()
};
let limiter = RateLimiter::new(config);
let ip: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
assert!(matches!(limiter.check(ip), RateLimitResult::Allowed { .. }));
assert!(matches!(limiter.check(ip), RateLimitResult::Limited { .. }));
}
#[test]
fn test_rate_limit_retry_after() {
let config = RateLimitConfig {
burst_size: 1,
requests_per_minute: 60, ..Default::default()
};
let limiter = RateLimiter::new(config);
let ip: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
limiter.check(ip);
if let RateLimitResult::Limited { retry_after } = limiter.check(ip) {
assert!(retry_after >= 1); } else {
panic!("Should be limited");
}
}
#[test]
fn test_rate_limiter_concurrent() {
use std::sync::Arc;
let config = RateLimitConfig {
burst_size: 100,
..Default::default()
};
let limiter = Arc::new(RateLimiter::new(config));
let ip: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
let handles: Vec<_> = (0..10)
.map(|_| {
let limiter = Arc::clone(&limiter);
thread::spawn(move || {
for _ in 0..10 {
let _ = limiter.check(ip);
}
})
})
.collect();
for handle in handles {
handle.join().unwrap();
}
assert_eq!(limiter.total_requests(), 100);
}
#[test]
fn test_rate_limiter_cleanup() {
let config = RateLimitConfig::default();
let limiter = RateLimiter::new(config);
let ip1: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
let ip2: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 2));
limiter.check(ip1);
limiter.check(ip2);
assert_eq!(limiter.tracked_ips(), 2);
limiter.cleanup_expired(Duration::from_nanos(1));
assert_eq!(limiter.tracked_ips(), 0);
}
#[test]
fn test_rate_limit_config_default() {
let config = RateLimitConfig::default();
assert_eq!(config.requests_per_minute, 60);
assert_eq!(config.burst_size, 10);
assert_eq!(config.refill_interval_ms, 1000);
assert!(config.whitelist.is_empty());
assert!(config.enabled);
}
#[test]
fn test_rate_limit_disabled() {
let config = RateLimitConfig {
enabled: false,
..Default::default()
};
let limiter = RateLimiter::new(config);
let ip: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
for _ in 0..1000 {
assert!(matches!(
limiter.check(ip),
RateLimitResult::Allowed { remaining: u32::MAX, .. }
));
}
}
#[test]
fn test_rate_limit_error_new() {
let error = RateLimitError::new(60);
assert_eq!(error.error, "rate_limit_exceeded");
assert_eq!(error.retry_after, 60);
assert!(error.message.contains("60"));
}
#[test]
fn test_rate_limit_statistics() {
let config = RateLimitConfig {
burst_size: 2,
..Default::default()
};
let limiter = RateLimiter::new(config);
let ip: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
limiter.check(ip); limiter.check(ip); limiter.check(ip);
assert_eq!(limiter.total_requests(), 3);
assert_eq!(limiter.rejected_requests(), 1);
}
}