use std::num::NonZeroU32;
use std::sync::Arc;
use std::time::{Duration, Instant};
use serde::{Deserialize, Serialize};
use gcra::{DefaultDirectRateLimiter, Jitter, Quota};
#[derive(Debug, Clone, Copy, PartialEq, Default, Serialize, Deserialize)]
pub struct RateLimitConfig {
#[serde(default)]
pub rps: u32,
#[serde(default)]
pub burst: u32,
}
impl RateLimitConfig {
#[must_use]
pub fn per_second(rps: u32) -> Self {
Self { rps, burst: rps }
}
#[must_use]
pub fn with_burst(mut self, burst: u32) -> Self {
self.burst = burst;
self
}
#[must_use]
pub fn is_enabled(&self) -> bool {
self.rps > 0
}
#[must_use]
fn capacity(self) -> u32 {
let cap = if self.burst == 0 {
self.rps
} else {
self.burst
};
cap.max(1)
}
}
#[derive(Clone)]
pub struct RateLimiter {
inner: Option<Arc<Inner>>,
label: &'static str,
}
struct Inner {
limiter: DefaultDirectRateLimiter,
jitter: Jitter,
}
impl RateLimiter {
#[must_use]
pub fn new(config: RateLimitConfig, label: &'static str) -> Self {
if !config.is_enabled() {
return Self { inner: None, label };
}
let rps = NonZeroU32::new(config.rps).unwrap_or(NonZeroU32::MIN);
let burst = NonZeroU32::new(config.capacity()).unwrap_or(NonZeroU32::MIN);
let quota = Quota::per_second(rps).allow_burst(burst);
let limiter = DefaultDirectRateLimiter::direct(quota);
let period = Duration::from_secs_f64(1.0 / f64::from(config.rps));
let jitter = Jitter::up_to(period);
Self {
inner: Some(Arc::new(Inner { limiter, jitter })),
label,
}
}
#[must_use]
pub fn is_enabled(&self) -> bool {
self.inner.is_some()
}
#[must_use]
pub fn try_acquire(&self) -> bool {
match &self.inner {
None => true,
Some(inner) => inner.limiter.check().is_ok(),
}
}
pub async fn acquire(&self) -> Duration {
let Some(inner) = &self.inner else {
return Duration::ZERO;
};
let started = Instant::now();
inner.limiter.until_ready_with_jitter(inner.jitter).await;
let waited = started.elapsed();
self.record(waited);
waited
}
#[cfg_attr(not(feature = "metrics"), allow(unused_variables))]
fn record(&self, waited: Duration) {
#[cfg(feature = "metrics")]
if waited > Duration::from_millis(1) {
::metrics::counter!("rate_limited_total", "limiter" => self.label).increment(1);
::metrics::histogram!("rate_limit_wait_seconds", "limiter" => self.label)
.record(waited.as_secs_f64());
}
}
}
impl std::fmt::Debug for RateLimiter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RateLimiter")
.field("label", &self.label)
.field("enabled", &self.is_enabled())
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn disabled_limiter_is_transparent() {
let rl = RateLimiter::new(RateLimitConfig::default(), "test");
assert!(!rl.is_enabled());
for _ in 0..1000 {
assert!(rl.try_acquire());
}
}
#[test]
fn try_acquire_drains_then_refuses_within_burst() {
let rl = RateLimiter::new(RateLimitConfig::per_second(1).with_burst(3), "test");
assert!(rl.try_acquire());
assert!(rl.try_acquire());
assert!(rl.try_acquire());
assert!(
!rl.try_acquire(),
"bucket exhausted -- fourth try must fail without waiting"
);
}
#[test]
fn config_capacity_defaults_burst_to_rps() {
let c = RateLimitConfig::per_second(50);
assert_eq!(c.burst, 50);
assert!(c.is_enabled());
let c2 = RateLimitConfig { rps: 10, burst: 0 };
assert_eq!(c2.capacity(), 10);
}
#[tokio::test]
async fn acquire_waits_when_empty_then_admits_after_refill() {
let rl = RateLimiter::new(RateLimitConfig::per_second(100).with_burst(1), "test");
let first = rl.acquire().await;
assert!(
first < Duration::from_millis(20),
"cold bucket admits ~immediately (jitter aside), got {first:?}"
);
let _ = rl.try_acquire();
let waited = rl.acquire().await;
assert!(
waited >= Duration::from_millis(8),
"empty bucket must wait ~10ms+ for a refill, waited {waited:?}"
);
}
#[tokio::test]
async fn enabled_limiter_paces_to_rate() {
let rl = RateLimiter::new(RateLimitConfig::per_second(100).with_burst(1), "test");
let _ = rl.acquire().await; let start = std::time::Instant::now();
for _ in 0..5 {
rl.acquire().await;
}
assert!(
start.elapsed() >= Duration::from_millis(40),
"5 acquires at 100rps must take ~50ms+, took {:?}",
start.elapsed()
);
}
}