Files
ObsiGate/tests/test_store_locks.py

124 lines
4.4 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Tests de non-régression — verrous des stores JSON (ROADMAP #85 T10a).
Sans verrou, les read-modify-write concurrents (créations de partages,
révocations de jetons) perdent des mises à jour : deux threads lisent le
même état, chacun écrit le sien, la première écriture est écrasée. Ces
tests martèlent les stores depuis plusieurs threads et exigent un compte
exact à la fin (aucune mise à jour perdue).
"""
from __future__ import annotations
import threading
N_THREADS = 8
N_OPS = 25
def test_concurrent_share_creations_lose_nothing(tmp_path, monkeypatch):
"""N threads × N partages → le store final contient tout."""
from backend import share as share_mod
monkeypatch.setattr(share_mod, "SHARES_FILE", tmp_path / "shares.json")
errors: list[BaseException] = []
def worker(n: int):
try:
for i in range(N_OPS):
share_mod.create_share("V", f"doc-{n}-{i}.md", "tester")
except BaseException as e: # pragma: no cover - diagnostic
errors.append(e)
threads = [threading.Thread(target=worker, args=(n,)) for n in range(N_THREADS)]
for t in threads:
t.start()
for t in threads:
t.join()
assert not errors
assert len(share_mod.list_shares()) == N_THREADS * N_OPS
def test_concurrent_share_access_and_revoke(tmp_path, monkeypatch):
"""Accès + révocations concurrents : compteurs et suppressions cohérents."""
from backend import share as share_mod
monkeypatch.setattr(share_mod, "SHARES_FILE", tmp_path / "shares.json")
tokens = [share_mod.create_share("V", f"doc-{i}.md", "tester")["token"] for i in range(20)]
def worker(n: int):
share_mod.record_access(tokens[2 * n])
share_mod.record_access(tokens[2 * n + 1])
share_mod.revoke_share(tokens[2 * n])
threads = [threading.Thread(target=worker, args=(n,)) for n in range(10)]
for t in threads:
t.start()
for t in threads:
t.join()
# Les 10 tokens pairs ont été révoqués ; les impairs subsistent avec
# leurs compteurs d'accès (aucune écriture perdue).
remaining = {s["token"] for s in share_mod.list_shares()}
assert len(remaining) == 10
assert all(share_mod.get_share_by_token(t)["access_count"] >= 1 for t in remaining)
def test_concurrent_token_revokes_lose_nothing(tmp_path, monkeypatch):
"""N threads × N révocations → toutes les JTIs sont persistées."""
import json
from backend.auth import jwt_handler
revoked_file = tmp_path / "revoked_tokens.json"
monkeypatch.setattr(jwt_handler, "REVOKED_TOKENS_FILE", revoked_file)
monkeypatch.setattr(jwt_handler, "_revoked_map", {})
monkeypatch.setattr(jwt_handler, "_revoked_loaded", True)
errors: list[BaseException] = []
def worker(n: int):
try:
for i in range(N_OPS):
jwt_handler.revoke_token(f"jti-{n}-{i}")
except BaseException as e: # pragma: no cover - diagnostic
errors.append(e)
threads = [threading.Thread(target=worker, args=(n,)) for n in range(N_THREADS)]
for t in threads:
t.start()
for t in threads:
t.join()
assert not errors
assert len(jwt_handler._revoked_map) == N_THREADS * N_OPS
# Le fichier reflète l'état mémoire (aucune écriture perdue).
on_disk = json.loads(revoked_file.read_text(encoding="utf-8"))
assert len(on_disk) == N_THREADS * N_OPS
assert jwt_handler.is_token_revoked("jti-0-0")
assert not jwt_handler.is_token_revoked("jti-absent")
def test_concurrent_webhook_and_tool_key_writes(tmp_path, monkeypatch):
"""Créations de webhooks + clés d'outils concurrentes : rien de perdu."""
from backend import webhooks as wh_mod
from backend.tools import secrets as sec_mod
monkeypatch.setattr(wh_mod, "WEBHOOKS_FILE", tmp_path / "webhooks.json")
monkeypatch.setattr(wh_mod, "WEBHOOK_SECRETS_FILE", tmp_path / "webhook_secrets.json")
monkeypatch.setattr(sec_mod, "_keys_file", lambda: tmp_path / "api_keys.json")
def worker(n: int):
for i in range(N_OPS):
wh_mod.create_webhook(f"hook-{n}-{i}", "https://example.com/hook", ["file_created"])
sec_mod.set_tool_key("OBSIGATE_GITHUB_TOKEN", f"tok-{n}-{i}")
threads = [threading.Thread(target=worker, args=(n,)) for n in range(N_THREADS)]
for t in threads:
t.start()
for t in threads:
t.join()
assert len(wh_mod.get_webhooks()) == N_THREADS * N_OPS