use std::time::Duration;
use crate::{
FixedWindowRateLimiter, LeakyBucketLimiter, SlidingWindowLogLimiter, SlidingWindowRateLimiter,
TokenBucketRateLimiter,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum LimitAlgorithm {
FixedWindow,
SlidingWindow,
SlidingWindowLog,
TokenBucket,
LeakyBucket,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct RateLimitConfig {
pub algorithm: LimitAlgorithm,
pub capacity: u64,
pub rate: f64,
pub window_ms: u64,
pub max_keys: usize,
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
algorithm: LimitAlgorithm::TokenBucket,
capacity: 100,
rate: 10.0,
window_ms: 1000,
max_keys: crate::DEFAULT_MAX_KEYS,
}
}
}
impl RateLimitConfig {
pub fn new() -> Self {
Self::default()
}
pub fn from_json_str(s: &str) -> Result<Self, ConfigError> {
serde_json::from_str(s).map_err(|e| ConfigError::Parse(e.to_string()))
}
pub fn to_json_string(&self) -> Result<String, ConfigError> {
serde_json::to_string(self).map_err(|e| ConfigError::Serialize(e.to_string()))
}
pub fn validate(&self) -> Result<(), ConfigError> {
if self.capacity == 0 {
return Err(ConfigError::InvalidCapacity);
}
match self.algorithm {
LimitAlgorithm::TokenBucket | LimitAlgorithm::LeakyBucket => {
if self.rate < 0.0 {
return Err(ConfigError::InvalidRate);
}
}
_ => {}
}
match self.algorithm {
LimitAlgorithm::FixedWindow
| LimitAlgorithm::SlidingWindow
| LimitAlgorithm::SlidingWindowLog => {
if self.window_ms == 0 {
return Err(ConfigError::InvalidWindow);
}
}
_ => {}
}
if self.max_keys == 0 {
return Err(ConfigError::InvalidMaxKeys);
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ConfigError {
Parse(String),
Serialize(String),
InvalidCapacity,
InvalidRate,
InvalidWindow,
InvalidMaxKeys,
}
impl std::fmt::Display for ConfigError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ConfigError::Parse(msg) => write!(f, "config parse error: {}", msg),
ConfigError::Serialize(msg) => write!(f, "config serialize error: {}", msg),
ConfigError::InvalidCapacity => write!(f, "capacity must be positive"),
ConfigError::InvalidRate => write!(f, "rate must be non-negative"),
ConfigError::InvalidWindow => write!(f, "window size must be positive"),
ConfigError::InvalidMaxKeys => write!(f, "max_keys must be positive"),
}
}
}
impl std::error::Error for ConfigError {}
pub struct RateLimitConfigBuilder {
config: RateLimitConfig,
}
impl Default for RateLimitConfigBuilder {
fn default() -> Self {
Self::new()
}
}
impl RateLimitConfigBuilder {
pub fn new() -> Self {
Self {
config: RateLimitConfig::default(),
}
}
pub fn algorithm(mut self, algo: LimitAlgorithm) -> Self {
self.config.algorithm = algo;
self
}
pub fn capacity(mut self, cap: u64) -> Self {
self.config.capacity = cap;
self
}
pub fn rate(mut self, rate: f64) -> Self {
self.config.rate = rate;
self
}
pub fn window_ms(mut self, ms: u64) -> Self {
self.config.window_ms = ms;
self
}
pub fn window_secs(mut self, secs: u64) -> Self {
self.config.window_ms = secs * 1000;
self
}
pub fn max_keys(mut self, max: usize) -> Self {
self.config.max_keys = max;
self
}
pub fn build(self) -> RateLimitConfig {
self.config
}
pub fn build_checked(self) -> Result<RateLimitConfig, ConfigError> {
self.config.validate()?;
Ok(self.config)
}
pub fn build_token_bucket(&self) -> TokenBucketRateLimiter {
TokenBucketRateLimiter::new(self.config.capacity, self.config.rate)
.with_max_keys(self.config.max_keys)
}
pub fn build_sliding_window(&self) -> SlidingWindowRateLimiter {
SlidingWindowRateLimiter::new(
self.config.capacity,
Duration::from_millis(self.config.window_ms),
)
.with_max_keys(self.config.max_keys)
}
pub fn build_fixed_window(&self) -> FixedWindowRateLimiter {
FixedWindowRateLimiter::new(
self.config.capacity,
Duration::from_millis(self.config.window_ms),
)
.with_max_keys(self.config.max_keys)
}
pub fn build_leaky_bucket(&self) -> LeakyBucketLimiter {
LeakyBucketLimiter::new(self.config.capacity, self.config.rate)
.with_max_keys(self.config.max_keys)
}
pub fn build_sliding_window_log(&self) -> SlidingWindowLogLimiter {
SlidingWindowLogLimiter::new(
self.config.capacity,
Duration::from_millis(self.config.window_ms),
)
.with_max_keys(self.config.max_keys)
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct TieredRateLimitConfig {
pub ip: RateLimitConfig,
pub user: RateLimitConfig,
pub api: RateLimitConfig,
pub global: RateLimitConfig,
}
impl Default for TieredRateLimitConfig {
fn default() -> Self {
Self {
ip: RateLimitConfig {
algorithm: LimitAlgorithm::SlidingWindow,
capacity: 1000,
rate: 0.0,
window_ms: 60_000,
max_keys: 10_000,
},
user: RateLimitConfig {
algorithm: LimitAlgorithm::TokenBucket,
capacity: 100,
rate: 10.0,
window_ms: 1000,
max_keys: 10_000,
},
api: RateLimitConfig {
algorithm: LimitAlgorithm::FixedWindow,
capacity: 500,
rate: 0.0,
window_ms: 60_000,
max_keys: 1000,
},
global: RateLimitConfig {
algorithm: LimitAlgorithm::TokenBucket,
capacity: 10_000,
rate: 100.0,
window_ms: 1000,
max_keys: 1,
},
}
}
}
impl TieredRateLimitConfig {
pub fn new() -> Self {
Self::default()
}
pub fn validate(&self) -> Result<(), ConfigError> {
self.ip.validate()?;
self.user.validate()?;
self.api.validate()?;
self.global.validate()?;
Ok(())
}
pub fn from_json_str(s: &str) -> Result<Self, ConfigError> {
serde_json::from_str(s).map_err(|e| ConfigError::Parse(e.to_string()))
}
pub fn to_json_string(&self) -> Result<String, ConfigError> {
serde_json::to_string(self).map_err(|e| ConfigError::Serialize(e.to_string()))
}
}
pub struct TieredConfigBuilder {
config: TieredRateLimitConfig,
}
impl Default for TieredConfigBuilder {
fn default() -> Self {
Self::new()
}
}
impl TieredConfigBuilder {
pub fn new() -> Self {
Self {
config: TieredRateLimitConfig::default(),
}
}
pub fn ip(mut self, config: RateLimitConfig) -> Self {
self.config.ip = config;
self
}
pub fn user(mut self, config: RateLimitConfig) -> Self {
self.config.user = config;
self
}
pub fn api(mut self, config: RateLimitConfig) -> Self {
self.config.api = config;
self
}
pub fn global(mut self, config: RateLimitConfig) -> Self {
self.config.global = config;
self
}
pub fn build(self) -> TieredRateLimitConfig {
self.config
}
pub fn build_checked(self) -> Result<TieredRateLimitConfig, ConfigError> {
self.config.validate()?;
Ok(self.config)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::RateLimiter;
#[test]
fn test_limit_algorithm_serde() {
let algo = LimitAlgorithm::TokenBucket;
let json = serde_json::to_string(&algo).unwrap();
let back: LimitAlgorithm = serde_json::from_str(&json).unwrap();
assert_eq!(algo, back);
}
#[test]
fn test_rate_limit_config_default() {
let config = RateLimitConfig::default();
assert_eq!(config.algorithm, LimitAlgorithm::TokenBucket);
assert_eq!(config.capacity, 100);
assert_eq!(config.rate, 10.0);
}
#[test]
fn test_rate_limit_config_validate_ok() {
let config = RateLimitConfig {
algorithm: LimitAlgorithm::TokenBucket,
capacity: 100,
rate: 10.0,
window_ms: 1000,
max_keys: 1000,
};
assert!(config.validate().is_ok());
}
#[test]
fn test_rate_limit_config_validate_zero_capacity() {
let config = RateLimitConfig {
capacity: 0,
..Default::default()
};
assert_eq!(config.validate(), Err(ConfigError::InvalidCapacity));
}
#[test]
fn test_rate_limit_config_validate_negative_rate() {
let config = RateLimitConfig {
algorithm: LimitAlgorithm::TokenBucket,
rate: -1.0,
..Default::default()
};
assert_eq!(config.validate(), Err(ConfigError::InvalidRate));
}
#[test]
fn test_rate_limit_config_validate_zero_window() {
let config = RateLimitConfig {
algorithm: LimitAlgorithm::FixedWindow,
window_ms: 0,
..Default::default()
};
assert_eq!(config.validate(), Err(ConfigError::InvalidWindow));
}
#[test]
fn test_rate_limit_config_validate_zero_max_keys() {
let config = RateLimitConfig {
max_keys: 0,
..Default::default()
};
assert_eq!(config.validate(), Err(ConfigError::InvalidMaxKeys));
}
#[test]
fn test_rate_limit_config_json_roundtrip() {
let config = RateLimitConfig::default();
let json = config.to_json_string().unwrap();
let back = RateLimitConfig::from_json_str(&json).unwrap();
assert_eq!(config.algorithm, back.algorithm);
assert_eq!(config.capacity, back.capacity);
}
#[test]
fn test_config_builder_basic() {
let config = RateLimitConfigBuilder::new()
.algorithm(LimitAlgorithm::SlidingWindow)
.capacity(200)
.window_secs(60)
.max_keys(5000)
.build();
assert_eq!(config.algorithm, LimitAlgorithm::SlidingWindow);
assert_eq!(config.capacity, 200);
assert_eq!(config.window_ms, 60_000);
assert_eq!(config.max_keys, 5000);
}
#[test]
fn test_config_builder_checked_ok() {
let config = RateLimitConfigBuilder::new()
.algorithm(LimitAlgorithm::TokenBucket)
.capacity(100)
.rate(10.0)
.build_checked();
assert!(config.is_ok());
}
#[test]
fn test_config_builder_checked_fail() {
let config = RateLimitConfigBuilder::new().capacity(0).build_checked();
assert!(config.is_err());
}
#[test]
fn test_config_builder_build_token_bucket() {
let builder = RateLimitConfigBuilder::new()
.algorithm(LimitAlgorithm::TokenBucket)
.capacity(10)
.rate(1.0)
.max_keys(100);
let limiter = builder.build_token_bucket();
assert_eq!(limiter.capacity(), 10);
let r = limiter.acquire("k").unwrap();
assert!(r.allowed);
}
#[test]
fn test_config_builder_build_sliding_window() {
let builder = RateLimitConfigBuilder::new()
.algorithm(LimitAlgorithm::SlidingWindow)
.capacity(10)
.window_secs(60);
let limiter = builder.build_sliding_window();
assert_eq!(limiter.max_requests(), 10);
let r = limiter.acquire("k").unwrap();
assert!(r.allowed);
}
#[test]
fn test_config_builder_build_fixed_window() {
let builder = RateLimitConfigBuilder::new()
.algorithm(LimitAlgorithm::FixedWindow)
.capacity(10)
.window_secs(60);
let limiter = builder.build_fixed_window();
assert_eq!(limiter.max_requests(), 10);
let r = limiter.acquire("k").unwrap();
assert!(r.allowed);
}
#[test]
fn test_config_builder_build_leaky_bucket() {
let builder = RateLimitConfigBuilder::new()
.algorithm(LimitAlgorithm::LeakyBucket)
.capacity(10)
.rate(1.0);
let limiter = builder.build_leaky_bucket();
assert_eq!(limiter.capacity(), 10);
let r = limiter.acquire("k").unwrap();
assert!(r.allowed);
}
#[test]
fn test_config_builder_build_sliding_window_log() {
let builder = RateLimitConfigBuilder::new()
.algorithm(LimitAlgorithm::SlidingWindowLog)
.capacity(10)
.window_secs(60);
let limiter = builder.build_sliding_window_log();
assert_eq!(limiter.max_requests(), 10);
let r = limiter.acquire("k").unwrap();
assert!(r.allowed);
}
#[test]
fn test_tiered_config_default() {
let config = TieredRateLimitConfig::default();
assert_eq!(config.ip.capacity, 1000);
assert_eq!(config.user.capacity, 100);
assert_eq!(config.api.capacity, 500);
assert_eq!(config.global.capacity, 10_000);
}
#[test]
fn test_tiered_config_validate() {
let config = TieredRateLimitConfig::default();
assert!(config.validate().is_ok());
}
#[test]
fn test_tiered_config_json_roundtrip() {
let config = TieredRateLimitConfig::default();
let json = config.to_json_string().unwrap();
let back = TieredRateLimitConfig::from_json_str(&json).unwrap();
assert_eq!(back.ip.capacity, config.ip.capacity);
}
#[test]
fn test_tiered_config_builder() {
let config = TieredConfigBuilder::new()
.ip(RateLimitConfig {
capacity: 2000,
..Default::default()
})
.user(RateLimitConfig {
capacity: 200,
..Default::default()
})
.build();
assert_eq!(config.ip.capacity, 2000);
assert_eq!(config.user.capacity, 200);
}
#[test]
fn test_tiered_config_builder_checked() {
let config = TieredConfigBuilder::new().build_checked();
assert!(config.is_ok());
}
#[test]
fn test_config_builder_window_ms() {
let config = RateLimitConfigBuilder::new().window_ms(500).build();
assert_eq!(config.window_ms, 500);
}
#[test]
fn test_config_builder_rate() {
let config = RateLimitConfigBuilder::new().rate(5.0).build();
assert_eq!(config.rate, 5.0);
}
}