use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum Acquire {
Ok,
Retry(Duration),
Unattainable { burst_capacity: u64 },
}
pub struct RateLimiter {
tat_ns: AtomicU64,
period_ns: u64,
burst_ns: u64,
origin: Instant,
}
impl RateLimiter {
pub fn new(rate_per_sec: f64, burst_capacity: u64) -> Self {
let period_ns = (1_000_000_000.0 / rate_per_sec) as u64;
let burst_ns = period_ns.saturating_mul(burst_capacity.max(1));
Self {
tat_ns: AtomicU64::new(0),
period_ns,
burst_ns,
origin: Instant::now(),
}
}
pub fn try_acquire(&self) -> bool {
self.try_acquire_at(self.now_ns())
}
pub fn try_acquire_at(&self, now: u64) -> bool {
loop {
let tat = self.tat_ns.load(Ordering::Acquire);
let new_tat = tat.max(now).saturating_add(self.period_ns);
if new_tat.saturating_sub(now) > self.burst_ns {
return false;
}
match self.tat_ns.compare_exchange_weak(
tat,
new_tat,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return true,
Err(_) => continue,
}
}
}
pub fn try_acquire_with_retry(&self) -> Acquire {
self.try_acquire_with_retry_at(self.now_ns())
}
pub fn try_acquire_with_retry_at(&self, now: u64) -> Acquire {
self.try_acquire_n_with_retry_at(now, 1)
}
pub fn try_acquire_n(&self, n: u64) -> bool {
self.try_acquire_n_at(self.now_ns(), n)
}
pub fn try_acquire_n_at(&self, now: u64, n: u64) -> bool {
matches!(self.try_acquire_n_with_retry_at(now, n), Acquire::Ok)
}
pub fn try_acquire_n_with_retry(&self, n: u64) -> Acquire {
self.try_acquire_n_with_retry_at(self.now_ns(), n)
}
pub fn try_acquire_n_with_retry_at(&self, now: u64, n: u64) -> Acquire {
if n == 0 {
return Acquire::Ok;
}
let cost = self.period_ns.saturating_mul(n);
if cost > self.burst_ns {
return Acquire::Unattainable {
burst_capacity: self.burst_capacity(),
};
}
loop {
let tat = self.tat_ns.load(Ordering::Acquire);
let new_tat = tat.max(now).saturating_add(cost);
if new_tat.saturating_sub(now) > self.burst_ns {
let wait = new_tat.saturating_sub(self.burst_ns).saturating_sub(now);
return Acquire::Retry(Duration::from_nanos(wait));
}
match self.tat_ns.compare_exchange_weak(
tat,
new_tat,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return Acquire::Ok,
Err(_) => continue,
}
}
}
pub fn time_until_ready(&self, n: u64) -> Option<Duration> {
self.time_until_ready_at(self.now_ns(), n)
}
pub fn time_until_ready_at(&self, now: u64, n: u64) -> Option<Duration> {
if n == 0 {
return Some(Duration::ZERO);
}
let cost = self.period_ns.saturating_mul(n);
if cost > self.burst_ns {
return None;
}
let tat = self.tat_ns.load(Ordering::Acquire);
let new_tat = tat.max(now).saturating_add(cost);
if new_tat.saturating_sub(now) > self.burst_ns {
let wait = new_tat.saturating_sub(self.burst_ns).saturating_sub(now);
Some(Duration::from_nanos(wait))
} else {
Some(Duration::ZERO)
}
}
pub fn acquire_within(&self, n: u64, timeout: Duration) -> bool {
let deadline = Instant::now() + timeout;
loop {
match self.try_acquire_n_with_retry(n) {
Acquire::Ok => return true,
Acquire::Unattainable { .. } => return false,
Acquire::Retry(wait) => {
let remaining = deadline.saturating_duration_since(Instant::now());
if wait > remaining {
return false;
}
std::thread::sleep(wait);
}
}
}
}
pub fn reset(&self) {
self.tat_ns.store(0, Ordering::Release);
}
pub fn now_ns(&self) -> u64 {
self.origin.elapsed().as_nanos() as u64
}
pub fn rate_per_sec(&self) -> f64 {
1_000_000_000.0 / self.period_ns as f64
}
pub fn burst_capacity(&self) -> u64 {
self.burst_ns.checked_div(self.period_ns).unwrap_or(0)
}
}
#[cfg(feature = "harness")]
pub mod recipe;
#[cfg(any(
feature = "token-bucket",
feature = "hierarchical",
feature = "distributed-backend",
feature = "metrics",
feature = "keyed",
))]
pub mod features;
#[cfg(any(
feature = "token-bucket",
feature = "hierarchical",
feature = "distributed-backend",
feature = "metrics",
feature = "keyed",
))]
pub use features::clock::{Clock, SystemClock, TestClock};
#[cfg(feature = "distributed-backend")]
pub use features::distributed_backend::{Backend, DistributedLimiter, InMemoryBackend};
#[cfg(feature = "hierarchical")]
pub use features::hierarchical::HierarchicalLimiter;
#[cfg(feature = "keyed")]
pub use features::keyed::KeyedRateLimiter;
#[cfg(feature = "metrics")]
pub use features::metrics::{MeteredTokenBucket, MetricsSnapshot};
#[cfg(feature = "token-bucket")]
pub use features::token_bucket::TokenBucket;
#[cfg(test)]
#[path = "rate_limiter_tests.rs"]
mod rate_limiter_tests;
#[cfg(test)]
#[path = "sample_app_tests.rs"]
mod sample_app_tests;