diff --git a/tests/test_rate_limiter.py b/tests/test_rate_limiter.py index dc0e58e..4eaa9e0 100644 --- a/tests/test_rate_limiter.py +++ b/tests/test_rate_limiter.py @@ -1,45 +1,80 @@ """Тесты для utils/rate_limiter.py — проверка логики токен-бакета.""" import asyncio -import time +from unittest.mock import MagicMock from utils.rate_limiter import RateLimiter +def _make_time() -> tuple[RateLimiter, list[float]]: + """Создать RateLimiter с контролируемой временной функцией.""" + times: list[float] = [0.0] + + def controlled_time() -> float: + return times[0] + + limiter = RateLimiter(rate=10.0, burst=5, _time_func=controlled_time) + return limiter, times + + async def test_initial_tokens_full() -> None: """Бакет заполнен до burst при создании.""" - limiter = RateLimiter(rate=2.0, burst=5) + limiter, _ = _make_time() assert limiter.tokens == 5.0 async def test_acquire_consumes_token() -> None: """acquire() уменьшает количество токенов.""" - limiter = RateLimiter(rate=1.0, burst=3) + limiter, _ = _make_time() await limiter.acquire() - assert limiter.tokens == 2.0 + assert limiter.tokens == 4.0 async def test_acquire_waits_when_empty() -> None: - """acquire() ждёт, когда токены закончились.""" - limiter = RateLimiter(rate=10.0, burst=1) # 10 токенов/сек - await limiter.acquire() # бакет пуст - start = time.monotonic() - await limiter.acquire() # должен ждать ~0.1 сек - elapsed = time.monotonic() - start - assert elapsed >= 0.05 # допускаем погрешность + """acquire() ждёт, когда токены закончились (контролируемое время).""" + limiter, times = _make_time() + # Потратить все 5 токенов + for _ in range(5): + await limiter.acquire() + assert limiter.tokens < 1.0 + + # Пропустить 0.2 сек -> должно пополниться 2 токена (rate=10) + times[0] = 0.2 + async with limiter.lock: + limiter._refill() + assert limiter.tokens >= 2.0 async def test_burst_cap() -> None: """Токены не превышают burst после долгого простоя.""" - limiter = RateLimiter(rate=100.0, burst=3) - await asyncio.sleep(0.1) # теоретически +10 токенов, но cap = 3 + limiter, times = _make_time() + times[0] = 10.0 # теоретически +100 токенов, но cap = 5 async with limiter.lock: limiter._refill() - assert limiter.tokens == 3.0 + assert limiter.tokens == 5.0 async def test_multiple_acquire() -> None: """Можно забрать несколько токенов за раз.""" - limiter = RateLimiter(rate=1.0, burst=10) - await limiter.acquire(token=5) - assert limiter.tokens == 5.0 + limiter, _ = _make_time() + await limiter.acquire(token=3) + assert limiter.tokens == 2.0 + + +async def test_refill_partial() -> None: + """Пополнение за малый интервал времени.""" + limiter, times = _make_time() + times[0] = 0.1 # 10 токенов/сек * 0.1 сек = 1 токен + async with limiter.lock: + limiter._refill() + assert limiter.tokens == 5.0 # был 5 + 1 = 6, но cap = 5 + + +async def test_refill_exact() -> None: + """Точное пополнение при частичном бакете.""" + limiter, times = _make_time() + await limiter.acquire(token=3) # осталось 2 + times[0] = 0.1 # +1 токен + async with limiter.lock: + limiter._refill() + assert limiter.tokens == 3.0 # 2 + 1 = 3 diff --git a/utils/rate_limiter.py b/utils/rate_limiter.py index c95e055..18e6553 100644 --- a/utils/rate_limiter.py +++ b/utils/rate_limiter.py @@ -10,7 +10,7 @@ import asyncio import logging import os import time -from typing import Final +from typing import Callable, Final logger = logging.getLogger(__name__) @@ -18,21 +18,28 @@ logger = logging.getLogger(__name__) class RateLimiter: """Токен-бакет: заполняется со скоростью rate токенов/сек, максимум burst.""" - def __init__(self, rate: float, burst: int) -> None: + def __init__( + self, + rate: float, + burst: int, + _time_func: Callable[[], float] | None = None, + ) -> None: """ Args: rate: Скорость пополнения токенов (токенов в секунду). burst: Максимальный размер бакета. + _time_func: Функция получения времени (для тестов). По умолчанию time.monotonic. """ self.rate: float = rate self.burst: int = burst self.tokens: float = float(burst) self.lock: asyncio.Lock = asyncio.Lock() - self._last_refill: float = time.monotonic() + self._time_func = _time_func or time.monotonic + self._last_refill: float = self._time_func() def _refill(self) -> None: """Пополнить токены за прошедшее время.""" - now: float = time.monotonic() + now: float = self._time_func() elapsed: float = now - self._last_refill self.tokens = min(self.burst, self.tokens + elapsed * self.rate) self._last_refill = now