use std::sync::atomic::{AtomicU64, Ordering};
pub struct TokenBucket {
capacity: u64,
refill_rate: f64,
tokens: AtomicU64,
last_refill: AtomicU64,
}
impl TokenBucket {
pub fn new(capacity: u64, refill_rate: f64) -> Self {
let now_ms = Self::now_ms();
Self {
capacity,
refill_rate,
tokens: AtomicU64::new(capacity * 1000), last_refill: AtomicU64::new(now_ms),
}
}
pub fn try_acquire(&self) -> bool {
self.try_acquire_n(1)
}
pub fn try_acquire_n(&self, n: u64) -> bool {
let now_ms = Self::now_ms();
let cost = n * 1000;
loop {
let last = self.last_refill.load(Ordering::Relaxed);
let current_tokens = self.tokens.load(Ordering::Relaxed);
let elapsed_ms = now_ms.saturating_sub(last);
let tokens_to_add = (elapsed_ms as f64 * self.refill_rate).round() as u64;
let new_tokens = (current_tokens + tokens_to_add).min(self.capacity * 1000);
if new_tokens < cost {
return false;
}
let final_tokens = new_tokens - cost;
if self
.tokens
.compare_exchange_weak(
current_tokens,
final_tokens,
Ordering::SeqCst,
Ordering::Relaxed,
)
.is_ok()
{
let _ = self.last_refill.compare_exchange(
last,
now_ms,
Ordering::Relaxed,
Ordering::Relaxed,
);
return true;
}
}
}
pub fn available(&self) -> u64 {
let now_ms = Self::now_ms();
let last = self.last_refill.load(Ordering::Relaxed);
let current_tokens = self.tokens.load(Ordering::Relaxed);
let elapsed_ms = now_ms.saturating_sub(last);
let tokens_to_add = (elapsed_ms as f64 * self.refill_rate).round() as u64;
(current_tokens + tokens_to_add).min(self.capacity * 1000) / 1000
}
fn now_ms() -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_token_bucket_basic() {
let bucket = TokenBucket::new(10, 1.0);
for _ in 0..10 {
assert!(bucket.try_acquire());
}
assert!(!bucket.try_acquire());
}
#[test]
fn test_token_bucket_burst() {
let bucket = TokenBucket::new(5, 10.0);
assert!(bucket.try_acquire_n(5));
assert!(!bucket.try_acquire());
}
}