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 or refill_rate <= 0: raise ValueError("capacity and refill_rate must be positive") self.capacity = float(capacity) self.refill_rate = float(refill_rate) self._tokens = float(capacity) self._clock = clock self._last_refill = clock() self._lock = threading.Lock() def _refill(self, now: float) -> None: elapsed = max(0.0, now - self._last_refill) self._tokens = min( self.capacity, self._tokens + elapsed * self.refill_rate, ) self._last_refill = now def try_acquire(self, tokens: float = 1.0) -> bool: """Consume tokens immediately, returning False if insufficient.""" 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 retry_after(self, tokens: float = 1.0) -> float: """Return seconds until the requested tokens should be available.""" if tokens <= 0: raise ValueError("tokens must be positive") if tokens > self.capacity: return float("inf") with self._lock: self._refill(self._clock()) return max(0.0, (tokens - self._tokens) / self.refill_rate) class FakeClock: def __init__(self) -> None: self.now = 0.0 def __call__(self) -> float: return self.now class TokenBucketTests(unittest.TestCase): def test_capacity_is_enforced(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)) clock.now += 0.5 self.assertTrue(bucket.try_acquire()) self.assertFalse(bucket.try_acquire()) self.assertAlmostEqual(bucket.retry_after(), 0.5) if __name__ == "__main__": unittest.main()