From 0f647e8bbbc9a2f29f98242bb45ba1e4725b89cb Mon Sep 17 00:00:00 2001 From: teknium1 <127238744+teknium1@users.noreply.github.com> Date: Wed, 23 Sep 2026 03:05:25 -0700 Subject: [PATCH] test: C14 multi-client session ownership model check over the real backend Class C14 (multi-client session ownership / routing / switching): prompts persisted to a stale session after a switch, events rendered twice or leaking into another chat, zombie leases ("already has a live owner"), a ws drop or client absence killing the in-flight turn, a live session hard-deleted. tests/e2e/core/terminal/_gateway_client.py spawns ONE real `python -m hermes_cli.main serve --port 0` per module (Desktop's argv) in an isolated HOME with only the recording fake provider, and drives it with Desktop-shaped WebSocket JSON-RPC clients (gateway.ready, client.capabilities, server->client request answers, abrupt TCP drop, reader pause for absence). test_multiclient_session_model.py runs a seeded fuzzer (seeds 11/23/37) over 3 clients: create, switch(session.resume), fast and slow prompts, interrupt mid-stream, delete of live (must be refused) and closed sessions, ws drop + reconnect + resume mid-turn, private-chat owner death (idle and running) + takeover, shared-chat owner death, client absence. A reference model checks after every step and at the end: prompts persisted exactly once, in order, to the session selected at send time; per-connection event seq strictly increasing (no duplicate delivery); no events for sessions a connection never attached to; canaries only in their own session; one live runtime per chat; orphaned runtimes released in bounded time and resumable/promptable by another client; drop/absence never cut a turn short (full reply persisted, message.complete status complete). Red-proof (each fails, restored after): duplicate event write in write_json; semantic revert of 67de93862c3 (#98028/#100325, orphan reaper ignores turn activity); semantic revert of de25545dce1 (rebind instead of fan-out); lease never released; session.delete of a live session not refused. (cherry picked from commit fe078a891231d3ba459da4ca46bf2a06979f4444) --- tests/e2e/core/terminal/__init__.py | 0 tests/e2e/core/terminal/_gateway_client.py | 308 +++++++++++ .../test_multiclient_session_model.py | 521 ++++++++++++++++++ 3 files changed, 829 insertions(+) create mode 100644 tests/e2e/core/terminal/__init__.py create mode 100644 tests/e2e/core/terminal/_gateway_client.py create mode 100644 tests/e2e/core/terminal/test_multiclient_session_model.py diff --git a/tests/e2e/core/terminal/__init__.py b/tests/e2e/core/terminal/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/e2e/core/terminal/_gateway_client.py b/tests/e2e/core/terminal/_gateway_client.py new file mode 100644 index 0000000000..4467edb182 --- /dev/null +++ b/tests/e2e/core/terminal/_gateway_client.py @@ -0,0 +1,308 @@ +"""Real ``hermes serve`` backend + WebSocket JSON-RPC clients for the terminal lane E2E suites. + +One backend per test module: a real ``python -m hermes_cli.main serve --port 0`` subprocess (the +exact argv Desktop spawns, see ``apps/desktop/electron/backend-command.ts``) with an isolated +HOME/HERMES_HOME, no inherited provider credentials, and only the recording fake LLM provider +configured. Clients speak the Desktop wire protocol over ``/api/ws`` (newline JSON-RPC, +``gateway.ready`` on accept, ``client.capabilities`` to answer server→client requests), so every +event, lease, reaper and approval path is the production one. +""" + +from __future__ import annotations + +import json +import os +import re +import secrets +import signal +import socket +import sqlite3 +import subprocess +import sys +import threading +import time +from pathlib import Path +from typing import Any, Callable, Iterable + +from websockets.sync.client import connect as ws_connect + +from tests.fakes.fake_llm_provider import FakeLLMServer, write_hermes_home + +REPO_ROOT = Path(__file__).resolve().parents[4] +READY_RE = re.compile(r"HERMES_BACKEND_READY port=(\d+)") + + +class RpcError(AssertionError): + def __init__(self, method: str, error: dict) -> None: + self.method = method + self.code = error.get("code") + self.message = str(error.get("message") or "") + self.data = error.get("data") + super().__init__(f"{method} -> {self.code}: {self.message}") + + +def poll_until(fn: Callable[[], Any], *, timeout: float, interval: float = 0.05, what: str = "condition") -> Any: + """Return the first truthy ``fn()`` before ``timeout``; AssertionError naming ``what`` otherwise.""" + deadline = time.monotonic() + timeout + last_exc: Exception | None = None + while True: + try: + value = fn() + if value: + return value + except AssertionError: + raise + except Exception as exc: # noqa: BLE001 - surfaced in the timeout message + last_exc = exc + if time.monotonic() >= deadline: + raise AssertionError(f"timed out after {timeout:.0f}s waiting for {what}" + + (f" (last error: {last_exc!r})" if last_exc else "")) + time.sleep(interval) + + +def _operator_home() -> str: + """The account's real home (passwd), immune to a test-time HOME override.""" + try: + import pwd + return pwd.getpwuid(os.getuid()).pw_dir + except (ImportError, KeyError): # pragma: no cover - non-POSIX + return os.path.expanduser("~") + + +def _minimal_env(home: Path, hermes_home: Path, token: str, extra: dict[str, str] | None) -> dict[str, str]: + """A clean child env: nothing Hermes-, provider- or credential-shaped leaks in from the test runner.""" + env = {k: os.environ[k] for k in ("PATH", "LANG", "LC_ALL", "TERM", "SYSTEMROOT") if k in os.environ} + tmpdir = home.parent / "tmp" + tmpdir.mkdir(parents=True, exist_ok=True) + env.update( + HOME=str(home), HERMES_HOME=str(hermes_home), PYTHONPATH=str(REPO_ROOT), + TMPDIR=str(tmpdir), HERMES_DASHBOARD_SESSION_TOKEN=token, PYTHONUNBUFFERED="1", + NO_COLOR="1", + # The child's HOME *is* the sandbox, so its "real" root (expanduser('~')/.hermes) is the tmp + # home and the pytest-ancestry live-DB guard would refuse every open. The guard exists to keep + # tests off the operator's state.db; Backend.start() asserts the sandbox is outside it. + HERMES_STATE_DB_GUARD_BYPASS="1", + ) + env.update(extra or {}) + return env + + +class Backend: + """One real ``hermes serve`` process wired to a :class:`FakeLLMServer`.""" + + def __init__(self, root: Path, responder, *, extra_config: str = "", env: dict[str, str] | None = None, + aux=None) -> None: + self.root = root + self.home = root / "home" + self.hermes_home = self.home / ".hermes" + self.llm = FakeLLMServer(responder, aux=aux) + self.extra_config = extra_config + self.extra_env = env or {} + self.token = secrets.token_hex(16) + self.proc: subprocess.Popen | None = None + self.port: int | None = None + self.clients: list[WSClient] = [] + + def start(self, *, timeout: float = 90.0) -> "Backend": + operator_root = (Path(_operator_home()) / ".hermes").resolve() + assert operator_root not in (self.hermes_home.resolve(), *self.hermes_home.resolve().parents), ( + f"sandbox {self.hermes_home} sits inside the operator's Hermes home {operator_root}") + self.llm.start() + write_hermes_home(self.hermes_home, self.llm.base_url, extra_config=self.extra_config) + self._stdout = open(self.root / "serve.stdout.log", "wb") + self._stderr = open(self.root / "serve.stderr.log", "wb") + # Own process group: stop() signals exactly the tree this fixture spawned, never a pattern match. + self.proc = subprocess.Popen( + [sys.executable, "-m", "hermes_cli.main", "serve", "--host", "127.0.0.1", "--port", "0"], + cwd=str(self.root), env=_minimal_env(self.home, self.hermes_home, self.token, self.extra_env), + stdin=subprocess.DEVNULL, stdout=self._stdout, stderr=self._stderr, start_new_session=True) + + proc = self.proc + + def ready() -> int | None: + if proc.poll() is not None: + raise AssertionError(f"hermes serve exited {proc.returncode} before ready:\n{self.logs()}") + match = READY_RE.search((self.root / "serve.stdout.log").read_text(errors="replace")) + return int(match.group(1)) if match else None + + try: + self.port = poll_until(ready, timeout=timeout, interval=0.1, what="HERMES_BACKEND_READY") + except BaseException: + self.stop() + raise + return self + + @property + def ws_url(self) -> str: + return f"ws://127.0.0.1:{self.port}/api/ws?token={self.token}" + + def connect(self, name: str) -> "WSClient": + client = WSClient(self.ws_url, name) + self.clients.append(client) + return client + + def logs(self, tail: int = 60) -> str: + parts = [] + for path in (self.root / "serve.stdout.log", self.root / "serve.stderr.log", + self.hermes_home / "logs" / "errors.log"): + name = path.name + if path.exists(): + lines = path.read_text(errors="replace").splitlines()[-tail:] + parts.append(f"--- {name} ---\n" + "\n".join(lines)) + return "\n".join(parts) + + def stop(self) -> None: + for client in self.clients: + client.close() + proc, self.proc = self.proc, None + if proc is not None and proc.poll() is None: + try: + os.killpg(proc.pid, signal.SIGTERM) + except ProcessLookupError: + pass + try: + proc.wait(timeout=15) + except subprocess.TimeoutExpired: + try: + os.killpg(proc.pid, signal.SIGKILL) + except ProcessLookupError: + pass + proc.wait(timeout=10) + self.llm.stop() + for fh in (getattr(self, "_stdout", None), getattr(self, "_stderr", None)): + if fh is not None: + fh.close() + + # state.db, read-only (never a second writer next to the backend's) + def db_rows(self, sql: str, args: Iterable[Any] = ()) -> list[tuple]: + path = self.hermes_home / "state.db" + if not path.exists(): + return [] + conn = sqlite3.connect(f"file:{path}?mode=ro", uri=True, timeout=30) + try: + conn.execute("PRAGMA busy_timeout=30000") + return conn.execute(sql, tuple(args)).fetchall() + finally: + conn.close() + + +class WSClient: + """One Desktop-shaped WebSocket connection. A reader thread records every inbound frame.""" + + def __init__(self, url: str, name: str, *, answers_requests: bool = True) -> None: + self.name = name + self.frames: list[dict] = [] + self._cond = threading.Condition() + self._reading = threading.Event() + self._reading.set() + self._next_id = 0 + self.closed = False + self.dropped = False + self._ws = ws_connect(url, max_size=None, ping_interval=None, open_timeout=30, close_timeout=2) + self._thread = threading.Thread(target=self._read_loop, name=f"ws-{name}", daemon=True) + self._thread.start() + self.wait_for(lambda f: etype(f) == "gateway.ready", timeout=60, what="gateway.ready") + if answers_requests: + self.call("client.capabilities", server_requests=True) + + def _read_loop(self) -> None: + while True: + self._reading.wait() + try: + raw = self._ws.recv() + except Exception: # noqa: BLE001 - closed/dropped socket ends the reader + break + for line in str(raw).splitlines(): + if line.strip(): + frame = json.loads(line) + frame["_t"] = time.monotonic() + with self._cond: + self.frames.append(frame) + self._cond.notify_all() + with self._cond: + self.closed = True + self._cond.notify_all() + + # -- absence: stop consuming; the websockets receive queue fills and TCP backpressures the backend + def pause(self) -> None: + self._reading.clear() + + def resume(self) -> None: + self._reading.set() + + def snapshot(self) -> list[dict]: + with self._cond: + return list(self.frames) + + def wait_for(self, pred: Callable[[dict], bool], *, timeout: float, start: int = 0, what: str = "frame") -> dict: + deadline = time.monotonic() + timeout + with self._cond: + idx = start + while True: + while idx < len(self.frames): + frame = self.frames[idx] + idx += 1 + if pred(frame): + return frame + if self.closed: + raise AssertionError(f"{self.name}: socket closed while waiting for {what}") + remaining = deadline - time.monotonic() + if remaining <= 0: + raise AssertionError(f"{self.name}: timed out after {timeout:.0f}s waiting for {what}") + self._cond.wait(min(remaining, 0.5)) + + def send(self, frame: dict) -> None: + self._ws.send(json.dumps(frame)) + + def request(self, method: str, params: dict, *, timeout: float = 60.0) -> dict: + """Send one RPC; return the raw response frame (``result`` or ``error``).""" + self._next_id += 1 + rid = self._next_id + start = len(self.snapshot()) + self.send({"jsonrpc": "2.0", "id": rid, "method": method, "params": params}) + return self.wait_for(lambda f: f.get("id") == rid and "method" not in f, timeout=timeout, + start=start, what=f"{method} reply") + + def call(self, method: str, *, timeout: float = 60.0, **params: Any) -> dict: + frame = self.request(method, params, timeout=timeout) + if "error" in frame: + raise RpcError(method, frame["error"]) + return frame.get("result") or {} + + def respond(self, request_id: str, result: dict) -> None: + """Answer a server→client request (clarify/approval/...) with a response frame.""" + self.send({"jsonrpc": "2.0", "id": request_id, "result": result}) + + def events(self, sid: str | None = None, etype: str | None = None) -> list[dict]: + return [f for f in self.snapshot() if f.get("method") == "event" + and (sid is None or f["params"].get("session_id") == sid) + and (etype is None or f["params"].get("type") == etype)] + + def server_requests(self, method: str | None = None) -> list[dict]: + return [f for f in self.snapshot() if isinstance(f.get("id"), str) and f.get("method") not in (None, "event") + and (method is None or f.get("method") == method)] + + def drop(self) -> None: + """Abrupt network loss: no close frame, the backend learns from the dead TCP stream.""" + self.dropped = True + self._reading.set() + sock = getattr(self._ws, "socket", None) + if sock is not None: + try: + sock.shutdown(socket.SHUT_RDWR) + except OSError: + pass + sock.close() + self._thread.join(timeout=10) + + def close(self) -> None: + self._reading.set() + try: + self._ws.close() + except Exception: # noqa: BLE001 - already dropped + pass + + +def etype(frame: dict) -> str: + """The event ``type`` of a notification frame ('' for replies and server→client requests).""" + return str((frame.get("params") or {}).get("type") or "") if frame.get("method") == "event" else "" diff --git a/tests/e2e/core/terminal/test_multiclient_session_model.py b/tests/e2e/core/terminal/test_multiclient_session_model.py new file mode 100644 index 0000000000..a5170fbffb --- /dev/null +++ b/tests/e2e/core/terminal/test_multiclient_session_model.py @@ -0,0 +1,521 @@ +"""C14 — multi-client session ownership, routing and switching over the REAL backend. + +Class: a message persisted to a stale session after a switch, session RPCs collapsing onto the +wrong connection, events rendered twice or leaking into another chat, a zombie lease locking a +chat ("already has a live owner"), a ws drop orphaning the in-flight turn, client absence killing +the turn, an active session hard-deleted. + +Harness: one real ``hermes serve`` (the Desktop's backend argv) per module, 3 Desktop-shaped +WebSocket clients, and a seeded fuzzer that interleaves create / switch(resume) / prompt (fast and +slow streams) / interrupt mid-stream / delete (inactive and active) / close / abrupt ws drop + +reconnect + resume / owner death + takeover / client absence. The fake provider echoes a per-prompt +canary so every assistant reply is attributable. A reference model tracks what each client has +selected and sent; invariants are checked after every step and at the end: + +1. every accepted prompt is persisted to the session its sender had selected, exactly once, in order; +2. each event reaches a subscribed connection at most once (per-session ``seq`` strictly increasing + per connection), never reaches a connection that did not attach to that session, and every + assistant canary appears only in its own session's events; +3. after the owning connection dies, the session is released in bounded time and another client can + resume and prompt it (no zombie lease), and no session key ever has two live runtimes; +4. deleting a live session is refused and destroys nothing; +5. a client that reconnects after a mid-turn drop gets the turn's completion (live event or history), + and neither a drop nor an absence interrupts the turn. +""" + +from __future__ import annotations + +import random +import re +import sys +import time +from collections import Counter, defaultdict +from dataclasses import dataclass, field + +import pytest + +from tests.e2e.core.terminal._gateway_client import Backend, RpcError, WSClient, etype, poll_until +from tests.fakes.fake_llm_provider import Text + +# Not marked ``integration`` (that marker means external services and is deselected by addopts): the backend is +# loopback-only. ``tests/e2e`` is outside default discovery; run with ``scripts/run_tests.sh --include-integration``. +pytestmark = pytest.mark.skipif(sys.platform == "win32", reason="POSIX process-group lifecycle for the spawned backend") + +CANARY_RE = re.compile(r"cnry-s\d+-\d{3}") +FILLER = " ".join(["stream"] * 45) +ABSENCE_S = 3.0 # far below the 30 s WS send deadline and the reaper's activity-stale threshold +GRACE_S = 2 # HERMES_TUI_WS_ORPHAN_REAP_GRACE_S: short so owner-death paths finish inside the test +STEP_TIMEOUT = 90.0 + + +def _text_of(message: dict) -> str: + content = message.get("content") + if isinstance(content, list): + return " ".join(str(part.get("text") or "") for part in content if isinstance(part, dict)) + return str(content or "") + + +def _responder(record: dict): + """Echo the newest canary: ``ACK ... DONE``; prompts containing ``slow`` stream for ~3 s, so a + reply that ends in ``DONE`` proves the turn was not cut short.""" + messages = record["body"].get("messages") or [] + last_user = next((m for m in reversed(messages) if m.get("role") == "user"), {}) + text = _text_of(last_user) + canaries = CANARY_RE.findall(text) + canary = canaries[-1] if canaries else "none" + if "slow" in text: + return Text(f"ACK {canary} {FILLER} DONE", chunk_chars=6, delay_per_chunk=0.07) + return Text(f"ACK {canary} DONE") + + +@pytest.fixture(scope="module") +def backend(tmp_path_factory): + root = tmp_path_factory.mktemp("c14") + be = Backend(root, _responder, extra_config="display:\n busy_input_mode: queue\n", + env={"HERMES_TUI_WS_ORPHAN_REAP_GRACE_S": str(GRACE_S)}) + be.start() + try: + yield be + finally: + be.stop() + + +@dataclass +class Turn: + canary: str + key: str + sid: str + client: str + slow: bool + attached: set[int] # id() of connections attached to ``sid`` when the prompt was accepted + interrupted: bool = False + + +@dataclass +class Sess: + sid: str | None + prompts: list[str] = field(default_factory=list) + deleted: bool = False + closed: bool = False + + +@dataclass +class Slot: + name: str + conn: WSClient | None = None + key: str | None = None + sid: str | None = None + + +class Model: + def __init__(self, backend: Backend, seed: int) -> None: + self.be = backend + self.seed = seed + self.rng = random.Random(seed) + self.slots = {n: Slot(n) for n in ("A", "B", "C")} + self.sessions: dict[str, Sess] = {} + self.sid_key: dict[str, str] = {} + self.turns: dict[str, Turn] = {} + self.conns: list[WSClient] = [] + # id(conn) -> sid -> frame index at the moment the attaching RPC was sent + self.attach_idx: dict[int, dict[str, int]] = defaultdict(dict) + self.n = 0 + self.log: list[str] = [] + + # -- plumbing ---------------------------------------------------------------------------- + def connect(self, slot: Slot) -> WSClient: + conn = self.be.connect(f"{slot.name}{len(self.conns)}-s{self.seed}") + self.conns.append(conn) + slot.conn = conn + return conn + + def live_conns(self) -> list[WSClient]: + return [c for c in self.conns if not c.dropped and not c.closed] + + def any_conn(self) -> WSClient: + for slot in self.slots.values(): + if slot.conn is not None and not slot.conn.dropped: + return slot.conn + return self.connect(self.slots["A"]) + + def attached_live(self, sid: str) -> list[WSClient]: + return [c for c in self.live_conns() if sid in self.attach_idx[id(c)]] + + def mark_attach(self, conn: WSClient, sid: str, idx: int) -> None: + self.attach_idx[id(conn)].setdefault(sid, idx) + + def active_rows(self) -> list[dict]: + return self.any_conn().call("session.active_list").get("sessions") or [] + + def wait_idle(self, key: str) -> None: + def idle() -> bool: + rows = [r for r in self.active_rows() if r.get("session_key") == key] + return all(r.get("status") == "idle" for r in rows) + poll_until(idle, timeout=STEP_TIMEOUT, interval=0.1, what=f"session {key} idle") + + def resume(self, slot: Slot, key: str) -> str: + conn = slot.conn + assert conn is not None + old = self.sessions[key].sid + holders = [c for c in self.attached_live(old)] if old else [] + + def attempt(): + idx = len(conn.snapshot()) + frame = conn.request("session.resume", {"session_id": key, "cols": 100}) + if "error" in frame: + if frame["error"].get("code") == 4009: # client-gone interrupt settling: documented retry + return None + raise RpcError("session.resume", frame["error"]) + return idx, frame["result"] + + idx, result = poll_until(attempt, timeout=STEP_TIMEOUT, interval=0.2, what=f"resume {key}") + sid = result["session_id"] + if holders: + assert sid == old, ( + f"resume of {key} minted runtime {sid} while connection(s) " + f"{[c.name for c in holders]} still hold live runtime {old}: two runtimes for one chat") + self.sessions[key].sid = sid + self.sessions[key].closed = False + self.sid_key[sid] = key + self.mark_attach(conn, sid, idx) + slot.key, slot.sid = key, sid + return sid + + def submit(self, slot: Slot, *, slow: bool) -> Turn: + conn, key = slot.conn, slot.key + assert conn is not None and key is not None + self.wait_idle(key) + self.n += 1 + canary = f"cnry-s{self.seed}-{self.n:03d}" + text = f"{'slow ' if slow else ''}please ack {canary}" + refusals: list[str] = [] + + def attempt(): + frame = conn.request("prompt.submit", {"session_id": slot.sid, "text": text}) + if "error" not in frame: + return frame["result"] + err = frame["error"] + code = err.get("code") + if code == 4009: # settling fence + return None + if code in (4001, 4007): # runtime reaped under us: the client re-resumes, as Desktop does + self.resume(slot, key) + return None + if code == 4090: # ownership refusal: retried until the deadline, then a zombie lease + refusals.append(str(err.get("message"))) + return None + raise RpcError("prompt.submit", err) + + try: + result = poll_until(attempt, timeout=45.0, interval=0.25, what=f"prompt.submit to {key}") + except AssertionError as exc: + raise AssertionError(f"{exc}; ownership refusals: {refusals[-3:]}") from None + assert result.get("status") == "streaming", f"prompt to idle {key} was not started: {result}" + sid = slot.sid + assert sid is not None + turn = Turn(canary, key, sid, slot.name, slow, {id(c) for c in self.attached_live(sid)}) + self.turns[canary] = turn + self.sessions[key].prompts.append(canary) + return turn + + def wait_first_delta(self, turn: Turn, conn: WSClient) -> None: + conn.wait_for(lambda f: etype(f) == "message.delta" and f["params"].get("session_id") == turn.sid, + timeout=STEP_TIMEOUT, what=f"first delta of {turn.canary}") + + def db_has_reply(self, turn: Turn) -> bool: + return bool(self.be.db_rows( + "SELECT 1 FROM messages WHERE session_id=? AND role='assistant' AND content LIKE ?", + (turn.key, f"%ACK {turn.canary} %DONE%"))) + + def complete_frames(self, conn: WSClient, turn: Turn) -> list[dict]: + return [f for f in conn.events(turn.sid, "message.complete") + if f"ACK {turn.canary}" in str((f["params"].get("payload") or {}).get("text") or "")] + + # -- ops --------------------------------------------------------------------------------- + def op_create(self, slot: Slot) -> None: + conn = slot.conn or self.connect(slot) + idx = len(conn.snapshot()) + result = conn.call("session.create", cols=100, source="desktop") + sid, key = result["session_id"], result["stored_session_id"] + self.sessions[key] = Sess(sid) + self.sid_key[sid] = key + self.mark_attach(conn, sid, idx) + slot.key, slot.sid = key, sid + + def resumable(self) -> list[str]: + out = [] + for key, sess in self.sessions.items(): + if sess.deleted: + continue + if sess.prompts or (sess.sid and self.attached_live(sess.sid) and not sess.closed): + out.append(key) + return out + + def op_switch(self, slot: Slot) -> None: + if slot.conn is None: + self.connect(slot) + others = [k for k in self.resumable() if k != slot.key] + if not others: + return self.op_create(slot) + self.resume(slot, self.rng.choice(sorted(others))) + + def ensure_selected(self, slot: Slot) -> None: + if slot.conn is None: + self.op_reconnect(slot) + if slot.key is None or self.sessions[slot.key].deleted or self.sessions[slot.key].closed: + self.op_create(slot) + + def op_prompt(self, slot: Slot, slow: bool = False) -> Turn: + self.ensure_selected(slot) + return self.submit(slot, slow=slow) + + def op_interrupt_mid_turn(self, slot: Slot) -> None: + turn = self.op_prompt(slot, slow=True) + assert slot.conn is not None + self.wait_first_delta(turn, slot.conn) + slot.conn.call("session.interrupt", session_id=turn.sid) + turn.interrupted = True + self.wait_idle(turn.key) + + def op_delete_active(self, slot: Slot) -> None: + self.ensure_selected(slot) + if not self.sessions[slot.key].prompts: + self.submit(slot, slow=False) + key = slot.key + assert slot.conn is not None and key is not None + before = self.be.db_rows("SELECT COUNT(*) FROM messages WHERE session_id=?", (key,)) + frame = slot.conn.request("session.delete", {"session_id": key}) + assert "error" in frame, f"session.delete of LIVE session {key} succeeded: {frame}" + assert self.be.db_rows("SELECT COUNT(*) FROM sessions WHERE id=?", (key,)) == [(1,)], ( + f"refused delete still removed live session {key}") + after = self.be.db_rows("SELECT COUNT(*) FROM messages WHERE session_id=?", (key,)) + assert after >= before, f"refused delete of {key} dropped messages: {before} -> {after}" + + def op_close_and_delete(self, slot: Slot) -> None: + # A private chat: create, prompt, close (the runtime ends), then delete the inactive row. + self.ensure_selected(slot) + self.op_create(slot) + turn = self.submit(slot, slow=False) + key, sid = turn.key, turn.sid + assert slot.conn is not None + self.wait_idle(key) + poll_until(lambda: self.db_has_reply(turn), timeout=STEP_TIMEOUT, what=f"reply persisted for {turn.canary}") + slot.conn.call("session.close", session_id=sid) + self.sessions[key].closed = True + self.sessions[key].sid = None + slot.key = slot.sid = None + poll_until(lambda: not any(r.get("session_key") == key for r in self.active_rows()), + timeout=STEP_TIMEOUT, what=f"closed session {key} gone from the live list") + slot.conn.call("session.delete", session_id=key) + self.sessions[key].deleted = True + assert self.be.db_rows("SELECT COUNT(*) FROM sessions WHERE id=?", (key,)) == [(0,)] + assert self.be.db_rows("SELECT COUNT(*) FROM messages WHERE session_id=?", (key,)) == [(0,)] + + def op_drop_mid_turn(self, slot: Slot) -> None: + turn = self.op_prompt(slot, slow=True) + assert slot.conn is not None + self.wait_first_delta(turn, slot.conn) + slot.conn.drop() + conn = self.connect(slot) + self.resume(slot, turn.key) + poll_until(lambda: self.complete_frames(conn, turn) or self.db_has_reply(turn), + timeout=STEP_TIMEOUT, what=f"completion of in-flight {turn.canary} after reconnect") + self.wait_idle(turn.key) + assert self.db_has_reply(turn), f"turn {turn.canary} was cut short by the ws drop (no full reply persisted)" + for frame in self.complete_frames(conn, turn): + assert (frame["params"].get("payload") or {}).get("status", "complete") == "complete", frame + + def orphan(self, slot: Slot, *, running: bool) -> None: + """The owner of a private chat dies (no reconnect); another client must take it over.""" + self.ensure_selected(slot) + self.op_create(slot) + self.submit(slot, slow=False) # a durable row: the chat is worth taking over + turn = self.submit(slot, slow=True) if running else None + key, sid = slot.key, slot.sid + assert slot.conn is not None and key is not None and sid is not None + if turn is not None: + self.wait_first_delta(turn, slot.conn) + slot.conn.drop() + slot.conn = None + assert not self.attached_live(sid) + # Nobody else watches: the runtime must be released in bounded time (grace + turn end). + poll_until(lambda: not any(r.get("session_key") == key for r in self.active_rows()), + timeout=STEP_TIMEOUT, interval=0.2, what=f"orphaned {key} released after its owner died") + if turn is not None: + assert self.db_has_reply(turn), f"owner death cut fresh turn {turn.canary} short (no full reply)" + taker = self.rng.choice([s for s in self.slots.values() if s is not slot]) + if taker.conn is None: + self.connect(taker) + self.resume(taker, key) + self.submit(taker, slow=False) # 45 s of 4090 refusals here = zombie lease + + def op_orphan_idle(self, slot: Slot) -> None: + self.orphan(slot, running=False) + + def op_orphan_running(self, slot: Slot) -> None: + self.orphan(slot, running=True) + + def op_shared_owner_drop(self, slot: Slot) -> None: + """Two windows on one chat; the prompting one dies mid-stream. The watcher keeps the turn.""" + self.ensure_selected(slot) + key = slot.key + assert key is not None + if not self.sessions[key].prompts: + self.submit(slot, slow=False) + watcher = self.rng.choice([s for s in self.slots.values() if s is not slot]) + if watcher.conn is None: + self.connect(watcher) + self.resume(watcher, key) + turn = self.submit(slot, slow=True) + assert watcher.conn is not None and id(watcher.conn) in turn.attached + self.wait_first_delta(turn, watcher.conn) + assert slot.conn is not None + slot.conn.drop() + slot.conn = None + poll_until(lambda: self.complete_frames(watcher.conn, turn), timeout=STEP_TIMEOUT, + what=f"watcher {watcher.name} gets {turn.canary} after the prompting window died") + (frame,) = self.complete_frames(watcher.conn, turn) + assert (frame["params"].get("payload") or {}).get("status", "complete") == "complete", frame + assert self.db_has_reply(turn) + self.submit(watcher, slow=False) + + def op_absence(self, slot: Slot) -> None: + turn = self.op_prompt(slot, slow=True) + conn = slot.conn + assert conn is not None + self.wait_first_delta(turn, conn) + conn.pause() + time.sleep(ABSENCE_S) # the scenario itself (laptop lid closed), not a synchronisation wait + conn.resume() + poll_until(lambda: self.complete_frames(conn, turn), timeout=STEP_TIMEOUT, + what=f"{turn.canary} completion after a {ABSENCE_S}s absence") + (frame,) = self.complete_frames(conn, turn) + assert (frame["params"].get("payload") or {}).get("status", "complete") == "complete", ( + f"client absence ended turn {turn.canary}: {frame['params'].get('payload')}") + + def op_reconnect(self, slot: Slot) -> None: + if slot.conn is not None and not slot.conn.dropped: + slot.conn.drop() + self.connect(slot) + key = slot.key + slot.sid = None + if key and key in self.resumable(): + self.resume(slot, key) + else: + slot.key = None + + # -- invariants -------------------------------------------------------------------------- + def check_events(self) -> None: + for conn in self.conns: + seqs: dict[str, list[int]] = defaultdict(list) + for idx, frame in enumerate(conn.snapshot()): + if frame.get("method") != "event": + continue + params = frame["params"] + sid = params.get("session_id") or "" + if not sid: + continue + first = self.attach_idx[id(conn)].get(sid) + assert first is not None and idx >= first, ( + f"{conn.name} received {params.get('type')} for session {sid} " + f"({self.sid_key.get(sid, 'foreign')}) it never attached to") + if isinstance(params.get("seq"), int): + seqs[sid].append(params["seq"]) + if params.get("type") == "message.complete": + text = str((params.get("payload") or {}).get("text") or "") + for canary in CANARY_RE.findall(text): + owner = self.turns.get(canary) + assert owner is not None and owner.key == self.sid_key.get(sid), ( + f"{conn.name}: reply for {canary} (sent to " + f"{owner.key if owner else '?'}) delivered as session {self.sid_key.get(sid)}") + for sid, values in seqs.items(): + dupes = [s for s, n in Counter(values).items() if n > 1] + assert not dupes, f"{conn.name} got session {sid} events twice: duplicate seq {dupes[:5]}" + assert values == sorted(values), f"{conn.name} got session {sid} events out of order" + + def check_single_runtime(self) -> None: + keys = Counter(r.get("session_key") for r in self.active_rows()) + mine = {k: n for k, n in keys.items() if k in self.sessions and n > 1} + assert not mine, f"chat(s) with more than one live runtime: {mine}" + + def check_final(self) -> None: + for key in list(self.sessions): + if not self.sessions[key].deleted: + self.wait_idle(key) + prefix = f"cnry-s{self.seed}-" + rows = self.be.db_rows( + "SELECT session_id, role, content FROM messages WHERE content LIKE ? ORDER BY id", (f"%{prefix}%",)) + user_by_key: dict[str, list[str]] = defaultdict(list) + replies: Counter = Counter() + full: Counter = Counter() + for key, role, content in rows: + found = [c for c in CANARY_RE.findall(str(content or "")) if c.startswith(prefix)] + if role == "user": + user_by_key[key].extend(found) + elif role == "assistant": + for canary in found: + replies[(key, canary)] += 1 + full[(key, canary)] += str(content).rstrip().endswith("DONE") + for key, sess in self.sessions.items(): + if sess.deleted: + assert not user_by_key.get(key), f"deleted session {key} still has rows" + continue + assert user_by_key.get(key, []) == sess.prompts, ( + f"session {key}: persisted prompts {user_by_key.get(key, [])} != sent {sess.prompts}") + stray = set(user_by_key) - set(self.sessions) + assert not stray, f"prompts persisted to sessions nobody selected: {stray}" + for canary, turn in self.turns.items(): + if self.sessions[turn.key].deleted: + continue + owner_replies = replies.get((turn.key, canary), 0) + elsewhere = {k: n for (k, c), n in replies.items() if c == canary and k != turn.key} + assert not elsewhere, f"reply to {canary} persisted into other session(s) {elsewhere}" + if turn.interrupted: + assert owner_replies <= 1, f"interrupted {canary} has {owner_replies} replies" + else: + assert owner_replies == 1, f"{canary} in {turn.key}: {owner_replies} persisted replies (want 1)" + assert full[(turn.key, canary)] == 1, f"{canary} in {turn.key}: reply was cut short" + survivors = {id(c) for c in self.live_conns()} + for canary, turn in self.turns.items(): + if turn.interrupted or self.sessions[turn.key].deleted: + continue + for conn in self.conns: + if id(conn) in turn.attached and id(conn) in survivors: + got = len(self.complete_frames(conn, turn)) + assert got == 1, f"{conn.name} saw {got} completions of {canary} (want exactly 1)" + + +MANDATORY = ["prompt", "prompt_slow", "interrupt_mid_turn", "switch", "delete_active", "close_and_delete", + "drop_mid_turn", "orphan_idle", "orphan_running", "shared_owner_drop", "absence", "reconnect"] +EXTRA = ["prompt", "prompt_slow", "switch", "switch", "prompt", "create", "delete_active"] + + +@pytest.mark.parametrize("seed", [11, 23, 37]) +def test_multiclient_session_model(backend: Backend, seed: int) -> None: + model = Model(backend, seed) + rng = model.rng + for slot in model.slots.values(): # every client starts with its own chat + model.connect(slot) + model.op_create(slot) + plan = MANDATORY + [rng.choice(EXTRA) for _ in range(6)] + rng.shuffle(plan) + try: + for step, op in enumerate(plan): + slot = model.slots[rng.choice(sorted(model.slots))] + model.log.append(f"{step}:{op}:{slot.name}") + t0 = time.monotonic() + if op == "prompt_slow": + model.op_prompt(slot, slow=True) + else: + getattr(model, f"op_{op}")(slot) + model.check_events() + model.check_single_runtime() + model.log.append(f"{time.monotonic() - t0:.1f}s") + model.check_final() + model.check_events() + print(f"seed={seed} steps={model.log} turns={len(model.turns)} sessions={len(model.sessions)}") + except AssertionError as exc: + raise AssertionError(f"seed={seed} steps={model.log}\n{exc}\n{backend.logs(25)}") from None + finally: + for conn in model.conns: + conn.close()