use std::collections::VecDeque;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CellDecision {
Allowed,
Rejected { retry_after_ms: u64 },
}
pub trait RateAlgorithm: Send {
fn try_admit(&mut self, now_ms: i64) -> CellDecision;
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub enum Algorithm {
#[default]
TokenBucket,
SlidingWindow,
LeakyBucket,
}
impl Algorithm {
pub fn parse(value: &str) -> Option<Self> {
match value.trim().to_ascii_lowercase().as_str() {
"token_bucket" | "token-bucket" | "tokenbucket" => Some(Self::TokenBucket),
"sliding_window" | "sliding-window" | "slidingwindow" => Some(Self::SlidingWindow),
"leaky_bucket" | "leaky-bucket" | "leakybucket" => Some(Self::LeakyBucket),
_ => None,
}
}
pub fn as_str(self) -> &'static str {
match self {
Self::TokenBucket => "token_bucket",
Self::SlidingWindow => "sliding_window",
Self::LeakyBucket => "leaky_bucket",
}
}
pub fn new_cell(self, rate_per_sec: f64, capacity: u32, now_ms: i64) -> Box<dyn RateAlgorithm> {
match self {
Self::TokenBucket => Box::new(TokenBucket::new(rate_per_sec, capacity, now_ms)),
Self::SlidingWindow => Box::new(SlidingWindow::new(
capacity,
window_ms_for_rate(rate_per_sec, capacity),
)),
Self::LeakyBucket => Box::new(LeakyBucket::new(rate_per_sec, capacity, now_ms)),
}
}
}
fn window_ms_for_rate(rate_per_sec: f64, capacity: u32) -> u64 {
if rate_per_sec <= 0.0 || capacity == 0 {
return 1_000;
}
let window_secs = capacity as f64 / rate_per_sec;
((window_secs * 1_000.0).round() as u64).max(1)
}
#[derive(Debug)]
pub struct TokenBucket {
rate_per_sec: f64,
capacity: f64,
tokens: f64,
last_refill_ms: i64,
}
impl TokenBucket {
pub fn new(rate_per_sec: f64, capacity: u32, now_ms: i64) -> Self {
Self {
rate_per_sec: rate_per_sec.max(0.0),
capacity: capacity as f64,
tokens: capacity as f64,
last_refill_ms: now_ms,
}
}
fn refill(&mut self, now_ms: i64) {
if now_ms <= self.last_refill_ms {
return;
}
let delta_ms = (now_ms - self.last_refill_ms) as f64;
let gained = (delta_ms / 1_000.0) * self.rate_per_sec;
self.tokens = (self.tokens + gained).min(self.capacity);
self.last_refill_ms = now_ms;
}
}
impl RateAlgorithm for TokenBucket {
fn try_admit(&mut self, now_ms: i64) -> CellDecision {
self.refill(now_ms);
if self.tokens >= 1.0 {
self.tokens -= 1.0;
return CellDecision::Allowed;
}
let deficit = 1.0 - self.tokens;
let retry_after_ms = if self.rate_per_sec > 0.0 {
((deficit / self.rate_per_sec) * 1_000.0).ceil() as u64
} else {
1_000
};
CellDecision::Rejected {
retry_after_ms: retry_after_ms.max(1),
}
}
}
#[derive(Debug)]
pub struct SlidingWindow {
capacity: usize,
window_ms: u64,
hits: VecDeque<i64>,
}
impl SlidingWindow {
pub fn new(capacity: u32, window_ms: u64) -> Self {
Self {
capacity: capacity as usize,
window_ms,
hits: VecDeque::with_capacity(capacity as usize),
}
}
fn trim(&mut self, now_ms: i64) {
let cutoff = now_ms.saturating_sub(self.window_ms as i64);
while let Some(&front) = self.hits.front() {
if front <= cutoff {
self.hits.pop_front();
} else {
break;
}
}
}
}
impl RateAlgorithm for SlidingWindow {
fn try_admit(&mut self, now_ms: i64) -> CellDecision {
self.trim(now_ms);
if self.capacity == 0 {
return CellDecision::Rejected {
retry_after_ms: self.window_ms.max(1),
};
}
if self.hits.len() < self.capacity {
self.hits.push_back(now_ms);
return CellDecision::Allowed;
}
let oldest = *self.hits.front().expect("hits non-empty when full");
let earliest_open_ms = oldest + self.window_ms as i64;
let retry_after_ms = ((earliest_open_ms - now_ms).max(1)) as u64;
CellDecision::Rejected { retry_after_ms }
}
}
#[derive(Debug)]
pub struct LeakyBucket {
rate_per_sec: f64,
capacity: f64,
level: f64,
last_drip_ms: i64,
}
impl LeakyBucket {
pub fn new(rate_per_sec: f64, capacity: u32, now_ms: i64) -> Self {
Self {
rate_per_sec: rate_per_sec.max(0.0),
capacity: capacity as f64,
level: 0.0,
last_drip_ms: now_ms,
}
}
fn drain(&mut self, now_ms: i64) {
if now_ms <= self.last_drip_ms {
return;
}
let delta_ms = (now_ms - self.last_drip_ms) as f64;
let drained = (delta_ms / 1_000.0) * self.rate_per_sec;
self.level = (self.level - drained).max(0.0);
self.last_drip_ms = now_ms;
}
}
impl RateAlgorithm for LeakyBucket {
fn try_admit(&mut self, now_ms: i64) -> CellDecision {
self.drain(now_ms);
if self.level + 1.0 <= self.capacity {
self.level += 1.0;
return CellDecision::Allowed;
}
let overflow = (self.level + 1.0) - self.capacity;
let retry_after_ms = if self.rate_per_sec > 0.0 {
((overflow / self.rate_per_sec) * 1_000.0).ceil() as u64
} else {
1_000
};
CellDecision::Rejected {
retry_after_ms: retry_after_ms.max(1),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn token_bucket_admits_up_to_capacity_then_throttles() {
let mut bucket = TokenBucket::new(1.0, 3, 0);
assert_eq!(bucket.try_admit(0), CellDecision::Allowed);
assert_eq!(bucket.try_admit(0), CellDecision::Allowed);
assert_eq!(bucket.try_admit(0), CellDecision::Allowed);
match bucket.try_admit(0) {
CellDecision::Rejected { retry_after_ms } => {
assert!((900..=1_100).contains(&retry_after_ms), "{retry_after_ms}");
}
other => panic!("expected reject, got {other:?}"),
}
assert_eq!(bucket.try_admit(1_000), CellDecision::Allowed);
}
#[test]
fn sliding_window_rejects_when_full_until_oldest_drops_off() {
let mut window = SlidingWindow::new(3, 1_000);
assert_eq!(window.try_admit(0), CellDecision::Allowed);
assert_eq!(window.try_admit(100), CellDecision::Allowed);
assert_eq!(window.try_admit(200), CellDecision::Allowed);
match window.try_admit(300) {
CellDecision::Rejected { retry_after_ms } => assert_eq!(retry_after_ms, 700),
other => panic!("expected reject, got {other:?}"),
}
assert_eq!(window.try_admit(1_001), CellDecision::Allowed);
}
#[test]
fn leaky_bucket_smooths_burst_to_drain_rate() {
let mut bucket = LeakyBucket::new(1.0, 2, 0);
assert_eq!(bucket.try_admit(0), CellDecision::Allowed);
assert_eq!(bucket.try_admit(0), CellDecision::Allowed);
match bucket.try_admit(0) {
CellDecision::Rejected { retry_after_ms } => {
assert!((900..=1_100).contains(&retry_after_ms), "{retry_after_ms}");
}
other => panic!("expected reject, got {other:?}"),
}
assert_eq!(bucket.try_admit(1_000), CellDecision::Allowed);
}
#[test]
fn parse_algorithm_accepts_friendly_aliases() {
assert_eq!(
Algorithm::parse("token_bucket"),
Some(Algorithm::TokenBucket)
);
assert_eq!(
Algorithm::parse("TOKEN-BUCKET"),
Some(Algorithm::TokenBucket)
);
assert_eq!(
Algorithm::parse("sliding_window"),
Some(Algorithm::SlidingWindow)
);
assert_eq!(
Algorithm::parse("leaky_bucket"),
Some(Algorithm::LeakyBucket)
);
assert_eq!(Algorithm::parse("nope"), None);
}
}