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 = 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 when 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 the seconds until the requested tokens are 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()) missing = max(0.0, tokens - self._tokens) return missing / self.refill_rate 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_rejects_when_empty(self) -> None: clock = FakeClock() bucket = TokenBucket(capacity=2, refill_rate=1, clock=clock) self.assertTrue(bucket.try_acquire(2)) self.assertFalse(bucket.try_acquire()) def test_refills_over_time(self) -> None: clock = FakeClock() bucket = TokenBucket(capacity=2, refill_rate=2, clock=clock) self.assertTrue(bucket.try_acquire(2)) self.assertAlmostEqual(bucket.retry_after(), 0.5) clock.advance(0.5) self.assertTrue(bucket.try_acquire()) self.assertFalse(bucket.try_acquire()) if __name__ == "__main__": unittest.main()