from __future__ import annotations import unittest from unittest.mock import patch from redis.exceptions import RedisError from govoplan_mail.backend.sending import rate_limit class RateLimitTests(unittest.TestCase): def setUp(self) -> None: rate_limit._local_next_allowed.clear() def test_local_fallback_waits_between_sends(self) -> None: now = [100.0] sleeps: list[float] = [] def fake_time() -> float: return now[0] def fake_sleep(seconds: float) -> None: sleeps.append(seconds) now[0] += seconds with ( patch.object(rate_limit, "_distributed_rate_limit_enabled", return_value=False), patch.object(rate_limit.time, "time", fake_time), patch.object(rate_limit.time, "sleep", fake_sleep), ): first = rate_limit.wait_for_rate_limit(key="tenant:t:campaign:c", messages_per_minute=4) second = rate_limit.wait_for_rate_limit(key="tenant:t:campaign:c", messages_per_minute=4) self.assertEqual(first.waited_seconds, 0.0) self.assertEqual(second.waited_seconds, 15.0) self.assertEqual(sleeps, [15.0]) def test_disabled_rate_limit_does_not_reserve_local_slot(self) -> None: now = [100.0] sleeps: list[float] = [] with ( patch.object(rate_limit, "_distributed_rate_limit_enabled", return_value=False), patch.object(rate_limit.time, "time", lambda: now[0]), patch.object(rate_limit.time, "sleep", lambda seconds: sleeps.append(seconds)), ): disabled = rate_limit.wait_for_rate_limit(key="k", messages_per_minute=4, enabled=False) enabled = rate_limit.wait_for_rate_limit(key="k", messages_per_minute=4) self.assertEqual(disabled.waited_seconds, 0.0) self.assertEqual(enabled.waited_seconds, 0.0) self.assertEqual(sleeps, []) def test_redis_failure_falls_back_to_local_wait(self) -> None: now = [100.0] sleeps: list[float] = [] def fail_redis_client(): raise RedisError("redis unavailable") def fake_sleep(seconds: float) -> None: sleeps.append(seconds) now[0] += seconds with ( patch.object(rate_limit, "_distributed_rate_limit_enabled", return_value=True), patch.object(rate_limit, "_redis_client", fail_redis_client), patch.object(rate_limit.time, "time", lambda: now[0]), patch.object(rate_limit.time, "sleep", fake_sleep), ): first = rate_limit.wait_for_rate_limit(key="k", messages_per_minute=60) second = rate_limit.wait_for_rate_limit(key="k", messages_per_minute=60) self.assertEqual(first.waited_seconds, 0.0) self.assertEqual(second.waited_seconds, 1.0) self.assertEqual(sleeps, [1.0]) if __name__ == "__main__": unittest.main()