use std::sync::Mutex; use std::time::Instant; /// A thread-safe token-bucket rate limiter. pub struct TokenBucket { capacity: f64, refill_per_second: f64, state: Mutex, } struct State { tokens: f64, last_refill: Instant, } impl TokenBucket { /// Creates a full bucket with the given capacity and refill rate. pub fn new(capacity: u64, refill_per_second: f64) -> Self { assert!(capacity > 0, "capacity must be positive"); assert!( refill_per_second.is_finite() && refill_per_second > 0.0, "refill rate must be finite and positive" ); Self { capacity: capacity as f64, refill_per_second, state: Mutex::new(State { tokens: capacity as f64, last_refill: Instant::now(), }), } } /// Attempts to consume one token without blocking. pub fn allow(&self) -> bool { self.try_take(1) } /// Attempts to consume `amount` tokens without blocking. pub fn try_take(&self, amount: u64) -> bool { if amount == 0 { return true; } let mut state = self.state.lock().unwrap_or_else(|e| e.into_inner()); self.refill(&mut state); let requested = amount as f64; if requested <= state.tokens { state.tokens -= requested; true } else { false } } /// Returns the currently available whole-token count. pub fn available(&self) -> u64 { let mut state = self.state.lock().unwrap_or_else(|e| e.into_inner()); self.refill(&mut state); state.tokens.floor() as u64 } fn refill(&self, state: &mut State) { let now = Instant::now(); let elapsed = now.duration_since(state.last_refill).as_secs_f64(); state.tokens = (state.tokens + elapsed * 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(2, 1.0); assert!(bucket.allow()); assert!(bucket.allow()); assert!(!bucket.allow()); assert_eq!(bucket.available(), 0); } #[test] fn refills_over_time_without_exceeding_capacity() { let bucket = TokenBucket::new(3, 10.0); assert!(bucket.try_take(3)); { let mut state = bucket.state.lock().unwrap(); state.last_refill -= Duration::from_millis(250); } assert_eq!(bucket.available(), 2); { let mut state = bucket.state.lock().unwrap(); state.last_refill -= Duration::from_secs(10); } assert_eq!(bucket.available(), 3); } }