use dashmap::DashMap;
use itoa::Buffer as ItoaBuffer;
use ntex::http::header::{HeaderName, HeaderValue};
use ntex::service::cfg::SharedCfg;
use ntex::{http::StatusCode, Middleware, ServiceCtx};
use std::net::IpAddr;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use ntex::{web, Service};
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
#[cfg(feature = "tokio")]
use tokio::time::interval;
#[cfg(feature = "smol")]
use smol::Timer;
#[cfg(feature = "json")]
use serde::{Deserialize, Serialize};
const HEADER_RATELIMIT_REMAINING: &str = "x-ratelimit-remaining";
const HEADER_RATELIMIT_LIMIT: &str = "x-ratelimit-limit";
const HEADER_RATELIMIT_RESET: &str = "x-ratelimit-reset";
const OVERFLOW_KEY: IpAddr = IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED);
#[derive(Debug)]
struct TokenBucket {
tokens: f64,
last_refill: Instant,
}
impl TokenBucket {
fn new(capacity: usize) -> Self {
Self {
tokens: capacity as f64,
last_refill: Instant::now(),
}
}
fn consume(&mut self, tokens: usize, now: Instant, config: &RateLimiterConfig) -> bool {
self.refill(now, config);
if self.tokens >= tokens as f64 {
self.tokens -= tokens as f64;
true
} else {
false
}
}
fn refill(&mut self, now: Instant, config: &RateLimiterConfig) {
let elapsed = now.duration_since(self.last_refill).as_secs_f64();
let refill_rate = config.capacity as f64 / config.window as f64;
let new_tokens = elapsed * refill_rate;
self.tokens = (self.tokens + new_tokens).min(config.capacity as f64);
self.last_refill = now;
}
fn is_stale(&self, now: Instant, stale_threshold: Duration) -> bool {
now.duration_since(self.last_refill) > stale_threshold
}
}
fn compute_reset_time(tokens: f64, config: &RateLimiterConfig) -> u64 {
let now_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
if tokens >= config.capacity as f64 {
return now_secs;
}
let missing_tokens = config.capacity as f64 - tokens;
let refill_rate = config.capacity as f64 / config.window as f64;
let seconds_to_refill = missing_tokens / refill_rate;
now_secs + seconds_to_refill.ceil() as u64
}
#[derive(Debug, Clone)]
pub struct RateLimiterConfig {
pub capacity: usize,
pub window: u64,
pub cleanup_interval: Duration,
pub stale_threshold: Duration,
pub trust_proxy_headers: bool,
pub max_entries: usize,
}
impl Default for RateLimiterConfig {
fn default() -> Self {
Self {
capacity: 100,
window: 60,
cleanup_interval: Duration::from_secs(300), stale_threshold: Duration::from_secs(3600), trust_proxy_headers: false,
max_entries: 100_000,
}
}
}
pub struct RateLimiter {
map: DashMap<IpAddr, TokenBucket>,
config: RateLimiterConfig,
entries: AtomicUsize,
}
impl RateLimiter {
pub fn new(capacity: usize, window: u64) -> Arc<Self> {
let config = RateLimiterConfig {
capacity,
window,
..Default::default()
};
Self::with_config(config)
}
pub fn with_config(config: RateLimiterConfig) -> Arc<Self> {
assert!(config.window > 0, "RateLimiter window must be greater than zero");
assert!(
!config.cleanup_interval.is_zero(),
"RateLimiter cleanup_interval must be greater than zero"
);
let limiter = Arc::new(RateLimiter {
map: DashMap::new(),
config,
entries: AtomicUsize::new(0),
});
#[cfg(any(feature = "tokio", feature = "smol"))]
Self::start_cleanup_task(Arc::clone(&limiter));
limiter
}
#[cfg(feature = "tokio")]
fn start_cleanup_task(limiter: Arc<RateLimiter>) {
let cleanup_interval = limiter.config.cleanup_interval;
let weak = Arc::downgrade(&limiter);
tokio::spawn(async move {
let mut interval = interval(cleanup_interval);
loop {
interval.tick().await;
let Some(limiter) = weak.upgrade() else {
break;
};
limiter.cleanup().await;
}
});
}
#[cfg(feature = "smol")]
fn start_cleanup_task(limiter: Arc<RateLimiter>) {
let cleanup_interval = limiter.config.cleanup_interval;
let weak = Arc::downgrade(&limiter);
smol::spawn(async move {
loop {
Timer::after(cleanup_interval).await;
let Some(limiter) = weak.upgrade() else {
break;
};
limiter.cleanup().await;
}
})
.detach();
}
pub fn check_rate_limit(&self, identifier: IpAddr) -> RateLimitResult {
let now = Instant::now();
let limit = self.config.capacity;
let key = if self.entries.load(Ordering::Relaxed) < self.config.max_entries
|| self.map.contains_key(&identifier)
{
identifier
} else {
OVERFLOW_KEY
};
let (allowed, tokens) = {
let mut bucket = self.map.entry(key).or_insert_with(|| {
self.entries.fetch_add(1, Ordering::Relaxed);
TokenBucket::new(limit)
});
let allowed = bucket.consume(1, now, &self.config);
(allowed, bucket.tokens)
};
RateLimitResult {
allowed,
remaining: tokens.floor() as u32,
reset: compute_reset_time(tokens, &self.config),
limit,
}
}
async fn cleanup(&self) {
let now = Instant::now();
let stale_threshold = self.config.stale_threshold;
let mut removed = 0usize;
self.map.retain(|_, bucket| {
let keep = !bucket.is_stale(now, stale_threshold);
if !keep {
removed += 1;
}
keep
});
if removed > 0 {
self.entries.fetch_sub(removed, Ordering::Relaxed);
if cfg!(debug_assertions) {
eprintln!("Cleaned {removed} stale rate limit entries");
}
}
}
pub fn stats(&self) -> RateLimiterStats {
RateLimiterStats {
active_entries: self.map.len(),
capacity: self.config.capacity,
window: self.config.window,
}
}
}
#[derive(Debug, Clone)]
pub struct RateLimitResult {
pub allowed: bool,
pub remaining: u32,
pub reset: u64,
pub limit: usize,
}
#[derive(Debug, Clone)]
pub struct RateLimiterStats {
pub active_entries: usize,
pub capacity: usize,
pub window: u64,
}
pub struct RateLimit {
pub limiter: Arc<RateLimiter>,
}
impl RateLimit {
pub fn new(limiter: Arc<RateLimiter>) -> Self {
Self { limiter }
}
}
impl<S> Middleware<S, SharedCfg> for RateLimit {
type Service = RateLimitMiddlewareService<S>;
fn create(&self, service: S, _cfg: SharedCfg) -> Self::Service {
RateLimitMiddlewareService {
service,
limiter: Arc::clone(&self.limiter),
}
}
}
pub struct RateLimitMiddlewareService<S> {
service: S,
limiter: Arc<RateLimiter>,
}
impl<S, Err> Service<web::WebRequest<Err>> for RateLimitMiddlewareService<S>
where
S: Service<web::WebRequest<Err>, Response = web::WebResponse, Error = web::Error> + 'static,
Err: web::ErrorRenderer,
{
type Response = web::WebResponse;
type Error = web::Error;
async fn call(
&self,
req: web::WebRequest<Err>,
ctx: ServiceCtx<'_, Self>,
) -> Result<Self::Response, Self::Error> {
let ip = extract_client_ip(&req, self.limiter.config.trust_proxy_headers);
let result = self.limiter.check_rate_limit(ip);
if !result.allowed {
return Err(RateLimitError::from(result).into());
}
let mut response = ctx.call(&self.service, req).await?;
add_rate_limit_headers(response.headers_mut(), &result);
Ok(response)
}
}
fn extract_client_ip<Err>(req: &web::WebRequest<Err>, trust_proxy_headers: bool) -> IpAddr {
if trust_proxy_headers {
if let Some(forwarded) = req.headers().get("x-forwarded-for") {
if let Ok(forwarded_str) = forwarded.to_str() {
if let Some(ip) = forwarded_str.split(',').next() {
let ip = ip.trim();
if let Ok(parsed_ip) = ip.parse::<IpAddr>() {
if !parsed_ip.is_unspecified() {
return parsed_ip;
}
}
}
}
}
if let Some(real_ip) = req.headers().get("x-real-ip") {
if let Ok(ip_str) = real_ip.to_str() {
let ip = ip_str.trim();
if let Ok(parsed_ip) = ip.parse::<IpAddr>() {
if !parsed_ip.is_unspecified() {
return parsed_ip;
}
}
}
}
}
if let Some(peer) = req.peer_addr() {
return peer.ip();
}
IpAddr::V4(std::net::Ipv4Addr::new(127, 0, 0, 1))
}
fn add_rate_limit_headers(headers: &mut ntex::http::HeaderMap, result: &RateLimitResult) {
let mut buf = ItoaBuffer::new();
if let Ok(value) = HeaderValue::from_str(buf.format(result.remaining)) {
headers.insert(HeaderName::from_static(HEADER_RATELIMIT_REMAINING), value);
}
if let Ok(value) = HeaderValue::from_str(buf.format(result.limit)) {
headers.insert(HeaderName::from_static(HEADER_RATELIMIT_LIMIT), value);
}
if let Ok(value) = HeaderValue::from_str(buf.format(result.reset)) {
headers.insert(HeaderName::from_static(HEADER_RATELIMIT_RESET), value);
}
}
#[derive(Debug)]
#[cfg_attr(feature = "json", derive(Serialize, Deserialize))]
struct RateLimitErrorData {
remaining: u32,
reset: u64,
limit: usize,
}
#[cfg(feature = "json")]
#[derive(Debug, Serialize, Deserialize)]
struct RateLimitErrorResponse {
code: u32,
message: String,
data: RateLimitErrorData,
}
#[derive(Debug)]
struct RateLimitError {
data: RateLimitErrorData,
}
impl From<RateLimitResult> for RateLimitError {
fn from(result: RateLimitResult) -> Self {
Self {
data: RateLimitErrorData {
remaining: result.remaining,
reset: result.reset,
limit: result.limit,
},
}
}
}
impl std::fmt::Display for RateLimitError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"Rate limit exceeded. Remaining: {}, Reset: {}, Limit: {}",
self.data.remaining, self.data.reset, self.data.limit
)
}
}
impl web::error::WebResponseError for RateLimitError {
fn error_response(&self, _: &ntex::web::HttpRequest) -> web::HttpResponse {
#[cfg(feature = "json")]
let body = {
let error_response = RateLimitErrorResponse {
code: 429,
message: "Rate limit exceeded".to_string(),
data: RateLimitErrorData {
remaining: self.data.remaining,
reset: self.data.reset,
limit: self.data.limit,
},
};
serde_json::to_string(&error_response)
.unwrap_or_else(|_| r#"{"code":429,"message":"Rate limit exceeded"}"#.to_string())
};
#[cfg(not(feature = "json"))]
let body = format!(
r#"{{"code":429,"message":"Rate limit exceeded","data":{{"remaining":{},"reset":{},"limit":{}}}}}"#,
self.data.remaining, self.data.reset, self.data.limit
);
let mut buf = ItoaBuffer::new();
web::HttpResponse::build(StatusCode::TOO_MANY_REQUESTS)
.set_header("content-type", "application/json")
.set_header(HEADER_RATELIMIT_REMAINING, buf.format(self.data.remaining))
.set_header(HEADER_RATELIMIT_LIMIT, buf.format(self.data.limit))
.set_header(HEADER_RATELIMIT_RESET, buf.format(self.data.reset))
.body(body)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_token_bucket_basic() {
let config = RateLimiterConfig {
capacity: 5,
window: 10,
..Default::default()
};
let mut bucket = TokenBucket::new(5);
let now = Instant::now();
for _ in 0..5 {
assert!(bucket.consume(1, now, &config));
}
assert!(!bucket.consume(1, now, &config));
assert_eq!(bucket.tokens.floor() as u32, 0);
}
#[test]
fn test_token_bucket_refill() {
let config = RateLimiterConfig {
capacity: 10,
window: 10, ..Default::default()
};
let mut bucket = TokenBucket::new(10);
let now = Instant::now();
for _ in 0..10 {
assert!(bucket.consume(1, now, &config));
}
assert!(!bucket.consume(1, now, &config));
let later = now + Duration::from_secs(5);
bucket.refill(later, &config);
assert_eq!(bucket.tokens.floor() as u32, 5);
for _ in 0..5 {
assert!(bucket.consume(1, later, &config));
}
assert!(!bucket.consume(1, later, &config));
}
fn check_capacity_5(limiter: &RateLimiter) {
let ip = "192.168.1.1".parse::<IpAddr>().unwrap();
for i in 0..5 {
let result = limiter.check_rate_limit(ip);
assert!(result.allowed, "Request {} should be allowed", i + 1);
assert_eq!(result.remaining, 4 - i as u32);
}
let result = limiter.check_rate_limit(ip);
assert!(!result.allowed);
assert_eq!(result.remaining, 0);
}
fn check_different_ips(limiter: &RateLimiter) {
let ip1 = "192.168.1.1".parse::<IpAddr>().unwrap();
let ip2 = "192.168.1.2".parse::<IpAddr>().unwrap();
let result1 = limiter.check_rate_limit(ip1);
let result2 = limiter.check_rate_limit(ip2);
assert!(result1.allowed);
assert!(result2.allowed);
assert_eq!(result1.remaining, 1);
assert_eq!(result2.remaining, 1);
}
#[cfg(feature = "tokio")]
#[tokio::test]
async fn test_rate_limiter() {
let limiter = RateLimiter::with_config(RateLimiterConfig {
capacity: 5,
window: 1,
..Default::default()
});
check_capacity_5(&limiter);
}
#[cfg(feature = "tokio")]
#[tokio::test]
async fn test_rate_limiter_different_ips() {
let limiter = RateLimiter::new(2, 60);
check_different_ips(&limiter);
}
#[cfg(feature = "smol")]
#[test]
fn test_rate_limiter() {
smol::block_on(async {
let limiter = RateLimiter::with_config(RateLimiterConfig {
capacity: 5,
window: 1,
..Default::default()
});
check_capacity_5(&limiter);
});
}
#[cfg(feature = "smol")]
#[test]
fn test_rate_limiter_different_ips() {
smol::block_on(async {
let limiter = RateLimiter::new(2, 60);
check_different_ips(&limiter);
});
}
fn check_overflow_routing(limiter: &RateLimiter) {
let ip1 = "10.0.0.1".parse::<IpAddr>().unwrap();
let ip2 = "10.0.0.2".parse::<IpAddr>().unwrap();
let ip3 = "10.0.0.3".parse::<IpAddr>().unwrap();
let ip4 = "10.0.0.4".parse::<IpAddr>().unwrap();
assert!(limiter.check_rate_limit(ip1).allowed);
assert!(limiter.check_rate_limit(ip2).allowed);
assert!(
limiter.check_rate_limit(ip3).allowed,
"first overflow hit should be allowed"
);
assert!(
!limiter.check_rate_limit(ip4).allowed,
"second overflow hit should be denied"
);
assert!(limiter.map.len() <= 3);
}
#[cfg(feature = "tokio")]
#[tokio::test]
async fn test_overflow_bucket() {
let limiter = RateLimiter::with_config(RateLimiterConfig {
capacity: 1,
window: 60,
max_entries: 2,
..Default::default()
});
check_overflow_routing(&limiter);
}
#[cfg(feature = "smol")]
#[test]
fn test_overflow_bucket() {
smol::block_on(async {
let limiter = RateLimiter::with_config(RateLimiterConfig {
capacity: 1,
window: 60,
max_entries: 2,
..Default::default()
});
check_overflow_routing(&limiter);
});
}
#[test]
#[should_panic(expected = "cleanup_interval must be greater than zero")]
fn test_zero_cleanup_interval_rejected() {
let _ = RateLimiter::with_config(RateLimiterConfig {
cleanup_interval: Duration::ZERO,
..Default::default()
});
}
#[test]
fn test_extract_client_ip_trust_proxy() {
use ntex::web::test::TestRequest;
let req = TestRequest::default()
.header("x-forwarded-for", "1.2.3.4")
.to_srv_request();
assert_eq!(
extract_client_ip(&req, true),
"1.2.3.4".parse::<IpAddr>().unwrap()
);
let req = TestRequest::default()
.header("x-forwarded-for", "1.2.3.4")
.to_srv_request();
assert_eq!(
extract_client_ip(&req, false),
"127.0.0.1".parse::<IpAddr>().unwrap()
);
let req = TestRequest::default()
.header("x-real-ip", "5.6.7.8")
.to_srv_request();
assert_eq!(
extract_client_ip(&req, true),
"5.6.7.8".parse::<IpAddr>().unwrap()
);
let req = TestRequest::default()
.header("x-forwarded-for", "not-an-ip")
.to_srv_request();
assert_eq!(
extract_client_ip(&req, true),
"127.0.0.1".parse::<IpAddr>().unwrap()
);
let req = TestRequest::default()
.header("x-forwarded-for", "0.0.0.0")
.to_srv_request();
assert_eq!(
extract_client_ip(&req, true),
"127.0.0.1".parse::<IpAddr>().unwrap()
);
}
}