124 lines
4.4 KiB
Python
124 lines
4.4 KiB
Python
"""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
|