from collections.abc import Generator from unittest.mock import patch import fakeredis import pytest from freezegun import freeze_time from assets import rate_limiting from assets.rate_limiting import RateLimitedResource @pytest.fixture def use_fake_redis_client() -> Generator[fakeredis.FakeRedis, None, None]: fake_client = fakeredis.FakeRedis() with patch( "assets.rate_limiting.redis_connector.client", new=fake_client, ): rate_limiting.get_rate_limiter.cache_clear() yield fake_client rate_limiting.get_rate_limiter.cache_clear() def test_unconfigured_resources_allow_all_requests() -> None: rate_limiter = rate_limiting.get_rate_limiter() # Should allow unlimited calls for _ in range(100): assert rate_limiter.hit("asdf", "qwer") == rate_limiting.Allowed() # type: ignore[arg-type] assert rate_limiter.hit("asdf", "oa:179") == rate_limiting.Allowed() # type: ignore[arg-type] def test_rate_limiter_enforces_limit( use_fake_redis_client: fakeredis.FakeRedis, ) -> None: rate_limiter = rate_limiting.get_rate_limiter() limit = rate_limiting.RATE_LIMITS[ (rate_limiting.RateLimitedResource.UPLOAD, "oa:179") ] with freeze_time("2025-10-29 01:01:00") as frozen_time: # Calls within the limit should succeed for _ in range(limit.amount): assert ( rate_limiter.hit(RateLimitedResource.UPLOAD, "oa:179") == rate_limiting.Allowed() ) # Next call should fail assert rate_limiter.hit( RateLimitedResource.UPLOAD, "oa:179" ) == rate_limiting.Exceeded(f"Rate limit exceeded. Limit: {limit}.") # Advance time to 1 second before the next rate limit window frozen_time.tick(limit.get_expiry()) # Should still fail assert rate_limiter.hit( RateLimitedResource.UPLOAD, "oa:179" ) == rate_limiting.Exceeded(f"Rate limit exceeded. Limit: {limit}.") # Advance time by 1 more second to enter the next rate limit window frozen_time.tick(1) # Calls within the limit should succeed for _ in range(limit.amount): assert ( rate_limiter.hit(RateLimitedResource.UPLOAD, "oa:179") == rate_limiting.Allowed() ) # Next call should fail assert rate_limiter.hit( RateLimitedResource.UPLOAD, "oa:179" ) == rate_limiting.Exceeded(f"Rate limit exceeded. Limit: {limit}.") def test_rate_limiter_error() -> None: rate_limiter = rate_limiting.get_rate_limiter() with patch.object(rate_limiter._rate_limiter, "hit") as mock_hit: mock_hit.side_effect = Exception("some error") assert rate_limiter.hit( RateLimitedResource.UPLOAD, "oa:179" ) == rate_limiting.Error("Exception: some error")