"""Unit tests for the token-bucket rate limiter.""" from __future__ import annotations import asyncio import pytest from contract_check.core.rate_limit import MemoryRateLimiter @pytest.fixture def limiter() -> MemoryRateLimiter: return MemoryRateLimiter() async def test_memory_rate_limiter_allows_first_requests(limiter: MemoryRateLimiter) -> None: # With limit=3, the first 3 requests should be allowed immediately. results = [await limiter.allow("key-a", 3) for _ in range(3)] assert all(r.allowed for r in results) assert all(r.retry_after_sec is None for r in results) async def test_memory_rate_limiter_blocks_excess(limiter: MemoryRateLimiter) -> None: # First 3 allowed, 4th blocked with a retry_after hint. for _ in range(3): assert (await limiter.allow("key-b", 3)).allowed blocked = await limiter.allow("key-b", 3) assert not blocked.allowed assert blocked.retry_after_sec is not None assert 0 < blocked.retry_after_sec <= 1 async def test_memory_rate_limiter_replenishes_after_time(limiter: MemoryRateLimiter) -> None: # Burn the bucket dry. for _ in range(3): assert (await limiter.allow("key-c", 3)).allowed assert not (await limiter.allow("key-c", 3)).allowed # Wait long enough for one token to refill at 3 tokens/second. await asyncio.sleep(0.5) assert (await limiter.allow("key-c", 3)).allowed async def test_memory_rate_limiter_keys_are_isolated(limiter: MemoryRateLimiter) -> None: assert (await limiter.allow("key-1", 1)).allowed assert not (await limiter.allow("key-1", 1)).allowed assert (await limiter.allow("key-2", 1)).allowed async def test_memory_rate_limiter_zero_limit_is_allowed(limiter: MemoryRateLimiter) -> None: # RedisRateLimiter treats non-positive limits as misconfigured and allows. assert (await limiter.allow("key-z", 0)).allowed