use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use crate::data::cache::RedisStorage;
#[async_trait::async_trait]
pub trait RateLimiter: Send + Sync {
async fn is_allowed(&self, key: &str, limit: u32) -> bool;
}
pub struct InMemoryRateLimiter {
inner: Mutex<HashMap<String, (u32, Instant)>>,
window: Duration,
}
impl InMemoryRateLimiter {
pub fn new() -> Self {
Self {
inner: Mutex::new(HashMap::new()),
window: Duration::from_secs(60),
}
}
pub fn new_with_cleanup() -> Arc<Self> {
let arc = Arc::new(Self {
inner: Mutex::new(HashMap::new()),
window: Duration::from_secs(60),
});
let weak = Arc::downgrade(&arc);
tokio::spawn(async move {
let mut interval = tokio::time::interval(Duration::from_secs(60));
loop {
interval.tick().await;
if let Some(limiter) = weak.upgrade() {
limiter.cleanup();
} else {
break;
}
}
});
arc
}
pub fn is_allowed_sync(&self, key: &str, limit: u32) -> bool {
let minute = chrono::Utc::now().timestamp() / 60;
let bucket_key = format!("{}:{}", key, minute);
let mut map = self.inner.lock().unwrap();
let now = Instant::now();
let entry = map.entry(bucket_key).or_insert((0, now + self.window));
if now > entry.1 {
*entry = (1, now + self.window);
return true;
}
if entry.0 < limit {
entry.0 += 1;
true
} else {
false
}
}
fn cleanup(&self) {
let now = Instant::now();
let mut map = self.inner.lock().unwrap();
map.retain(|_, (_, expiry)| *expiry > now);
}
}
impl Default for InMemoryRateLimiter {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
impl RateLimiter for InMemoryRateLimiter {
async fn is_allowed(&self, key: &str, limit: u32) -> bool {
self.is_allowed_sync(key, limit)
}
}
pub struct RedisBackedRateLimiter {
storage: Option<RedisStorage>,
fallback: Arc<InMemoryRateLimiter>,
}
impl RedisBackedRateLimiter {
pub fn new() -> Self {
let storage = RedisStorage::from_env().ok();
Self {
storage,
fallback: InMemoryRateLimiter::new_with_cleanup(),
}
}
pub fn with_storage(storage: RedisStorage) -> Self {
Self {
storage: Some(storage),
fallback: InMemoryRateLimiter::new_with_cleanup(),
}
}
pub fn from_url(url: &str) -> Self {
let storage = RedisStorage::new(url).ok();
Self {
storage,
fallback: InMemoryRateLimiter::new_with_cleanup(),
}
}
}
impl Default for RedisBackedRateLimiter {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
impl RateLimiter for RedisBackedRateLimiter {
async fn is_allowed(&self, key: &str, limit: u32) -> bool {
let Some(storage) = &self.storage else {
return self.fallback.is_allowed(key, limit).await;
};
let minute = chrono::Utc::now().timestamp() / 60;
let redis_key = format!("rate_limit:{}:{}", key, minute);
match storage.increment_value(&redis_key).await {
Ok(count) => {
if count == 1 {
let _ = storage.set_expiration(&redis_key, 60).await;
return true;
}
count <= limit as i64
}
Err(e) => {
tracing::warn!(backend = "redis", fallback_backend = "in_memory", error = %e, "Rate limiter failed; using fallback");
self.fallback.is_allowed(key, limit).await
}
}
}
}
pub enum LimiterKind {
InMemory,
Redis,
}
pub fn create(kind: LimiterKind) -> Arc<dyn RateLimiter> {
match kind {
LimiterKind::InMemory => InMemoryRateLimiter::new_with_cleanup(),
LimiterKind::Redis => Arc::new(RedisBackedRateLimiter::new()),
}
}
pub fn create_from_env() -> Arc<dyn RateLimiter> {
if std::env::var("REDIS_URL").is_ok() || std::env::var("REDIS_URI").is_ok() {
tracing::info!(component = "rate_limiter", backend = "redis", "Configured rate limiter");
Arc::new(RedisBackedRateLimiter::new())
} else {
tracing::info!(component = "rate_limiter", backend = "in_memory", "Configured rate limiter");
InMemoryRateLimiter::new_with_cleanup()
}
}