use std::fmt;
#[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, Clone, PartialEq, Eq)]
pub struct RateLimitError(pub String);
impl fmt::Display for RateLimitError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "rate limiter error: {}", self.0)
}
}
impl std::error::Error for RateLimitError {}
pub trait RateLimiter: Send + Sync {
fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rate_limit_result_allowed() {
let r = RateLimitResult::allowed(9, 1_700_000_000_000);
assert!(r.allowed);
assert_eq!(r.remaining, 9);
}
#[test]
fn test_rate_limit_result_rejected() {
let r = RateLimitResult::rejected(0, 1_700_000_000_000);
assert!(!r.allowed);
assert_eq!(r.remaining, 0);
}
#[test]
fn test_rate_limit_error_display() {
let e = RateLimitError("backend down".to_string());
assert!(e.to_string().contains("backend down"));
}
struct AlwaysAllow;
impl RateLimiter for AlwaysAllow {
fn try_acquire(&self, _key: &str) -> Result<RateLimitResult, RateLimitError> {
Ok(RateLimitResult::allowed(1, 0))
}
}
#[test]
fn test_trait_dispatch() {
let limiter: &dyn RateLimiter = &AlwaysAllow;
let r = limiter.try_acquire("test-key").unwrap();
assert!(r.allowed);
}
}