#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TokenBucket {
capacity_milli: u64,
refill_per_second: u64,
available_milli: u64,
last_ms: u64,
}
impl TokenBucket {
#[must_use]
pub fn new(capacity: u32, refill_per_second: u32, now_ms: u64) -> Self {
let capacity_milli = u64::from(capacity) * 1000;
Self {
capacity_milli,
refill_per_second: u64::from(refill_per_second),
available_milli: capacity_milli,
last_ms: now_ms,
}
}
pub fn try_acquire(&mut self, now_ms: u64) -> bool {
self.try_acquire_n(now_ms, 1)
}
pub fn try_acquire_n(&mut self, now_ms: u64, tokens: u32) -> bool {
self.refill(now_ms);
let cost = u64::from(tokens) * 1000;
if self.available_milli < cost {
return false;
}
self.available_milli -= cost;
true
}
pub fn available(&mut self, now_ms: u64) -> u32 {
self.refill(now_ms);
u32::try_from(self.available_milli / 1000).unwrap_or(u32::MAX)
}
fn refill(&mut self, now_ms: u64) {
let elapsed = now_ms.saturating_sub(self.last_ms);
self.last_ms = self.last_ms.max(now_ms);
let earned = elapsed.saturating_mul(self.refill_per_second);
self.available_milli = self
.available_milli
.saturating_add(earned)
.min(self.capacity_milli);
}
}
#[cfg(test)]
#[path = "token_bucket.test.rs"]
mod tests;
#[cfg(test)]
#[path = "token_bucket.spec.rs"]
mod spec;