70 lines
2.4 KiB
Python
70 lines
2.4 KiB
Python
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
|