use serde::{Deserialize, Serialize};
use std::{
collections::{HashMap, VecDeque},
sync::{Arc, Mutex, MutexGuard, PoisonError},
time::Duration,
};
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RateLimitPolicy {
#[default]
Off,
Wait {
max_wait: Duration,
},
Fail,
}
const DEFAULT_WINDOW_MS: i64 = 900_000;
const DEFAULT_COST: i64 = 2;
const WAIT_MARGIN_MS: i64 = 50;
#[derive(Debug)]
struct Bucket {
window_ms: i64,
remaining: i64,
cost: i64,
last_response_ms: i64,
spent: VecDeque<(i64, i64)>,
in_flight: i64,
blocked_until_ms: i64,
seen: bool,
}
impl Default for Bucket {
fn default() -> Self {
Bucket {
window_ms: DEFAULT_WINDOW_MS,
remaining: 0,
cost: DEFAULT_COST,
last_response_ms: 0,
spent: VecDeque::new(),
in_flight: 0,
blocked_until_ms: 0,
seen: false,
}
}
}
impl Bucket {
fn wait_ms(&self, now: i64) -> i64 {
if !self.seen {
return 0;
}
let needed = (self.in_flight + 1) * self.cost;
let mut available = self.remaining;
let mut ready_at = now;
if available < needed {
ready_at = self.last_response_ms + self.window_ms;
for (at, cost) in &self.spent {
let returns_at = at + self.window_ms;
if returns_at <= self.last_response_ms {
continue;
}
available += cost;
if available >= needed {
ready_at = returns_at;
break;
}
}
}
let wait = (ready_at - now).max(self.blocked_until_ms - now).max(0);
if wait > 0 {
wait + WAIT_MARGIN_MS
} else {
0
}
}
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum Acquire {
Granted,
Wait(i64),
}
#[derive(Debug, Default)]
pub(crate) struct RateLimiter {
buckets: Mutex<HashMap<String, Bucket>>,
}
impl RateLimiter {
fn lock(&self) -> MutexGuard<'_, HashMap<String, Bucket>> {
self.buckets.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(crate) fn key(group: &str, token: Option<&str>) -> String {
format!("{group}|{}", token.unwrap_or(""))
}
pub(crate) fn group_of(key: &str) -> &str {
key.split('|').next().unwrap_or(key)
}
pub(crate) fn try_acquire(&self, key: &str, now: i64) -> Acquire {
let mut buckets = self.lock();
let bucket = buckets.entry(key.to_owned()).or_default();
match bucket.wait_ms(now) {
0 => {
bucket.in_flight += 1;
Acquire::Granted
}
wait => Acquire::Wait(wait),
}
}
pub(crate) fn release(&self, key: &str) {
if let Some(bucket) = self.lock().get_mut(key) {
bucket.in_flight = (bucket.in_flight - 1).max(0);
}
}
pub(crate) fn record(
&self,
key: &str,
remaining: i64,
used: i64,
window_ms: i64,
success: bool,
now: i64,
) {
let mut buckets = self.lock();
let bucket = buckets.entry(key.to_owned()).or_default();
if window_ms > 0 {
bucket.window_ms = window_ms;
}
bucket.remaining = remaining;
bucket.last_response_ms = now;
bucket.seen = true;
if success && used > 0 {
bucket.cost = used;
}
if used > 0 {
bucket.spent.push_back((now, used));
}
let window = bucket.window_ms;
while bucket
.spent
.front()
.is_some_and(|(at, _)| at + window <= now)
{
bucket.spent.pop_front();
}
buckets.retain(|_, b| b.in_flight > 0 || b.last_response_ms + b.window_ms > now);
}
pub(crate) fn block_until(&self, key: &str, until_ms: i64) {
let mut buckets = self.lock();
let bucket = buckets.entry(key.to_owned()).or_default();
bucket.blocked_until_ms = bucket.blocked_until_ms.max(until_ms);
bucket.seen = true;
}
}
pub(crate) struct Permit {
limiter: Arc<RateLimiter>,
key: String,
}
impl Permit {
pub(crate) fn new(limiter: Arc<RateLimiter>, key: &str) -> Self {
Permit {
limiter,
key: key.to_owned(),
}
}
}
impl Drop for Permit {
fn drop(&mut self) {
self.limiter.release(&self.key);
}
}
#[cfg(test)]
mod tests {
use super::*;
const KEY: &str = "g|t";
#[test]
fn test_unseen_groups_are_not_throttled() {
let limiter = RateLimiter::default();
for _ in 0..10 {
assert_eq!(limiter.try_acquire(KEY, 0), Acquire::Granted);
}
}
#[test]
fn test_requests_that_fit_are_granted_and_in_flight_ones_count() {
let limiter = RateLimiter::default();
limiter.record(KEY, 5, 2, 10_000, true, 0);
assert_eq!(limiter.try_acquire(KEY, 1), Acquire::Granted);
assert_eq!(limiter.try_acquire(KEY, 1), Acquire::Granted);
assert!(matches!(limiter.try_acquire(KEY, 1), Acquire::Wait(_)));
limiter.release(KEY);
assert_eq!(limiter.try_acquire(KEY, 1), Acquire::Granted);
}
#[test]
fn test_tokens_return_when_their_window_ends() {
let limiter = RateLimiter::default();
limiter.record(KEY, 0, 2, 1_000, true, 0);
let Acquire::Wait(wait) = limiter.try_acquire(KEY, 100) else {
panic!("should wait");
};
assert_eq!(wait, 900 + WAIT_MARGIN_MS);
assert_eq!(limiter.try_acquire(KEY, 1_000), Acquire::Granted);
}
#[test]
fn test_the_wait_ends_at_the_first_return_that_is_enough() {
let limiter = RateLimiter::default();
for at in [0, 400, 800] {
limiter.record(KEY, 0, 2, 1_000, true, at);
}
assert_eq!(
limiter.try_acquire(KEY, 800),
Acquire::Wait(200 + WAIT_MARGIN_MS)
);
let limiter = RateLimiter::default();
for at in [0, 400, 800] {
limiter.record(KEY, 0, 2, 1_000, true, at);
}
limiter.record(KEY, 0, 0, 1_000, true, 800);
assert_eq!(limiter.try_acquire(KEY, 1_000), Acquire::Granted);
assert_eq!(
limiter.try_acquire(KEY, 1_000),
Acquire::Wait(400 + WAIT_MARGIN_MS)
);
}
#[test]
fn test_retry_after_blocks_the_group() {
let limiter = RateLimiter::default();
limiter.record(KEY, 100, 2, 10_000, true, 0);
limiter.block_until(KEY, 5_000);
assert_eq!(
limiter.try_acquire(KEY, 1_000),
Acquire::Wait(4_000 + WAIT_MARGIN_MS)
);
assert_eq!(limiter.try_acquire(KEY, 5_000), Acquire::Granted);
}
#[test]
fn test_buckets_are_independent_per_group_and_token() {
let limiter = RateLimiter::default();
limiter.record("a|x", 0, 2, 10_000, true, 0);
assert!(matches!(limiter.try_acquire("a|x", 1), Acquire::Wait(_)));
assert_eq!(limiter.try_acquire("a|y", 1), Acquire::Granted);
assert_eq!(limiter.try_acquire("b|x", 1), Acquire::Granted);
assert_eq!(RateLimiter::key("a", Some("x")), "a|x");
assert_eq!(RateLimiter::group_of("a|x"), "a");
}
#[test]
fn test_idle_buckets_are_dropped() {
let limiter = RateLimiter::default();
limiter.record("old|x", 10, 2, 1_000, true, 0);
limiter.record("new|x", 10, 2, 1_000, true, 5_000);
assert_eq!(limiter.lock().len(), 1);
}
#[test]
fn test_a_permit_releases_on_drop() {
let limiter = Arc::new(RateLimiter::default());
limiter.record(KEY, 2, 2, 10_000, true, 0);
assert_eq!(limiter.try_acquire(KEY, 1), Acquire::Granted);
let permit = Permit::new(Arc::clone(&limiter), KEY);
assert!(matches!(limiter.try_acquire(KEY, 1), Acquire::Wait(_)));
drop(permit);
assert_eq!(limiter.try_acquire(KEY, 1), Acquire::Granted);
}
#[test]
fn test_the_policy_serializes_in_snake_case() {
assert_eq!(
serde_json::to_string(&RateLimitPolicy::Off).unwrap(),
"\"off\""
);
let wait = RateLimitPolicy::Wait {
max_wait: Duration::from_secs(3),
};
let json = serde_json::to_string(&wait).unwrap();
assert_eq!(
serde_json::from_str::<RateLimitPolicy>(&json).unwrap(),
wait
);
}
}