use redis::aio::ConnectionManager;
use crate::config::BotConfig;
use crate::error::{BotError, Result};
pub const DEFAULT_LIMIT: u64 = 30;
pub const DEFAULT_WINDOW_SECS: u64 = 60;
#[derive(Clone)]
pub struct RateLimiter {
conn: ConnectionManager,
limit: u64,
window_secs: u64,
}
impl RateLimiter {
pub async fn connect(config: &BotConfig) -> Result<Self> {
let client = redis::Client::open(config.redis_url.as_str())?;
let conn = ConnectionManager::new(client).await?;
Ok(Self {
conn,
limit: DEFAULT_LIMIT,
window_secs: DEFAULT_WINDOW_SECS,
})
}
fn user_key(fichub_user_id: i64) -> String {
format!("archivist:ratelimit:{fichub_user_id}")
}
fn platform_key(platform: &str, ext_id: &str) -> String {
format!("archivist:ratelimit:{platform}:{ext_id}")
}
pub async fn check(&self, platform: &str, ext_id: &str) -> Result<()> {
let key = Self::platform_key(platform, ext_id);
self.check_key(&key).await
}
pub async fn check_user(&self, platform: &str, ext_id: &str, fichub_user_id: i64) -> Result<()> {
let key = Self::user_key(fichub_user_id);
self.check_key(&key).await?;
self.check_key(&Self::platform_key(platform, ext_id)).await
}
async fn check_key(&self, key: &str) -> Result<()> {
let count: u64 = redis::cmd("INCR")
.arg(key)
.query_async(&mut self.conn.clone())
.await?;
if count == 1 {
let _: () = redis::cmd("EXPIRE")
.arg(key)
.arg(self.window_secs)
.query_async(&mut self.conn.clone())
.await?;
}
if count > self.limit {
let ttl: i64 = redis::cmd("TTL")
.arg(key)
.query_async(&mut self.conn.clone())
.await
.unwrap_or(self.window_secs as i64);
return Err(BotError::RateLimited {
retry_after_secs: ttl.max(1) as u64,
});
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn user_key_format() {
assert_eq!(
RateLimiter::user_key(42),
"archivist:ratelimit:42"
);
}
#[test]
fn platform_key_format() {
assert_eq!(
RateLimiter::platform_key("matrix", "@u:server"),
"archivist:ratelimit:matrix:@u:server"
);
assert_eq!(
RateLimiter::platform_key("telegram", "123"),
"archivist:ratelimit:telegram:123"
);
}
}