use std::sync::Mutex; use std::time::Instant; /// A thread-safe token-bucket rate limiter. /// /// Tokens accumulate at `refill_per_second` up to `capacity`. /// Each successful acquisition consumes the requested number of tokens. #[derive(Debug)] pub struct TokenBucket { capacity: f64, refill_per_second: f64, state: Mutex, } #[derive(Debug)] struct State { tokens: f64, last_refill: Instant, } impl TokenBucket { /// Creates a full token bucket. pub fn new(capacity: u64, refill_per_second: f64) -> Self { assert!(capacity > 0, "capacity must be greater than zero"); assert!( refill_per_second.is_finite() && refill_per_second > 0.0, "refill rate must be finite and greater than zero" ); Self { capacity: capacity as f64, refill_per_second, state: Mutex::new(State { tokens: capacity as f64, last_refill: Instant::now(), }), } } /// Attempts to consume `amount` tokens without blocking. pub fn try_acquire(&self, amount: u64) -> bool { self.try_acquire_at(amount, Instant::now()) } fn try_acquire_at(&self, amount: u64, now: Instant) -> bool { if amount == 0 { return true; } let required = amount as f64; if required > self.capacity { return false; } let mut state = self.state.lock().unwrap_or_else(|e| e.into_inner()); self.refill(&mut state, now); if state.tokens >= required { state.tokens -= required; true } else { false } } fn refill(&self, state: &mut State, now: Instant) { if let Some(elapsed) = now.checked_duration_since(state.last_refill) { state.tokens = (state.tokens + elapsed.as_secs_f64() * self.refill_per_second) .min(self.capacity); state.last_refill = now; } } } #[cfg(test)] mod tests { use super::*; use std::time::Duration; #[test] fn consumes_until_empty() { let bucket = TokenBucket::new(3, 1.0); assert!(bucket.try_acquire(2)); assert!(bucket.try_acquire(1)); assert!(!bucket.try_acquire(1)); assert!(bucket.try_acquire(0)); } #[test] fn refills_over_time_and_respects_capacity() { let bucket = TokenBucket::new(4, 2.0); let start = bucket.state.lock().unwrap().last_refill; assert!(bucket.try_acquire_at(4, start)); assert!(!bucket.try_acquire_at(1, start)); let one_second_later = start + Duration::from_secs(1); assert!(bucket.try_acquire_at(2, one_second_later)); assert!(!bucket.try_acquire_at(1, one_second_later)); let much_later = start + Duration::from_secs(10); assert!(bucket.try_acquire_at(4, much_later)); assert!(!bucket.try_acquire_at(1, much_later)); } }