diff --git a/backend/app/api/auth.py b/backend/app/api/auth.py index 291350c..f6e1941 100644 --- a/backend/app/api/auth.py +++ b/backend/app/api/auth.py @@ -56,14 +56,13 @@ _vanished: set[str] = set() AUTH_GENERATION_SYNC_SECONDS = 1.0 AUTH_GENERATION_MAX_AGE_SECONDS = 2.0 _generations_synced_at: float | None = None -# One lock per login this node has not synced yet, so a concurrent HTTP -# request and WS handshake for the same just-migrated login collapse into a -# single SELECT instead of one each (lct-42). +# Один lock на ещё неизвестный узлу логин: параллельные HTTP и WS входы +# разделяют один запрос к БД до очередной синхронизации (lct-42). _lookup_locks: dict[str, asyncio.Lock] = {} class _AuthStateUnavailable(Exception): - """The one-shot lookup for a login unknown to this node could not reach PostgreSQL.""" + """Разовая проверка неизвестного узлу логина не смогла обратиться к БД.""" async def _close_revoked(ws: WebSocket) -> None: @@ -116,15 +115,18 @@ async def load_generations() -> None: async def sync_generations() -> None: - """Refresh shared account epochs and close sockets revoked on peer nodes.""" + """Сверить версии полномочий и закрыть отозванные на другом узле сокеты.""" + known_before_query = set(_generations) async with get_sessionmaker()() as db: rows = (await db.execute(select(User.login, User.auth_version))).all() current = {login: version for login, version in rows} for login, version in current.items(): previous = _generations.get(login) + # Снимок БД мог устареть за время запроса; меньшая версия не должна + # отменять локальный отзыв, который уже поднял поколение. if previous is None: _generations[login] = version - elif previous != version: + elif version > previous: invalidate_login(login, version) # Account deletion is not exposed by the application. Still close active # sockets if an operator removes one directly from the shared directory DB. @@ -135,7 +137,9 @@ async def sync_generations() -> None: # lookup fall back to 0 and re-accept cookies issued before the revocation. synthetic = {"dev"} if get_settings().dev_auth_bypass else set() _vanished.intersection_update(_generations.keys() - current.keys()) - for login in _generations.keys() - current.keys() - synthetic - _vanished: + # Разовая проверка могла найти учётку после снимка этого запроса. + # Её отсутствие в старом снимке не означает отзыв. + for login in known_before_query - current.keys() - synthetic - _vanished: invalidate_login(login) _vanished.add(login) global _generations_synced_at @@ -143,32 +147,28 @@ async def sync_generations() -> None: async def _resolve_unknown_login(login: str) -> int | None: - """Look up a login absent from this node's cache without waiting a tick. + """Проверить неизвестный узлу логин сразу, не дожидаясь опроса БД. - A revoked login stays cached in `_generations` (see `_vanished` above), - so absence from `_generations` unambiguously means "this node has not - synced this login yet" — an unrecognized login is checked directly - rather than treated as revoked. Returns the current epoch, or `None` if - the login is not (or no longer) in `users`. Raises - `_AuthStateUnavailable` if PostgreSQL cannot be reached; the caller - fails closed. + Отозванный логин остаётся в `_generations` (см. `_vanished`), поэтому + отсутствие в кэше означает, что этот узел ещё не видел учётку. + Возвращает версию или None, если строки в `users` нет. При недоступной + БД вызывает `_AuthStateUnavailable`, чтобы вход был закрыт с 503/1013. """ lock = _lookup_locks.setdefault(login, asyncio.Lock()) async with lock: cached = _generations.get(login) if cached is not None: - return cached # a concurrent request already resolved it + return cached # параллельный запрос уже получил версию try: async with get_sessionmaker()() as db: version = await db.scalar( select(User.auth_version).where(User.login == login) ) - except Exception as exc: # noqa: BLE001 — fail closed, not FORBIDDEN + except Exception as exc: # noqa: BLE001 — ошибка БД должна закрыть вход log.error("разовая проверка полномочий не удалась (%s)", type(exc).__name__) raise _AuthStateUnavailable from exc - # A concurrent local revoke or sync tick may have written a newer - # generation while the SELECT above was in flight; the stale read - # must never clobber it back to an older, already-revoked epoch. + # Пока шёл SELECT, локальный отзыв или синхронизация могли записать + # новую версию. Старый ответ не должен вернуть отозванную cookie. cached = _generations.get(login) if cached is not None: return cached diff --git a/backend/tests/test_auth_hardening.py b/backend/tests/test_auth_hardening.py index d39c759..e484462 100644 --- a/backend/tests/test_auth_hardening.py +++ b/backend/tests/test_auth_hardening.py @@ -135,6 +135,76 @@ def test_generation_sync_preserves_synthetic_dev_account(monkeypatch): asyncio.run(run()) +@pytest.mark.asyncio +async def test_stale_generation_snapshot_does_not_revoke_newly_resolved_login(monkeypatch): + snapshot_read = asyncio.Event() + release_snapshot = asyncio.Event() + + class FakeResult: + def all(self): + return [] + + class FakeDb: + async def execute(self, _query): + snapshot_read.set() + await release_snapshot.wait() + return FakeResult() + + async def scalar(self, _query): + return 0 + + class FakeSession: + async def __aenter__(self): + return FakeDb() + + async def __aexit__(self, *_args): + return None + + auth.prime_generations({}) + monkeypatch.setattr(auth, "get_settings", lambda: SimpleNamespace(dev_auth_bypass=False)) + monkeypatch.setattr(auth, "get_sessionmaker", lambda: lambda: FakeSession()) + sync = asyncio.create_task(auth.sync_generations()) + await snapshot_read.wait() + assert await auth._resolve_unknown_login("just-created") == 0 + release_snapshot.set() + await sync + assert auth._generations["just-created"] == 0 + assert "just-created" not in auth._vanished + + +@pytest.mark.asyncio +async def test_stale_generation_snapshot_does_not_restore_revoked_cookie(monkeypatch): + snapshot_read = asyncio.Event() + release_snapshot = asyncio.Event() + + class FakeResult: + def all(self): + return [("revoked", 0)] + + class FakeDb: + async def execute(self, _query): + snapshot_read.set() + await release_snapshot.wait() + return FakeResult() + + class FakeSession: + async def __aenter__(self): + return FakeDb() + + async def __aexit__(self, *_args): + return None + + auth.prime_generations({"revoked": 0}) + monkeypatch.setattr(auth, "get_settings", lambda: SimpleNamespace(dev_auth_bypass=False)) + monkeypatch.setattr(auth, "get_sessionmaker", lambda: lambda: FakeSession()) + sync = asyncio.create_task(auth.sync_generations()) + await snapshot_read.wait() + auth.invalidate_login("revoked", 1) + release_snapshot.set() + await sync + assert auth._generations["revoked"] == 1 + + def test_cross_origin_browser_websocket_is_rejected_before_handshake(client): assert client.post("/api/auth/dev-token").status_code == 200 with pytest.raises(WebSocketDisconnect) as exc: diff --git a/scripts/test_cluster_failover.py b/scripts/test_cluster_failover.py index 0ececc8..a907266 100644 --- a/scripts/test_cluster_failover.py +++ b/scripts/test_cluster_failover.py @@ -72,7 +72,7 @@ def login(base: str, login_name: str, password: str) -> str: return cookie -async def send_control(ws_base: str, session_id: UUID, cookie: str, payload: dict) -> None: +async def send_control(ws_base: str, session_id: UUID, cookie: str, payload: dict) -> float: async with websockets.connect( f"{ws_base}/ws/control/{session_id}", additional_headers={"Cookie": cookie}, @@ -81,7 +81,9 @@ async def send_control(ws_base: str, session_id: UUID, cookie: str, payload: dic ping_interval=None, ) as socket: await socket.send(json.dumps(payload, ensure_ascii=False)) + sent_at = time.monotonic() await asyncio.sleep(0.15) + return sent_at async def probe_ws(ws_base: str, path: str, cookie: str) -> dict: @@ -220,7 +222,7 @@ async def main(args: argparse.Namespace) -> int: factory = async_sessionmaker(engine, expire_on_commit=False) observer_engine = create_async_engine(database_url, poolclass=NullPool) run_id = uuid4().hex[:12] - session_id = uuid4() + session_id = args.session_id or uuid4() trainee = Trainee(name=f"Failover smoke {run_id}") instructor_login = f"fo-i-{run_id}" trainee_login = f"fo-t-{run_id}" @@ -232,6 +234,7 @@ async def main(args: argparse.Namespace) -> int: fault_stopped = False result: dict = {"session_id": str(session_id), "checks": {}} backend = f"http://127.0.0.1:{args.backend_port}" + result["login_node"] = "backend-a" ws_base = args.frontend_url.replace("https://", "wss://").replace("http://", "ws://") instructor_cookie = trainee_cookie = None try: @@ -249,12 +252,18 @@ async def main(args: argparse.Namespace) -> int: await db.commit() instructor_cookie = login(backend, instructor_login, instructor_password) + login_completed_at = time.monotonic() trainee_cookie = login(backend, trainee_login, trainee_password) - await send_control(ws_base, session_id, instructor_cookie, { + control_sent_at = await send_control(ws_base, session_id, instructor_cookie, { "type": "scenario.start", "scenario_id": args.scenario, "trainee": trainee.name, "trainee_id": str(trainee.id), "mode": "training", "exercise": "dds", }) + result["login_to_control_seconds"] = round(control_sent_at - login_completed_at, 4) + if args.expect_initial_owner == "backend-b": + result["checks"]["control_sent_within_first_second"] = ( + result["login_to_control_seconds"] < 1 + ) async def session_row(): async with factory() as db: @@ -289,6 +298,12 @@ async def main(args: argparse.Namespace) -> int: old_epoch = row.backend_fencing_epoch result["initial_owner"] = row.backend_node_id result["initial_epoch"] = old_epoch + if args.expect_initial_owner and row.backend_node_id != args.expect_initial_owner: + raise RuntimeError( + f"session routed to {row.backend_node_id}, expected {args.expect_initial_owner}" + ) + if args.expect_initial_owner == "backend-b": + result["checks"]["login_a_initial_control_b"] = True # Commit real DDS work before the fault. These actions must survive the # checkpoint handoff and remain part of the final scored report. @@ -642,4 +657,6 @@ if __name__ == "__main__": parser.add_argument("--reconnect-attempts", type=int, default=5) parser.add_argument("--rest-timeout", type=float, default=60) parser.add_argument("--failure-mode", choices=("kill", "partition"), default="kill") + parser.add_argument("--expect-initial-owner", choices=("backend-a", "backend-b")) + parser.add_argument("--session-id", type=UUID) raise SystemExit(asyncio.run(main(parser.parse_args())))