from __future__ import annotations from unittest.mock import AsyncMock from uuid import UUID import pytest from yuxi.channel.security.rate_limit import RateLimiter @pytest.fixture def redis() -> AsyncMock: return AsyncMock() @pytest.fixture def limiter(redis: AsyncMock) -> RateLimiter: return RateLimiter(redis) @pytest.mark.unit class TestRateLimiter: async def test_is_allowed_when_redis_returns_one(self, limiter: RateLimiter, redis: AsyncMock, monkeypatch) -> None: redis.eval.return_value = 1 monkeypatch.setattr("yuxi.channel.security.rate_limit.time.time", lambda: 1000.0) monkeypatch.setattr( "yuxi.channel.security.rate_limit.uuid.uuid4", lambda: UUID("12345678-1234-5678-1234-567812345678"), ) allowed = await limiter.is_allowed("rate:key", max_requests=5, window_seconds=60) assert allowed is True redis.eval.assert_awaited_once() script, num_keys, key, window_start, max_requests_arg, now, member, window_seconds = redis.eval.call_args.args assert script == RateLimiter._ALLOW_SCRIPT assert num_keys == 1 assert key == "rate:key" assert window_start == 940.0 assert max_requests_arg == 5 assert now == 1000.0 assert member == "1000.0:12345678123456781234567812345678" assert window_seconds == 60 async def test_is_not_allowed_when_redis_returns_zero(self, limiter: RateLimiter, redis: AsyncMock) -> None: redis.eval.return_value = 0 allowed = await limiter.is_allowed("rate:key", max_requests=5, window_seconds=60) assert allowed is False redis.eval.assert_awaited_once() async def test_zero_max_requests_always_blocked(self, limiter: RateLimiter, redis: AsyncMock) -> None: redis.eval.return_value = 0 allowed = await limiter.is_allowed("rate:key", max_requests=0, window_seconds=60) assert allowed is False async def test_window_seconds_passed_to_redis(self, limiter: RateLimiter, redis: AsyncMock) -> None: redis.eval.return_value = 1 await limiter.is_allowed("rate:key", max_requests=10, window_seconds=120) _, _, _, _, _, _, _, window_seconds = redis.eval.call_args.args assert window_seconds == 120 async def test_redis_client_stored_on_instance(self, redis: AsyncMock) -> None: limiter = RateLimiter(redis) assert limiter.redis is redis