use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{Mutex, Semaphore};
use tokio::time::Instant;
#[derive(Debug)]
struct Bucket {
rate: f64,
capacity: f64,
tokens: f64,
last_refill: Instant,
}
impl Bucket {
fn new(rate_bps: u64, now: Instant) -> Self {
let rate = rate_bps as f64;
Bucket {
rate,
capacity: rate,
tokens: rate,
last_refill: now,
}
}
fn refill(&mut self, now: Instant) {
let elapsed = now
.saturating_duration_since(self.last_refill)
.as_secs_f64();
if elapsed > 0.0 {
self.tokens = (self.tokens + elapsed * self.rate).min(self.capacity);
self.last_refill = now;
}
}
fn wait_for(&mut self, need: f64, now: Instant) -> Duration {
if self.rate <= 0.0 {
return Duration::ZERO; }
self.refill(now);
let effective = need.min(self.capacity);
if self.tokens >= effective {
Duration::ZERO
} else {
Duration::from_secs_f64((effective - self.tokens) / self.rate)
}
}
fn consume(&mut self, need: f64) {
if self.rate > 0.0 {
self.tokens -= need;
}
}
}
#[derive(Debug)]
pub struct FcfsRateLimiter {
gate: Semaphore,
global: Mutex<Bucket>,
per_conn: Mutex<HashMap<String, Bucket>>,
per_conn_bps: u64,
}
impl FcfsRateLimiter {
pub fn new(global_bps: u64, per_conn_bps: u64) -> Arc<Self> {
let now = Instant::now();
Arc::new(FcfsRateLimiter {
gate: Semaphore::new(1),
global: Mutex::new(Bucket::new(global_bps, now)),
per_conn: Mutex::new(HashMap::new()),
per_conn_bps,
})
}
pub async fn acquire(&self, conn_key: &str, bytes: u64) {
if bytes == 0 {
return;
}
let need = bytes as f64;
let _permit = self
.gate
.acquire()
.await
.expect("rate-limiter gate never closed");
loop {
let now = Instant::now();
let wait_global = self.global.lock().await.wait_for(need, now);
let wait_conn = {
let mut map = self.per_conn.lock().await;
let bucket = map
.entry(conn_key.to_string())
.or_insert_with(|| Bucket::new(self.per_conn_bps, now));
bucket.wait_for(need, now)
};
let wait = wait_global.max(wait_conn);
if wait.is_zero() {
self.global.lock().await.consume(need);
if let Some(bucket) = self.per_conn.lock().await.get_mut(conn_key) {
bucket.consume(need);
}
return;
}
tokio::time::sleep(wait).await;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test(start_paused = true)]
async fn unlimited_limiter_never_waits() {
let rl = FcfsRateLimiter::new(0, 0);
let start = Instant::now();
for _ in 0..100 {
rl.acquire("peer", 1_000_000).await;
}
assert_eq!(start.elapsed(), Duration::ZERO, "no cap → no wait");
}
#[tokio::test(start_paused = true)]
async fn global_cap_paces_total_throughput() {
let rl = FcfsRateLimiter::new(1000, 0);
let start = Instant::now();
rl.acquire("a", 1000).await; rl.acquire("a", 1000).await; let elapsed = start.elapsed();
assert!(
elapsed >= Duration::from_millis(900),
"second 1000B should wait ~1s for refill, waited {elapsed:?}"
);
}
#[tokio::test(start_paused = true)]
async fn per_conn_cap_isolates_connections() {
let rl = FcfsRateLimiter::new(0, 1000);
rl.acquire("A", 1000).await; let start = Instant::now();
rl.acquire("B", 1000).await; assert_eq!(start.elapsed(), Duration::ZERO, "B has an independent cap");
let start_a = Instant::now();
rl.acquire("A", 1000).await;
assert!(start_a.elapsed() >= Duration::from_millis(900));
}
#[tokio::test(start_paused = true)]
async fn fcfs_preserves_arrival_order() {
let rl = FcfsRateLimiter::new(1000, 0);
let order = Arc::new(Mutex::new(Vec::<u32>::new()));
let mut handles = Vec::new();
for i in 0..3u32 {
let rl = rl.clone();
let order = order.clone();
tokio::time::sleep(Duration::from_millis(1)).await;
handles.push(tokio::spawn(async move {
rl.acquire("peer", 1000).await;
order.lock().await.push(i);
}));
}
for h in handles {
h.await.unwrap();
}
assert_eq!(
*order.lock().await,
vec![0, 1, 2],
"served in arrival order"
);
}
#[tokio::test(start_paused = true)]
async fn zero_bytes_is_instant_and_free() {
let rl = FcfsRateLimiter::new(1, 1);
let start = Instant::now();
rl.acquire("peer", 0).await;
assert_eq!(start.elapsed(), Duration::ZERO);
}
#[tokio::test(start_paused = true)]
async fn oversized_request_does_not_deadlock() {
let rl = FcfsRateLimiter::new(1000, 0);
rl.acquire("peer", 5000).await; rl.acquire("peer", 10).await;
}
}