131 lines
4.4 KiB
Python
131 lines
4.4 KiB
Python
from __future__ import annotations
|
|
|
|
import unittest
|
|
|
|
from redis.exceptions import RedisError
|
|
|
|
from govoplan_core.core.throttling import (
|
|
FixedWindowBucket,
|
|
FixedWindowThrottle,
|
|
InMemoryFixedWindowStore,
|
|
ResilientFixedWindowStore,
|
|
ThrottleDimension,
|
|
build_fixed_window_throttle,
|
|
)
|
|
|
|
|
|
class _Clock:
|
|
def __init__(self) -> None:
|
|
self.value = 100.0
|
|
|
|
def __call__(self) -> float:
|
|
return self.value
|
|
|
|
|
|
class _RecordingStore:
|
|
def __init__(self) -> None:
|
|
self.values: dict[str, int] = {}
|
|
|
|
def read(self, key: str) -> FixedWindowBucket:
|
|
return FixedWindowBucket(self.values.get(key, 0), 60)
|
|
|
|
def increment(self, key: str, *, window_seconds: int) -> FixedWindowBucket:
|
|
del window_seconds
|
|
self.values[key] = self.values.get(key, 0) + 1
|
|
return FixedWindowBucket(self.values[key], 60)
|
|
|
|
def delete(self, key: str) -> None:
|
|
self.values.pop(key, None)
|
|
|
|
|
|
class _FailingStore(_RecordingStore):
|
|
def read(self, key: str) -> FixedWindowBucket:
|
|
del key
|
|
raise RedisError("offline")
|
|
|
|
def increment(self, key: str, *, window_seconds: int) -> FixedWindowBucket:
|
|
del key, window_seconds
|
|
raise RedisError("offline")
|
|
|
|
def delete(self, key: str) -> None:
|
|
del key
|
|
raise RedisError("offline")
|
|
|
|
|
|
class FixedWindowThrottleTests(unittest.TestCase):
|
|
def test_blocks_at_limit_and_expires_without_sleeping(self) -> None:
|
|
clock = _Clock()
|
|
store = InMemoryFixedWindowStore(clock=clock)
|
|
throttle = FixedWindowThrottle(store, window_seconds=60)
|
|
dimensions = (ThrottleDimension("password", "secret-subject", 3),)
|
|
|
|
self.assertTrue(throttle.check(dimensions).allowed)
|
|
self.assertTrue(throttle.record(dimensions).allowed)
|
|
self.assertTrue(throttle.record(dimensions).allowed)
|
|
blocked = throttle.record(dimensions)
|
|
self.assertFalse(blocked.allowed)
|
|
self.assertEqual(blocked.retry_after_seconds, 60)
|
|
self.assertFalse(throttle.check(dimensions).allowed)
|
|
|
|
clock.value += 61
|
|
self.assertTrue(throttle.check(dimensions).allowed)
|
|
|
|
def test_subject_is_hashed_and_success_can_reset_it(self) -> None:
|
|
store = _RecordingStore()
|
|
throttle = FixedWindowThrottle(store, window_seconds=60)
|
|
dimensions = (ThrottleDimension("poll-participation-password", "tenant:token-secret", 2),)
|
|
|
|
throttle.record(dimensions)
|
|
|
|
self.assertEqual(len(store.values), 1)
|
|
stored_key = next(iter(store.values))
|
|
self.assertNotIn("tenant:token-secret", stored_key)
|
|
self.assertIn(":poll-participation-password:", stored_key)
|
|
throttle.reset(dimensions)
|
|
self.assertEqual(store.values, {})
|
|
|
|
def test_resilient_store_uses_bounded_local_fallback(self) -> None:
|
|
clock = _Clock()
|
|
fallback = InMemoryFixedWindowStore(max_entries=2, clock=clock)
|
|
store = ResilientFixedWindowStore(
|
|
_FailingStore(),
|
|
fallback,
|
|
retry_seconds=30,
|
|
clock=clock,
|
|
)
|
|
|
|
first = store.increment("first", window_seconds=60)
|
|
second = store.increment("first", window_seconds=60)
|
|
store.increment("second", window_seconds=60)
|
|
store.increment("third", window_seconds=60)
|
|
|
|
self.assertEqual(first.count, 1)
|
|
self.assertEqual(second.count, 2)
|
|
self.assertEqual(fallback.read("first").count, 0)
|
|
self.assertEqual(fallback.read("third").count, 1)
|
|
|
|
def test_builder_operates_without_redis(self) -> None:
|
|
throttle = build_fixed_window_throttle(
|
|
redis_url=None,
|
|
window_seconds=30,
|
|
max_local_entries=10,
|
|
)
|
|
dimensions = (ThrottleDimension("test", "subject", 1),)
|
|
|
|
self.assertFalse(throttle.record(dimensions).allowed)
|
|
self.assertFalse(throttle.check(dimensions).allowed)
|
|
|
|
def test_rejects_invalid_dimensions(self) -> None:
|
|
throttle = FixedWindowThrottle(_RecordingStore(), window_seconds=60)
|
|
|
|
with self.assertRaisesRegex(ValueError, "At least one"):
|
|
throttle.check(())
|
|
with self.assertRaisesRegex(ValueError, "positive"):
|
|
throttle.check((ThrottleDimension("test", "subject", 0),))
|
|
with self.assertRaisesRegex(ValueError, "namespaces"):
|
|
throttle.check((ThrottleDimension("Invalid namespace", "subject", 1),))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|