from __future__ import annotations import time import unittest from collections.abc import Callable class TokenBucket: """Token-bucket rate limiter using a monotonic clock.""" def __init__( self, rate: float, capacity: float, clock: Callable[[], float] = time.monotonic, ) -> None: if rate <= 0: raise ValueError("rate must be positive") if capacity <= 0: raise ValueError("capacity must be positive") self.rate = rate self.capacity = capacity self._tokens = capacity self._clock = clock self._updated_at = clock() def _refill(self) -> None: now = self._clock() elapsed = max(0.0, now - self._updated_at) self._tokens = min(self.capacity, self._tokens + elapsed * self.rate) self._updated_at = now def consume(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 self._refill() 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: raise ValueError("tokens cannot exceed capacity") self._refill() return max(0.0, (tokens - self._tokens) / self.rate) class TokenBucketTests(unittest.TestCase): def setUp(self) -> None: self.now = 0.0 self.bucket = TokenBucket(2.0, 3.0, lambda: self.now) def test_capacity_is_enforced(self) -> None: self.assertTrue(self.bucket.consume(3.0)) self.assertFalse(self.bucket.consume()) def test_tokens_refill_over_time(self) -> None: self.assertTrue(self.bucket.consume(3.0)) self.now += 0.5 self.assertTrue(self.bucket.consume()) self.assertFalse(self.bucket.consume()) self.assertAlmostEqual(self.bucket.retry_after(), 0.5) if __name__ == "__main__": unittest.main()