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 67de93862c (#98028/#100325, orphan reaper ignores turn
activity); semantic revert of de25545dce (rebind instead of fan-out); lease never
released; session.delete of a live session not refused.
(cherry picked from commit fe078a891231d3ba459da4ca46bf2a06979f4444)
This commit is contained in:
0
tests/e2e/core/terminal/__init__.py
Normal file
0
tests/e2e/core/terminal/__init__.py
Normal file
308
tests/e2e/core/terminal/_gateway_client.py
Normal file
308
tests/e2e/core/terminal/_gateway_client.py
Normal file
@@ -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 ""
|
||||
521
tests/e2e/core/terminal/test_multiclient_session_model.py
Normal file
521
tests/e2e/core/terminal/test_multiclient_session_model.py
Normal file
@@ -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 <canary> ... 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()
|
||||
Reference in New Issue
Block a user