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() @property def tokens(self) -> float: self._refill() return self._tokens def allow(self, cost: float = 1.0) -> bool: """Consume tokens and return True, or reject without consuming.""" if cost <= 0: raise ValueError("cost must be positive") if cost > self.capacity: return False self._refill() if self._tokens < cost: return False self._tokens -= cost return True def retry_after(self, cost: float = 1.0) -> float: """Return seconds until the requested cost can be consumed.""" if cost <= 0: raise ValueError("cost must be positive") if cost > self.capacity: return float("inf") self._refill() return max(0.0, (cost - self._tokens) / self.rate) 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 class TokenBucketTests(unittest.TestCase): def test_rejects_when_empty_and_refills_over_time(self) -> None: now = [0.0] bucket = TokenBucket(2.0, 2.0, lambda: now[0]) self.assertTrue(bucket.allow(2.0)) self.assertFalse(bucket.allow()) now[0] += 0.5 self.assertTrue(bucket.allow()) def test_retry_after(self) -> None: now = [0.0] bucket = TokenBucket(4.0, 2.0, lambda: now[0]) self.assertTrue(bucket.allow(2.0)) self.assertAlmostEqual(bucket.retry_after(), 0.25) if __name__ == "__main__": unittest.main()