use std::sync::atomic::{AtomicU32, Ordering};
const SCALE: u32 = 1_000;
#[derive(Debug)]
pub(crate) struct TokenBucket {
tokens: AtomicU32,
max_scaled_tokens: u32,
refill_amount: u32,
}
impl TokenBucket {
pub(crate) fn new(max_tokens: u32, refill_ratio: f32) -> Self {
let max_scaled_tokens = max_tokens * SCALE;
let refill_amount = (refill_ratio * SCALE as f32).round() as u32;
Self {
tokens: AtomicU32::new(0),
max_scaled_tokens,
refill_amount,
}
}
pub(crate) fn try_acquire(&self) -> bool {
self.tokens
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
if current >= SCALE {
Some(current - SCALE)
} else {
None
}
})
.is_ok()
}
pub(crate) fn refill(&self) {
let _ = self
.tokens
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
if current >= self.max_scaled_tokens {
None } else {
Some(
current
.saturating_add(self.refill_amount)
.min(self.max_scaled_tokens),
)
}
});
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::task::{JoinSet, yield_now};
impl TokenBucket {
fn available_scaled_tokens(&self) -> u32 {
self.tokens.load(Ordering::Acquire)
}
}
#[test]
fn starts_empty_and_cannot_acquire() {
let bucket = TokenBucket::new(10, 0.1);
assert_eq!(bucket.available_scaled_tokens(), 0);
assert!(!bucket.try_acquire());
}
#[test]
fn refill_and_acquire() {
let bucket = TokenBucket::new(10, 0.1);
bucket.refill();
assert_eq!(bucket.available_scaled_tokens(), 100);
assert!(!bucket.try_acquire());
for _ in 0..9 {
bucket.refill();
}
assert_eq!(bucket.available_scaled_tokens(), 1000);
assert!(bucket.try_acquire());
assert_eq!(bucket.available_scaled_tokens(), 0);
assert!(!bucket.try_acquire());
}
#[test]
fn capped_at_max_tokens() {
let bucket = TokenBucket::new(2, 0.5); for _ in 0..5 {
bucket.refill();
}
assert_eq!(bucket.available_scaled_tokens(), 2000);
for _ in 0..5 {
bucket.refill();
}
assert_eq!(bucket.available_scaled_tokens(), 2000);
assert!(bucket.try_acquire());
assert_eq!(bucket.available_scaled_tokens(), 1000);
assert!(bucket.try_acquire());
assert_eq!(bucket.available_scaled_tokens(), 0);
assert!(!bucket.try_acquire());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_refill_and_acquire() {
use std::sync::Arc;
let bucket = Arc::new(TokenBucket::new(1000, 1.0));
let num_tasks = 4;
let tokens_per_task = 250;
let mut refillers = JoinSet::new();
for _ in 0..num_tasks {
let bucket = Arc::clone(&bucket);
refillers.spawn(async move {
for _ in 0..tokens_per_task {
bucket.refill();
}
yield_now().await;
});
}
let mut acquirers = JoinSet::new();
for _ in 0..num_tasks {
let bucket = Arc::clone(&bucket);
acquirers.spawn(async move {
let mut total_acquired = 0;
loop {
if total_acquired == tokens_per_task {
break;
}
if bucket.try_acquire() {
total_acquired += 1;
}
yield_now().await;
}
});
}
refillers.join_all().await;
acquirers.join_all().await;
assert_eq!(
bucket.available_scaled_tokens(),
0,
"all tokens should have been acquired"
);
}
}