use serde::{Deserialize, Serialize};
use std::collections::{HashMap, VecDeque};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
#[derive(Debug)]
pub struct TokenBucket {
capacity: u64,
tokens: f64,
fill_rate: f64, last_update: Instant,
}
impl TokenBucket {
pub fn new(capacity: u64, fill_rate: f64) -> Self {
Self {
capacity,
tokens: capacity as f64,
fill_rate,
last_update: Instant::now(),
}
}
pub fn try_consume(&mut self, tokens: u64) -> bool {
self.refill();
if self.tokens >= tokens as f64 {
self.tokens -= tokens as f64;
true
} else {
false
}
}
fn refill(&mut self) {
let now = Instant::now();
let elapsed = now.duration_since(self.last_update).as_secs_f64();
let new_tokens = elapsed * self.fill_rate;
self.tokens = (self.tokens + new_tokens).min(self.capacity as f64);
self.last_update = now;
}
pub fn available_tokens(&mut self) -> f64 {
self.refill();
self.tokens
}
pub fn reset(&mut self) {
self.tokens = self.capacity as f64;
self.last_update = Instant::now();
}
pub fn time_until_available(&mut self, tokens: u64) -> Option<Duration> {
self.refill();
if self.tokens >= tokens as f64 {
return Some(Duration::ZERO);
}
let tokens_needed = tokens as f64 - self.tokens;
let seconds = tokens_needed / self.fill_rate;
Some(Duration::from_secs_f64(seconds))
}
}
#[derive(Debug)]
pub struct LeakyBucket {
capacity: usize,
leak_rate: f64, queue: VecDeque<Instant>,
last_leak: Instant,
}
impl LeakyBucket {
pub fn new(capacity: usize, leak_rate: f64) -> Self {
Self {
capacity,
leak_rate,
queue: VecDeque::new(),
last_leak: Instant::now(),
}
}
pub fn try_acquire(&mut self) -> bool {
self.leak();
if self.queue.len() < self.capacity {
self.queue.push_back(Instant::now());
true
} else {
false
}
}
fn leak(&mut self) {
let now = Instant::now();
let elapsed = now.duration_since(self.last_leak).as_secs_f64();
let requests_to_leak = (elapsed * self.leak_rate) as usize;
for _ in 0..requests_to_leak.min(self.queue.len()) {
self.queue.pop_front();
}
self.last_leak = now;
}
pub fn queue_size(&mut self) -> usize {
self.leak();
self.queue.len()
}
pub fn reset(&mut self) {
self.queue.clear();
self.last_leak = Instant::now();
}
pub fn estimated_wait_time(&mut self) -> Duration {
self.leak();
if self.queue.is_empty() {
return Duration::ZERO;
}
let seconds = self.queue.len() as f64 / self.leak_rate;
Duration::from_secs_f64(seconds)
}
}
#[derive(Debug)]
pub struct SlidingWindow {
max_requests: u64,
window_size: Duration,
requests: VecDeque<Instant>,
}
impl SlidingWindow {
pub fn new(max_requests: u64, window_size: Duration) -> Self {
Self {
max_requests,
window_size,
requests: VecDeque::new(),
}
}
pub fn try_acquire(&mut self) -> bool {
self.cleanup();
if self.requests.len() < self.max_requests as usize {
self.requests.push_back(Instant::now());
true
} else {
false
}
}
fn cleanup(&mut self) {
let now = Instant::now();
let cutoff = now - self.window_size;
while let Some(&request_time) = self.requests.front() {
if request_time < cutoff {
self.requests.pop_front();
} else {
break;
}
}
}
pub fn current_count(&mut self) -> usize {
self.cleanup();
self.requests.len()
}
pub fn remaining(&mut self) -> u64 {
self.cleanup();
self.max_requests.saturating_sub(self.requests.len() as u64)
}
pub fn reset(&mut self) {
self.requests.clear();
}
pub fn time_until_available(&mut self) -> Duration {
self.cleanup();
if (self.requests.len() as u64) < self.max_requests {
return Duration::ZERO;
}
if let Some(&oldest) = self.requests.front() {
let elapsed = Instant::now().duration_since(oldest);
self.window_size.saturating_sub(elapsed)
} else {
Duration::ZERO
}
}
}
#[derive(Debug, Clone)]
pub struct DistributedRateLimiter {
inner: Arc<Mutex<DistributedLimiterState>>,
}
#[derive(Debug)]
struct DistributedLimiterState {
counters: HashMap<String, SlidingWindow>,
max_requests: u64,
window_size: Duration,
}
impl DistributedRateLimiter {
pub fn new(max_requests: u64, window_size: Duration) -> Self {
Self {
inner: Arc::new(Mutex::new(DistributedLimiterState {
counters: HashMap::new(),
max_requests,
window_size,
})),
}
}
pub fn try_acquire(&self, key: &str) -> bool {
let mut state = self.inner.lock().unwrap();
let max_requests = state.max_requests;
let window_size = state.window_size;
let limiter = state
.counters
.entry(key.to_string())
.or_insert_with(|| SlidingWindow::new(max_requests, window_size));
limiter.try_acquire()
}
pub fn current_count(&self, key: &str) -> usize {
let mut state = self.inner.lock().unwrap();
state
.counters
.get_mut(key)
.map(|l| l.current_count())
.unwrap_or(0)
}
pub fn remaining(&self, key: &str) -> u64 {
let mut state = self.inner.lock().unwrap();
state
.counters
.get_mut(key)
.map(|l| l.remaining())
.unwrap_or(state.max_requests)
}
pub fn reset_key(&self, key: &str) {
let mut state = self.inner.lock().unwrap();
state.counters.remove(key);
}
pub fn cleanup(&self) {
let mut state = self.inner.lock().unwrap();
state.counters.retain(|_, limiter| {
limiter.cleanup();
limiter.current_count() > 0
});
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RateLimitConfig {
pub algorithm: RateLimitAlgorithm,
pub max_requests: u64,
pub window_seconds: u64,
pub burst_size: Option<u64>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RateLimitAlgorithm {
TokenBucket,
LeakyBucket,
SlidingWindow,
Distributed,
}
impl RateLimitConfig {
pub fn token_bucket(capacity: u64, _fill_rate: f64) -> Self {
Self {
algorithm: RateLimitAlgorithm::TokenBucket,
max_requests: capacity,
window_seconds: 1,
burst_size: Some(capacity),
}
}
pub fn leaky_bucket(capacity: usize, _leak_rate: f64) -> Self {
Self {
algorithm: RateLimitAlgorithm::LeakyBucket,
max_requests: capacity as u64,
window_seconds: 1,
burst_size: None,
}
}
pub fn sliding_window(max_requests: u64, window_seconds: u64) -> Self {
Self {
algorithm: RateLimitAlgorithm::SlidingWindow,
max_requests,
window_seconds,
burst_size: None,
}
}
pub fn distributed(max_requests: u64, window_seconds: u64) -> Self {
Self {
algorithm: RateLimitAlgorithm::Distributed,
max_requests,
window_seconds,
burst_size: None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::thread;
#[test]
fn test_token_bucket_basic() {
let mut bucket = TokenBucket::new(10, 1.0);
assert!(bucket.try_consume(10));
assert!(!bucket.try_consume(1));
thread::sleep(Duration::from_secs(1));
assert!(bucket.try_consume(1)); }
#[test]
fn test_token_bucket_refill() {
let mut bucket = TokenBucket::new(10, 10.0);
assert!(bucket.try_consume(10));
thread::sleep(Duration::from_secs(1));
assert!(bucket.available_tokens() >= 9.0);
}
#[test]
fn test_token_bucket_partial_consume() {
let mut bucket = TokenBucket::new(100, 10.0);
assert!(bucket.try_consume(30));
assert!(bucket.try_consume(30));
assert!(bucket.try_consume(30));
assert!(!bucket.try_consume(20)); }
#[test]
fn test_leaky_bucket_basic() {
let mut bucket = LeakyBucket::new(10, 1.0);
for _ in 0..10 {
assert!(bucket.try_acquire());
}
assert!(!bucket.try_acquire());
}
#[test]
fn test_leaky_bucket_leak() {
let mut bucket = LeakyBucket::new(10, 10.0);
for _ in 0..10 {
assert!(bucket.try_acquire());
}
thread::sleep(Duration::from_millis(500));
assert!(bucket.queue_size() < 10);
}
#[test]
fn test_sliding_window_basic() {
let mut window = SlidingWindow::new(5, Duration::from_secs(1));
for _ in 0..5 {
assert!(window.try_acquire());
}
assert!(!window.try_acquire());
assert_eq!(window.current_count(), 5);
}
#[test]
fn test_sliding_window_expiration() {
let mut window = SlidingWindow::new(5, Duration::from_millis(100));
for _ in 0..5 {
assert!(window.try_acquire());
}
thread::sleep(Duration::from_millis(150));
assert!(window.try_acquire());
assert_eq!(window.current_count(), 1);
}
#[test]
fn test_sliding_window_remaining() {
let mut window = SlidingWindow::new(10, Duration::from_secs(1));
assert_eq!(window.remaining(), 10);
window.try_acquire();
window.try_acquire();
assert_eq!(window.remaining(), 8);
}
#[test]
fn test_distributed_limiter() {
let limiter = DistributedRateLimiter::new(5, Duration::from_secs(1));
assert!(limiter.try_acquire("user1"));
assert!(limiter.try_acquire("user2"));
assert_eq!(limiter.current_count("user1"), 1);
assert_eq!(limiter.current_count("user2"), 1);
}
#[test]
fn test_distributed_limiter_per_key_limit() {
let limiter = DistributedRateLimiter::new(3, Duration::from_secs(1));
for _ in 0..3 {
assert!(limiter.try_acquire("user1"));
}
assert!(!limiter.try_acquire("user1"));
assert!(limiter.try_acquire("user2"));
}
#[test]
fn test_distributed_limiter_reset() {
let limiter = DistributedRateLimiter::new(2, Duration::from_secs(1));
limiter.try_acquire("user1");
limiter.try_acquire("user1");
assert!(!limiter.try_acquire("user1"));
limiter.reset_key("user1");
assert!(limiter.try_acquire("user1"));
}
#[test]
fn test_rate_limit_config() {
let token_config = RateLimitConfig::token_bucket(100, 10.0);
assert!(matches!(
token_config.algorithm,
RateLimitAlgorithm::TokenBucket
));
let sliding_config = RateLimitConfig::sliding_window(100, 60);
assert!(matches!(
sliding_config.algorithm,
RateLimitAlgorithm::SlidingWindow
));
}
#[test]
fn test_token_bucket_time_until_available() {
let mut bucket = TokenBucket::new(10, 10.0);
bucket.try_consume(10);
let wait_time = bucket.time_until_available(5).unwrap();
assert!(wait_time >= Duration::from_millis(400));
assert!(wait_time <= Duration::from_millis(600));
}
#[test]
fn test_leaky_bucket_estimated_wait() {
let mut bucket = LeakyBucket::new(10, 10.0);
for _ in 0..10 {
bucket.try_acquire();
}
let wait = bucket.estimated_wait_time();
assert!(wait >= Duration::from_millis(900));
}
}