from __future__ import annotations import threading import time import unittest from collections.abc import Callable class TokenBucket: """Thread-safe token-bucket rate limiter.""" def __init__( self, capacity: float, refill_rate: float, *, clock: Callable[[], float] = time.monotonic, ) -> None: if capacity <= 0: raise ValueError("capacity must be positive") if refill_rate <= 0: raise ValueError("refill_rate must be positive") self.capacity = float(capacity) self.refill_rate = float(refill_rate) self._tokens = self.capacity self._clock = clock self._updated_at = clock() self._lock = threading.Lock() def _refill(self, now: float) -> None: elapsed = max(0.0, now - self._updated_at) self._tokens = min( self.capacity, self._tokens + elapsed * self.refill_rate, ) self._updated_at = now def try_acquire(self, tokens: float = 1.0) -> bool: """Consume tokens immediately, returning False if too few are available.""" if tokens <= 0: raise ValueError("tokens must be positive") if tokens > self.capacity: return False with self._lock: self._refill(self._clock()) if self._tokens < tokens: return False self._tokens -= tokens return True def wait_time(self, tokens: float = 1.0) -> float: """Return the seconds until the requested tokens will be available.""" if tokens <= 0: raise ValueError("tokens must be positive") if tokens > self.capacity: raise ValueError("requested tokens exceed bucket capacity") with self._lock: self._refill(self._clock()) return max(0.0, (tokens - self._tokens) / self.refill_rate) def acquire(self, tokens: float = 1.0) -> None: """Block until the requested tokens can be consumed.""" if tokens <= 0: raise ValueError("tokens must be positive") if tokens > self.capacity: raise ValueError("requested tokens exceed bucket capacity") while not self.try_acquire(tokens): time.sleep(self.wait_time(tokens)) class FakeClock: def __init__(self) -> None: self.now = 0.0 def __call__(self) -> float: return self.now def advance(self, seconds: float) -> None: self.now += seconds class TokenBucketTests(unittest.TestCase): def test_consumption_and_exhaustion(self) -> None: bucket = TokenBucket(capacity=2, refill_rate=1) self.assertTrue(bucket.try_acquire(2)) self.assertFalse(bucket.try_acquire()) def test_tokens_refill_over_time(self) -> None: clock = FakeClock() bucket = TokenBucket(capacity=3, refill_rate=2, clock=clock) self.assertTrue(bucket.try_acquire(3)) self.assertFalse(bucket.try_acquire()) clock.advance(0.5) self.assertTrue(bucket.try_acquire()) self.assertAlmostEqual(bucket.wait_time(), 0.5) if __name__ == "__main__": unittest.main()