#[derive(Debug, Clone)]
pub struct RateLimitResult {
pub allowed: bool,
pub remaining: u64,
pub reset_at: i64,
}
impl RateLimitResult {
pub fn allowed(remaining: u64, reset_at: i64) -> Self {
Self {
allowed: true,
remaining,
reset_at,
}
}
pub fn rejected(remaining: u64, reset_at: i64) -> Self {
Self {
allowed: false,
remaining,
reset_at,
}
}
}
#[derive(Debug)]
pub enum RateLimitError {
KeyNotFound(String),
Internal(String),
}
impl std::fmt::Display for RateLimitError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
RateLimitError::KeyNotFound(key) => write!(f, "Key not found: {}", key),
RateLimitError::Internal(msg) => write!(f, "Internal error: {}", msg),
}
}
}
impl std::error::Error for RateLimitError {}
impl serde::Serialize for RateLimitError {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.serialize_str(&self.to_string())
}
}
pub trait RateLimiter: Send + Sync {
fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
self.acquire(key)
}
fn reset(&self, key: &str) -> Result<(), RateLimitError>;
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
#[test]
fn test_rate_limit_result_allowed() {
let result = RateLimitResult::allowed(5, 1000);
assert!(result.allowed);
assert_eq!(result.remaining, 5);
assert_eq!(result.reset_at, 1000);
}
#[test]
fn test_rate_limit_result_rejected() {
let result = RateLimitResult::rejected(0, 2000);
assert!(!result.allowed);
assert_eq!(result.remaining, 0);
}
#[test]
fn test_rate_limit_error_display() {
let e = RateLimitError::KeyNotFound("user-1".to_string());
assert!(format!("{}", e).contains("user-1"));
let e = RateLimitError::Internal("lock poisoned".to_string());
assert!(format!("{}", e).contains("lock poisoned"));
}
struct CounterLimiter {
max: u64,
count: AtomicU64,
}
impl CounterLimiter {
fn new(max: u64) -> Self {
Self {
max,
count: AtomicU64::new(0),
}
}
}
impl RateLimiter for CounterLimiter {
fn acquire(&self, _key: &str) -> Result<RateLimitResult, RateLimitError> {
let prev = self.count.fetch_add(1, Ordering::SeqCst);
if prev < self.max {
Ok(RateLimitResult::allowed(self.max - prev - 1, 0))
} else {
Ok(RateLimitResult::rejected(0, 0))
}
}
fn reset(&self, _key: &str) -> Result<(), RateLimitError> {
self.count.store(0, Ordering::SeqCst);
Ok(())
}
}
#[test]
fn test_counter_limiter_via_trait() {
let limiter: Arc<dyn RateLimiter> = Arc::new(CounterLimiter::new(2));
let r1 = limiter.acquire("k").unwrap();
assert!(r1.allowed);
let r2 = limiter.acquire("k").unwrap();
assert!(r2.allowed);
let r3 = limiter.acquire("k").unwrap();
assert!(!r3.allowed);
}
#[test]
fn test_counter_limiter_reset() {
let limiter: Arc<dyn RateLimiter> = Arc::new(CounterLimiter::new(1));
assert!(limiter.acquire("k").unwrap().allowed);
assert!(!limiter.acquire("k").unwrap().allowed);
limiter.reset("k").unwrap();
assert!(limiter.acquire("k").unwrap().allowed);
}
#[test]
fn test_default_try_acquire_delegates_to_acquire() {
let limiter: Arc<dyn RateLimiter> = Arc::new(CounterLimiter::new(1));
let r1 = limiter.try_acquire("k").unwrap();
assert!(r1.allowed);
let r2 = limiter.try_acquire("k").unwrap();
assert!(!r2.allowed);
}
#[test]
fn test_rate_limiter_is_send_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Arc<dyn RateLimiter>>();
}
}