use async_trait::async_trait;
use super::{KeyedLimiter, build_keyed_limiter};
use crate::errors::OrionError;
#[async_trait]
pub trait RateLimitBackend: Send + Sync {
async fn check(&self, key: String) -> Result<bool, OrionError>;
}
pub struct LocalRateLimitBackend {
limiter: KeyedLimiter,
}
impl LocalRateLimitBackend {
pub fn new(rps: u32, burst: u32) -> Self {
Self {
limiter: build_keyed_limiter(rps, burst),
}
}
}
#[async_trait]
impl RateLimitBackend for LocalRateLimitBackend {
async fn check(&self, key: String) -> Result<bool, OrionError> {
Ok(self.limiter.check_key(&key).is_ok())
}
}
pub struct RedisRateLimitBackend {
conn: redis::aio::ConnectionManager,
scope: String,
limit_per_window: u32,
script: redis::Script,
}
const FIXED_WINDOW_SCRIPT: &str = r#"
local t = redis.call('TIME')
local key = KEYS[1] .. ':' .. t[1]
local count = redis.call('INCR', key)
redis.call('EXPIRE', key, 2)
return count
"#;
impl RedisRateLimitBackend {
pub fn new(conn: redis::aio::ConnectionManager, scope: String, rps: u32, burst: u32) -> Self {
Self {
conn,
scope,
limit_per_window: rps.saturating_add(burst).max(1),
script: redis::Script::new(FIXED_WINDOW_SCRIPT),
}
}
}
#[async_trait]
impl RateLimitBackend for RedisRateLimitBackend {
async fn check(&self, key: String) -> Result<bool, OrionError> {
let base_key = format!("orion:rl:{}:{}", self.scope, key);
let mut conn = self.conn.clone();
let result: Result<i64, redis::RedisError> =
self.script.key(&base_key).invoke_async(&mut conn).await;
match result {
Ok(count) => Ok(count <= i64::from(self.limit_per_window)),
Err(e) => Err(OrionError::internal(format!(
"rate-limit backend '{}': {e}",
self.scope
))),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_local_backend_enforces_burst() {
let backend = LocalRateLimitBackend::new(1, 2);
assert!(backend.check("ip-1".to_string()).await.expect("test"));
assert!(backend.check("ip-1".to_string()).await.expect("test"));
assert!(!backend.check("ip-1".to_string()).await.expect("test"));
assert!(backend.check("ip-2".to_string()).await.expect("test"));
}
}