Files
sales-trainer/backend/tests/test_rate_limit_processes.py

181 lines
5.7 KiB
Python

"""Cross-process rate-limit regression tests."""
from __future__ import annotations
import multiprocessing
import os
from pathlib import Path
from app.storage.store import JsonStore
def _hold_record_lock(data_dir: str, key: str, ready, release) -> None:
from app.storage.store import JsonStore
store = JsonStore(Path(data_dir))
with store.record_lock(key):
ready.set()
release.wait(timeout=10)
def _mutate_record(data_dir: str, key: str, operation: str, done, results) -> None:
from app.storage.store import JsonStore
store = JsonStore(Path(data_dir))
try:
if operation == "create":
store.create({"id": key, "value": "created"}, key=key)
elif operation == "replace":
store.replace(key, {"id": key, "value": "replaced"})
else:
store.delete(key)
except BaseException as exc: # pragma: no cover - child diagnostic
results.put((operation, "error", type(exc).__name__))
else:
results.put((operation, "ok"))
finally:
done.set()
def _read_record_while_paused(data_dir: str, key: str, ready, release, results) -> None:
import app.storage.store as store_module
original_read_json = store_module._read_json
def paused_read(path):
ready.set()
release.wait(timeout=10)
return original_read_json(path)
store_module._read_json = paused_read
try:
value = JsonStore(Path(data_dir)).get(key)
except BaseException as exc: # pragma: no cover - child diagnostic
results.put(("read", "error", type(exc).__name__))
else:
results.put(("read", "ok", value["value"]))
def _check_worker(data_dir: str, barrier, results) -> None:
os.environ["DATA_DIR"] = data_dir
from app.config import Config
Config.DATA_DIR = Path(data_dir)
from app.services.rate_limit import check
barrier.wait()
results.put(check("login:user", "same-user", limit=5, window=300))
def test_rate_limit_is_atomic_across_worker_processes(tmp_path, monkeypatch):
monkeypatch.setenv("DATA_DIR", str(tmp_path))
context = multiprocessing.get_context("spawn")
worker_count = 6
barrier = context.Barrier(worker_count)
results = context.Queue()
workers = [
context.Process(
target=_check_worker,
args=(str(tmp_path), barrier, results),
)
for _ in range(worker_count)
]
for worker in workers:
worker.start()
for worker in workers:
worker.join(timeout=30)
assert not worker.is_alive()
assert worker.exitcode == 0
observed = sorted(results.get(timeout=5) for _ in workers)
assert observed == [False, True, True, True, True, True]
def test_rate_limit_fails_closed_on_malformed_persisted_state(tmp_path, monkeypatch):
monkeypatch.setattr("app.config.Config.DATA_DIR", tmp_path)
from app.services import rate_limit
store = rate_limit._record_store()
record_id = rate_limit._record_id(rate_limit._key("login:user", "corrupt-user"))
store.create({"id": record_id, "timestamps": "not-a-list"}, key=record_id)
assert rate_limit.check("login:user", "corrupt-user", limit=5, window=300) is False
store.update(record_id, timestamps=["not-a-number"])
assert rate_limit.check("login:user", "corrupt-user", limit=5, window=300) is False
store._path(record_id).write_text("{", encoding="utf-8")
assert rate_limit.check("login:user", "corrupt-user", limit=5, window=300) is False
store._path(record_id).write_text("[]", encoding="utf-8")
assert rate_limit.check("login:user", "corrupt-user", limit=5, window=300) is False
def test_jsonstore_mutations_honor_cross_process_record_lock(tmp_path):
context = multiprocessing.get_context("spawn")
for operation in ("create", "replace", "delete"):
root = tmp_path / operation
store = JsonStore(root)
key = f"item-{operation}"
if operation != "create":
store.create({"id": key, "value": "initial"}, key=key)
ready = context.Event()
release = context.Event()
done = context.Event()
results = context.Queue()
holder = context.Process(
target=_hold_record_lock,
args=(str(root), key, ready, release),
)
mutator = context.Process(
target=_mutate_record,
args=(str(root), key, operation, done, results),
)
holder.start()
assert ready.wait(timeout=5)
mutator.start()
assert not done.wait(timeout=0.5)
release.set()
holder.join(timeout=5)
mutator.join(timeout=5)
assert holder.exitcode == 0
assert mutator.exitcode == 0
assert results.get(timeout=5) == (operation, "ok")
def test_jsonstore_reader_holds_cross_process_lock_against_delete(tmp_path):
context = multiprocessing.get_context("spawn")
root = tmp_path / "reader-delete"
store = JsonStore(root)
key = "item-reader-delete"
store.create({"id": key, "value": "initial"}, key=key)
ready = context.Event()
release = context.Event()
done = context.Event()
results = context.Queue()
reader = context.Process(
target=_read_record_while_paused,
args=(str(root), key, ready, release, results),
)
deleter = context.Process(
target=_mutate_record,
args=(str(root), key, "delete", done, results),
)
reader.start()
assert ready.wait(timeout=5)
deleter.start()
assert not done.wait(timeout=0.5)
release.set()
reader.join(timeout=5)
deleter.join(timeout=5)
assert reader.exitcode == 0
assert deleter.exitcode == 0
assert sorted(results.get(timeout=5) for _ in range(2)) == [
("delete", "ok"),
("read", "ok", "initial"),
]