use std::collections::HashMap;
use std::sync::Arc;
use limiteron::limiters::{Limiter, TokenBucketLimiter};
use tokio::sync::Mutex;
use crate::rate_limit::{RateLimitConfig, RateLimitDimension, RateLimitStatus};
pub struct LimiteronAdapter {
buckets: Mutex<HashMap<String, Arc<TokenBucketLimiter>>>,
config: RateLimitConfig,
}
impl LimiteronAdapter {
pub fn new(config: RateLimitConfig) -> Self {
Self {
buckets: Mutex::new(HashMap::new()),
config,
}
}
pub fn with_default_config() -> Self {
Self::new(RateLimitConfig::default())
}
pub async fn check_rate_limit(&self, dimensions: Vec<RateLimitDimension>) -> bool {
for dimension in dimensions {
let key = dimension_to_key(&dimension);
let max_requests = dimension_limit(&self.config, &dimension);
if max_requests == 0 {
return false;
}
let bucket = self.get_or_create_bucket(&key, max_requests).await;
match bucket.allow(1).await {
Ok(true) => continue,
_ => return false,
}
}
true
}
pub async fn get_status(&self, dimension: RateLimitDimension) -> RateLimitStatus {
let key = dimension_to_key(&dimension);
let max_requests = dimension_limit(&self.config, &dimension);
let bucket = self.get_or_create_bucket(&key, max_requests).await;
let tokens = bucket.tokens();
RateLimitStatus {
dimension: format!("{:?}", dimension),
max_requests,
current_count: max_requests.saturating_sub(tokens),
remaining: tokens,
window_secs: 60, algorithm: "token_bucket".to_string(),
}
}
pub async fn get_remaining(&self, dimension: RateLimitDimension) -> u64 {
self.get_status(dimension).await.remaining
}
async fn get_or_create_bucket(&self, key: &str, max_requests: u64) -> Arc<TokenBucketLimiter> {
let mut buckets = self.buckets.lock().await;
buckets
.entry(key.to_string())
.or_insert_with(|| {
let refill_rate = (max_requests / 60).max(1);
Arc::new(TokenBucketLimiter::new(max_requests, refill_rate))
})
.clone()
}
}
fn dimension_to_key(dimension: &RateLimitDimension) -> String {
match dimension {
RateLimitDimension::Global => "global".to_string(),
RateLimitDimension::Ip(ip) => format!("ip:{}", ip),
RateLimitDimension::User(u) => format!("user:{}", u),
RateLimitDimension::ApiKey(k) => format!("apikey:{}", k),
}
}
fn dimension_limit(config: &RateLimitConfig, dimension: &RateLimitDimension) -> u64 {
match dimension {
RateLimitDimension::Global => config.global_requests_per_minute,
RateLimitDimension::Ip(_) => config.ip_requests_per_minute,
RateLimitDimension::User(_) => config.user_requests_per_minute,
RateLimitDimension::ApiKey(_) => config.api_key_requests_per_minute,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn small_config() -> RateLimitConfig {
RateLimitConfig {
global_requests_per_minute: 5,
ip_requests_per_minute: 3,
user_requests_per_minute: 4,
api_key_requests_per_minute: 2,
..RateLimitConfig::default()
}
}
#[tokio::test]
async fn test_check_rate_limit_allows_under_limit() {
let adapter = LimiteronAdapter::new(small_config());
for _ in 0..5 {
assert!(
adapter
.check_rate_limit(vec![RateLimitDimension::Global])
.await
);
}
}
#[tokio::test]
async fn test_check_rate_limit_blocks_over_limit() {
let adapter = LimiteronAdapter::new(small_config());
assert!(
adapter
.check_rate_limit(vec![RateLimitDimension::ApiKey("k1".into())])
.await
);
assert!(
adapter
.check_rate_limit(vec![RateLimitDimension::ApiKey("k1".into())])
.await
);
assert!(
!adapter
.check_rate_limit(vec![RateLimitDimension::ApiKey("k1".into())])
.await
);
}
#[tokio::test]
async fn test_multiple_dimensions_independent() {
let adapter = LimiteronAdapter::new(small_config());
let ip = RateLimitDimension::Ip("1.2.3.4".into());
let user = RateLimitDimension::User("alice".into());
for _ in 0..3 {
assert!(adapter.check_rate_limit(vec![ip.clone()]).await);
}
assert!(!adapter.check_rate_limit(vec![ip.clone()]).await);
assert!(adapter.check_rate_limit(vec![user.clone()]).await);
}
#[tokio::test]
async fn test_check_rate_limit_all_dimensions_must_pass() {
let adapter = LimiteronAdapter::new(small_config());
let api_key = RateLimitDimension::ApiKey("key".into());
assert!(adapter.check_rate_limit(vec![api_key.clone()]).await);
assert!(adapter.check_rate_limit(vec![api_key.clone()]).await);
assert!(
!adapter
.check_rate_limit(vec![RateLimitDimension::Global, api_key.clone()])
.await
);
}
#[tokio::test]
async fn test_get_status_returns_correct_info() {
let adapter = LimiteronAdapter::new(small_config());
let dim = RateLimitDimension::Ip("10.0.0.1".into());
assert!(adapter.check_rate_limit(vec![dim.clone()]).await);
let status = adapter.get_status(dim.clone()).await;
assert_eq!(status.max_requests, 3);
assert_eq!(status.remaining, 2);
assert_eq!(status.current_count, 1);
assert_eq!(status.algorithm, "token_bucket");
assert!(status.dimension.contains("Ip"));
}
#[tokio::test]
async fn test_get_remaining_returns_tokens() {
let adapter = LimiteronAdapter::new(small_config());
let dim = RateLimitDimension::User("bob".into());
assert_eq!(adapter.get_remaining(dim.clone()).await, 4);
adapter.check_rate_limit(vec![dim.clone()]).await;
adapter.check_rate_limit(vec![dim.clone()]).await;
assert_eq!(adapter.get_remaining(dim.clone()).await, 2);
}
#[tokio::test]
async fn test_global_dimension_key_isolation() {
let adapter = LimiteronAdapter::new(small_config());
for _ in 0..5 {
assert!(
adapter
.check_rate_limit(vec![RateLimitDimension::Global])
.await
);
}
assert!(
!adapter
.check_rate_limit(vec![RateLimitDimension::Global])
.await
);
assert!(
adapter
.check_rate_limit(vec![RateLimitDimension::Ip("9.9.9.9".into())])
.await
);
}
#[tokio::test]
async fn test_token_refill_restores_capacity() {
let adapter = LimiteronAdapter::new(small_config());
let dim = RateLimitDimension::ApiKey("refill".into());
adapter.check_rate_limit(vec![dim.clone()]).await;
adapter.check_rate_limit(vec![dim.clone()]).await;
assert!(!adapter.check_rate_limit(vec![dim.clone()]).await);
tokio::time::sleep(tokio::time::Duration::from_millis(1200)).await;
assert!(
adapter.check_rate_limit(vec![dim.clone()]).await,
"expected refill to allow request after wait"
);
}
#[tokio::test]
async fn test_limit_zero_rejects_all() {
let mut config = small_config();
config.global_requests_per_minute = 0;
let adapter = LimiteronAdapter::new(config);
assert!(
!adapter
.check_rate_limit(vec![RateLimitDimension::Global])
.await,
"limit=0 must reject all requests"
);
}
#[tokio::test]
async fn test_window_secs_is_60() {
let adapter = LimiteronAdapter::new(small_config());
let status = adapter.get_status(RateLimitDimension::Global).await;
assert_eq!(
status.window_secs, 60,
"window_secs must be 60 (per-minute window)"
);
}
}