use std::sync::Mutex;
use std::time::Instant;
pub struct TokenBucket {
inner: Mutex<Inner>,
}
struct Inner {
burst: f64,
refill: f64, tokens: f64, last: Instant,
}
impl TokenBucket {
pub fn new(burst: u64, refill_per_sec: u64) -> Self {
assert!(burst > 0, "TokenBucket::new: burst must be > 0");
TokenBucket {
inner: Mutex::new(Inner {
burst: burst as f64,
refill: refill_per_sec as f64,
tokens: burst as f64,
last: Instant::now(),
}),
}
}
fn accrue(inner: &mut Inner, now: Instant) {
if inner.refill == 0.0 { return; }
let dt = now.saturating_duration_since(inner.last).as_secs_f64();
if dt <= 0.0 { return; }
inner.tokens = (inner.tokens + dt * inner.refill).min(inner.burst);
inner.last = now;
}
pub fn take(&self, n: u64) -> bool {
if n == 0 { return true; }
let mut inner = self.inner.lock().expect("token-bucket mutex poisoned");
let now = Instant::now();
Self::accrue(&mut inner, now);
if inner.tokens < n as f64 {
return false;
}
inner.tokens -= n as f64;
true
}
pub fn refill(&self, n: u64) {
if n == 0 { return; }
let mut inner = self.inner.lock().expect("token-bucket mutex poisoned");
inner.tokens = (inner.tokens + n as f64).min(inner.burst);
}
pub fn peek(&self) -> u64 {
let mut inner = self.inner.lock().expect("token-bucket mutex poisoned");
let now = Instant::now();
Self::accrue(&mut inner, now);
inner.tokens as u64
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::thread;
use std::time::Duration;
#[test]
fn init_starts_full() {
let b = TokenBucket::new(100, 10);
assert_eq!(b.peek(), 100);
}
#[test]
fn take_drains() {
let b = TokenBucket::new(10, 0);
for _ in 0..10 { assert!(b.take(1)); }
assert!(!b.take(1));
assert_eq!(b.peek(), 0);
}
#[test]
fn refill_clips_at_burst() {
let b = TokenBucket::new(5, 0);
assert!(b.take(5));
assert_eq!(b.peek(), 0);
b.refill(1000);
assert_eq!(b.peek(), 5);
}
#[test]
fn time_refill_accrues() {
let b = TokenBucket::new(100, 100);
assert!(b.take(100));
assert_eq!(b.peek(), 0);
thread::sleep(Duration::from_millis(250));
let p = b.peek();
assert!(p >= 20, "after >=250ms at 100/s, peek = {} (expected >= 20)", p);
assert!(p <= 100, "burst cap violated, peek = {} (expected <= 100)", p);
assert!(b.take(10));
}
#[test]
fn thread_safety_exact_balance() {
let b = Arc::new(TokenBucket::new(500, 0));
let mut handles = Vec::new();
for _ in 0..8 {
let b = Arc::clone(&b);
handles.push(thread::spawn(move || {
let mut ok = 0;
for _ in 0..1000 {
if b.take(1) { ok += 1; }
}
ok
}));
}
let total: u32 = handles.into_iter().map(|h| h.join().unwrap()).sum();
assert_eq!(total, 500, "expected exactly 500 takes across 8 threads");
assert_eq!(b.peek(), 0);
}
#[test]
fn zero_cost_take_is_free() {
let b = TokenBucket::new(1, 0);
assert!(b.take(1));
assert_eq!(b.peek(), 0);
assert!(b.take(0)); assert_eq!(b.peek(), 0);
}
#[test]
#[should_panic(expected = "burst must be > 0")]
fn zero_burst_panics() {
let _ = TokenBucket::new(0, 1);
}
}