use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tracing::{debug, trace, warn};
use crate::errors::{LimitType, ZentinelError, ZentinelResult};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Limits {
pub max_header_size_bytes: usize,
pub max_header_count: usize,
pub max_header_name_bytes: usize,
pub max_header_value_bytes: usize,
pub max_body_size_bytes: usize,
pub max_body_buffer_bytes: usize,
pub max_body_inspection_bytes: usize,
pub max_decompression_ratio: f32,
pub max_decompressed_size_bytes: usize,
pub max_connections_per_client: usize,
pub max_connections_per_route: usize,
pub max_total_connections: usize,
pub max_idle_connections_per_upstream: usize,
pub max_in_flight_requests: usize,
pub max_in_flight_requests_per_worker: usize,
pub max_queued_requests: usize,
pub max_agent_queue_depth: usize,
pub max_agent_body_bytes: usize,
pub max_agent_response_bytes: usize,
pub max_requests_per_second_global: Option<u32>,
pub max_requests_per_second_per_client: Option<u32>,
pub max_requests_per_second_per_route: Option<u32>,
pub max_memory_bytes: Option<usize>,
pub max_memory_percent: Option<f32>,
}
impl Default for Limits {
fn default() -> Self {
Self {
max_header_size_bytes: 8192, max_header_count: 100, max_header_name_bytes: 256, max_header_value_bytes: 4096,
max_body_size_bytes: 10 * 1024 * 1024,
max_body_buffer_bytes: 1024 * 1024,
max_body_inspection_bytes: 1024 * 1024,
max_decompression_ratio: 100.0,
max_decompressed_size_bytes: 100 * 1024 * 1024,
max_connections_per_client: 100,
max_connections_per_route: 1000,
max_total_connections: 10000,
max_idle_connections_per_upstream: 100,
max_in_flight_requests: 10000,
max_in_flight_requests_per_worker: 1000,
max_queued_requests: 1000,
max_agent_queue_depth: 100,
max_agent_body_bytes: 1024 * 1024, max_agent_response_bytes: 10 * 1024,
max_requests_per_second_global: None,
max_requests_per_second_per_client: None,
max_requests_per_second_per_route: None,
max_memory_bytes: None,
max_memory_percent: None,
}
}
}
impl Limits {
pub fn for_testing() -> Self {
Self {
max_header_size_bytes: 16384,
max_header_count: 200,
max_body_size_bytes: 100 * 1024 * 1024, max_in_flight_requests: 100000,
..Default::default()
}
}
pub fn for_production() -> Self {
Self {
max_header_size_bytes: 4096,
max_header_count: 50,
max_body_size_bytes: 1024 * 1024, max_in_flight_requests: 5000,
max_requests_per_second_global: Some(10000),
max_requests_per_second_per_client: Some(100),
max_memory_percent: Some(80.0),
..Default::default()
}
}
pub fn validate(&self) -> ZentinelResult<()> {
if self.max_header_size_bytes == 0 {
return Err(ZentinelError::Config {
message: "max_header_size_bytes must be greater than 0".to_string(),
source: None,
});
}
if self.max_header_count == 0 {
return Err(ZentinelError::Config {
message: "max_header_count must be greater than 0".to_string(),
source: None,
});
}
if self.max_body_buffer_bytes > self.max_body_size_bytes {
return Err(ZentinelError::Config {
message: "max_body_buffer_bytes cannot exceed max_body_size_bytes".to_string(),
source: None,
});
}
if self.max_decompression_ratio <= 0.0 {
return Err(ZentinelError::Config {
message: "max_decompression_ratio must be positive".to_string(),
source: None,
});
}
if let Some(pct) = self.max_memory_percent {
if pct <= 0.0 || pct > 100.0 {
return Err(ZentinelError::Config {
message: "max_memory_percent must be between 0 and 100".to_string(),
source: None,
});
}
}
Ok(())
}
pub fn check_header_size(&self, size: usize) -> ZentinelResult<()> {
if size > self.max_header_size_bytes {
return Err(ZentinelError::limit_exceeded(
LimitType::HeaderSize,
size,
self.max_header_size_bytes,
));
}
Ok(())
}
pub fn check_header_count(&self, count: usize) -> ZentinelResult<()> {
if count > self.max_header_count {
return Err(ZentinelError::limit_exceeded(
LimitType::HeaderCount,
count,
self.max_header_count,
));
}
Ok(())
}
pub fn check_body_size(&self, size: usize) -> ZentinelResult<()> {
if size > self.max_body_size_bytes {
return Err(ZentinelError::limit_exceeded(
LimitType::BodySize,
size,
self.max_body_size_bytes,
));
}
Ok(())
}
}
#[derive(Debug)]
pub struct RateLimiter {
capacity: u32,
tokens: Arc<RwLock<f64>>,
refill_rate: f64,
last_refill: Arc<RwLock<Instant>>,
}
impl RateLimiter {
pub fn new(capacity: u32, refill_per_second: u32) -> Self {
trace!(
capacity = capacity,
refill_per_second = refill_per_second,
"Creating rate limiter"
);
Self {
capacity,
tokens: Arc::new(RwLock::new(capacity as f64)),
refill_rate: refill_per_second as f64,
last_refill: Arc::new(RwLock::new(Instant::now())),
}
}
pub fn try_acquire(&self, tokens: u32) -> bool {
self.refill();
let mut available_tokens = self.tokens.write();
if *available_tokens >= tokens as f64 {
*available_tokens -= tokens as f64;
trace!(
tokens_requested = tokens,
tokens_remaining = *available_tokens as u32,
"Rate limiter: tokens acquired"
);
true
} else {
trace!(
tokens_requested = tokens,
tokens_available = *available_tokens as u32,
"Rate limiter: insufficient tokens"
);
false
}
}
pub fn check(&self, tokens: u32) -> bool {
self.refill();
let available_tokens = self.tokens.read();
*available_tokens >= tokens as f64
}
pub fn available(&self) -> u32 {
self.refill();
let tokens = self.tokens.read();
*tokens as u32
}
fn refill(&self) {
let now = Instant::now();
let mut last_refill = self.last_refill.write();
let elapsed = now.duration_since(*last_refill).as_secs_f64();
if elapsed > 0.0 {
let mut tokens = self.tokens.write();
let tokens_to_add = elapsed * self.refill_rate;
*tokens = (*tokens + tokens_to_add).min(self.capacity as f64);
*last_refill = now;
}
}
pub fn reset(&self) {
let mut tokens = self.tokens.write();
*tokens = self.capacity as f64;
let mut last_refill = self.last_refill.write();
*last_refill = Instant::now();
}
pub fn last_accessed(&self) -> Instant {
*self.last_refill.read()
}
}
pub struct MultiRateLimiter {
global: Option<RateLimiter>,
per_client: Arc<RwLock<HashMap<String, RateLimiter>>>,
per_route: Arc<RwLock<HashMap<String, RateLimiter>>>,
client_limit: Option<(u32, u32)>, route_limit: Option<(u32, u32)>, }
impl MultiRateLimiter {
pub fn new(limits: &Limits) -> Self {
let global = limits
.max_requests_per_second_global
.map(|rps| RateLimiter::new(rps * 10, rps));
let client_limit = limits
.max_requests_per_second_per_client
.map(|rps| (rps * 10, rps));
let route_limit = limits
.max_requests_per_second_per_route
.map(|rps| (rps * 10, rps));
Self {
global,
per_client: Arc::new(RwLock::new(HashMap::new())),
per_route: Arc::new(RwLock::new(HashMap::new())),
client_limit,
route_limit,
}
}
pub fn check_request(&self, client_id: &str, route: &str) -> ZentinelResult<()> {
trace!(
client_id = %client_id,
route = %route,
"Checking rate limits"
);
if let Some(ref limiter) = self.global {
if !limiter.try_acquire(1) {
warn!(
client_id = %client_id,
route = %route,
"Global rate limit exceeded"
);
return Err(ZentinelError::RateLimit {
message: "Global rate limit exceeded".to_string(),
limit: limiter.capacity,
window_seconds: 10,
retry_after_seconds: Some(1),
});
}
}
if let Some((capacity, refill)) = self.client_limit {
let mut limiters = self.per_client.write();
let limiter = limiters
.entry(client_id.to_string())
.or_insert_with(|| RateLimiter::new(capacity, refill));
if !limiter.try_acquire(1) {
warn!(
client_id = %client_id,
route = %route,
"Per-client rate limit exceeded"
);
return Err(ZentinelError::RateLimit {
message: format!("Rate limit exceeded for client {}", client_id),
limit: capacity,
window_seconds: 10,
retry_after_seconds: Some(1),
});
}
}
if let Some((capacity, refill)) = self.route_limit {
let mut limiters = self.per_route.write();
let limiter = limiters
.entry(route.to_string())
.or_insert_with(|| RateLimiter::new(capacity, refill));
if !limiter.try_acquire(1) {
warn!(
client_id = %client_id,
route = %route,
"Per-route rate limit exceeded"
);
return Err(ZentinelError::RateLimit {
message: format!("Rate limit exceeded for route {}", route),
limit: capacity,
window_seconds: 10,
retry_after_seconds: Some(1),
});
}
}
trace!(
client_id = %client_id,
route = %route,
"Rate limits check passed"
);
Ok(())
}
pub fn cleanup(&self, max_age: Duration) -> (usize, usize) {
let now = Instant::now();
let clients_before = self.per_client.read().len();
self.per_client.write().retain(|client_id, limiter| {
let age = now.duration_since(limiter.last_accessed());
let keep = age < max_age;
if !keep {
trace!(
client_id = %client_id,
age_secs = age.as_secs(),
"Removing idle client rate limiter"
);
}
keep
});
let clients_removed = clients_before - self.per_client.read().len();
let routes_before = self.per_route.read().len();
self.per_route.write().retain(|route, limiter| {
let age = now.duration_since(limiter.last_accessed());
let keep = age < max_age;
if !keep {
trace!(
route = %route,
age_secs = age.as_secs(),
"Removing idle route rate limiter"
);
}
keep
});
let routes_removed = routes_before - self.per_route.read().len();
if clients_removed > 0 || routes_removed > 0 {
debug!(
clients_removed = clients_removed,
routes_removed = routes_removed,
clients_remaining = self.per_client.read().len(),
routes_remaining = self.per_route.read().len(),
"Rate limiter cleanup completed"
);
}
(clients_removed, routes_removed)
}
pub fn entry_counts(&self) -> (usize, usize) {
(self.per_client.read().len(), self.per_route.read().len())
}
}
pub struct ConnectionLimiter {
per_client: Arc<RwLock<HashMap<String, usize>>>,
per_route: Arc<RwLock<HashMap<String, usize>>>,
total: Arc<RwLock<usize>>,
limits: Limits,
}
impl ConnectionLimiter {
pub fn new(limits: Limits) -> Self {
debug!(
max_total = limits.max_total_connections,
max_per_client = limits.max_connections_per_client,
max_per_route = limits.max_connections_per_route,
"Creating connection limiter"
);
Self {
per_client: Arc::new(RwLock::new(HashMap::new())),
per_route: Arc::new(RwLock::new(HashMap::new())),
total: Arc::new(RwLock::new(0)),
limits,
}
}
pub fn try_acquire(&self, client_id: &str, route: &str) -> ZentinelResult<ConnectionGuard<'_>> {
trace!(
client_id = %client_id,
route = %route,
"Attempting to acquire connection slot"
);
{
let mut total = self.total.write();
if *total >= self.limits.max_total_connections {
warn!(
current = *total,
max = self.limits.max_total_connections,
"Total connection limit exceeded"
);
return Err(ZentinelError::limit_exceeded(
LimitType::ConnectionCount,
*total,
self.limits.max_total_connections,
));
}
*total += 1;
}
{
let mut per_client = self.per_client.write();
let client_count = per_client.entry(client_id.to_string()).or_insert(0);
if *client_count >= self.limits.max_connections_per_client {
*self.total.write() -= 1;
warn!(
client_id = %client_id,
current = *client_count,
max = self.limits.max_connections_per_client,
"Per-client connection limit exceeded"
);
return Err(ZentinelError::limit_exceeded(
LimitType::ConnectionCount,
*client_count,
self.limits.max_connections_per_client,
));
}
*client_count += 1;
}
{
let mut per_route = self.per_route.write();
let route_count = per_route.entry(route.to_string()).or_insert(0);
if *route_count >= self.limits.max_connections_per_route {
*self.total.write() -= 1;
*self.per_client.write().get_mut(client_id).unwrap() -= 1;
warn!(
route = %route,
current = *route_count,
max = self.limits.max_connections_per_route,
"Per-route connection limit exceeded"
);
return Err(ZentinelError::limit_exceeded(
LimitType::ConnectionCount,
*route_count,
self.limits.max_connections_per_route,
));
}
*route_count += 1;
}
trace!(
client_id = %client_id,
route = %route,
"Connection slot acquired"
);
Ok(ConnectionGuard {
limiter: self,
client_id: client_id.to_string(),
route: route.to_string(),
})
}
fn release(&self, client_id: &str, route: &str) {
trace!(
client_id = %client_id,
route = %route,
"Releasing connection slot"
);
*self.total.write() -= 1;
if let Some(count) = self.per_client.write().get_mut(client_id) {
*count = count.saturating_sub(1);
}
if let Some(count) = self.per_route.write().get_mut(route) {
*count = count.saturating_sub(1);
}
}
pub fn stats(&self) -> ConnectionStats {
ConnectionStats {
total: *self.total.read(),
per_client_count: self.per_client.read().len(),
per_route_count: self.per_route.read().len(),
}
}
}
pub struct ConnectionGuard<'a> {
limiter: &'a ConnectionLimiter,
client_id: String,
route: String,
}
impl Drop for ConnectionGuard<'_> {
fn drop(&mut self) {
self.limiter.release(&self.client_id, &self.route);
}
}
#[derive(Debug, Clone, Serialize)]
pub struct ConnectionStats {
pub total: usize,
pub per_client_count: usize,
pub per_route_count: usize,
}
#[cfg(test)]
mod tests {
use super::*;
use std::thread;
use std::time::Duration;
#[test]
fn test_limits_validation() {
let mut limits = Limits::default();
assert!(limits.validate().is_ok());
limits.max_header_size_bytes = 0;
assert!(limits.validate().is_err());
limits = Limits::default();
limits.max_body_buffer_bytes = limits.max_body_size_bytes + 1;
assert!(limits.validate().is_err());
}
#[test]
fn test_rate_limiter() {
let limiter = RateLimiter::new(10, 10);
for _ in 0..10 {
assert!(limiter.try_acquire(1));
}
assert!(!limiter.try_acquire(1));
thread::sleep(Duration::from_millis(200));
assert!(limiter.try_acquire(1));
assert!(limiter.available() > 0);
}
#[test]
fn test_connection_limiter() {
let limits = Limits {
max_total_connections: 100,
max_connections_per_client: 10,
max_connections_per_route: 50,
..Default::default()
};
let limiter = ConnectionLimiter::new(limits);
let _guard1 = limiter.try_acquire("client1", "route1").unwrap();
let _guard2 = limiter.try_acquire("client1", "route1").unwrap();
let stats = limiter.stats();
assert_eq!(stats.total, 2);
}
#[test]
fn test_rate_limiter_last_accessed() {
let limiter = RateLimiter::new(10, 10);
let before = Instant::now();
limiter.try_acquire(1);
let last_accessed = limiter.last_accessed();
assert!(last_accessed >= before);
assert!(last_accessed <= Instant::now());
}
#[test]
fn test_multi_rate_limiter_entry_counts() {
let limits = Limits {
max_requests_per_second_per_client: Some(100),
max_requests_per_second_per_route: Some(1000),
..Default::default()
};
let limiter = MultiRateLimiter::new(&limits);
assert_eq!(limiter.entry_counts(), (0, 0));
let _ = limiter.check_request("client1", "route1");
let _ = limiter.check_request("client2", "route1");
let _ = limiter.check_request("client1", "route2");
assert_eq!(limiter.entry_counts(), (2, 2));
}
#[test]
fn test_multi_rate_limiter_cleanup() {
let limits = Limits {
max_requests_per_second_per_client: Some(100),
max_requests_per_second_per_route: Some(1000),
..Default::default()
};
let limiter = MultiRateLimiter::new(&limits);
let _ = limiter.check_request("client1", "route1");
let _ = limiter.check_request("client2", "route2");
assert_eq!(limiter.entry_counts(), (2, 2));
let (clients_removed, routes_removed) = limiter.cleanup(Duration::from_secs(3600));
assert_eq!(clients_removed, 0);
assert_eq!(routes_removed, 0);
assert_eq!(limiter.entry_counts(), (2, 2));
thread::sleep(Duration::from_millis(50));
let (clients_removed, routes_removed) = limiter.cleanup(Duration::from_millis(10));
assert_eq!(clients_removed, 2);
assert_eq!(routes_removed, 2);
assert_eq!(limiter.entry_counts(), (0, 0));
}
#[test]
fn test_multi_rate_limiter_cleanup_partial() {
let limits = Limits {
max_requests_per_second_per_client: Some(100),
max_requests_per_second_per_route: Some(1000),
..Default::default()
};
let limiter = MultiRateLimiter::new(&limits);
let _ = limiter.check_request("old_client", "old_route");
thread::sleep(Duration::from_millis(60));
let _ = limiter.check_request("new_client", "new_route");
assert_eq!(limiter.entry_counts(), (2, 2));
let (clients_removed, routes_removed) = limiter.cleanup(Duration::from_millis(30));
assert_eq!(clients_removed, 1);
assert_eq!(routes_removed, 1);
assert_eq!(limiter.entry_counts(), (1, 1));
}
}