test(tui): synchronize compute-host protocol frames
Compute-host protocol tests polled a thread-unsafe StringIO on a five-second wall-clock deadline while unrelated provider discovery and persistence work competed under CI load. Isolate those side paths and wait causally for complete JSON-line frames so strict turn ordering remains observable without scheduling races.
This commit is contained in:
59
tests/tui_gateway/_compute_host_frames.py
Normal file
59
tests/tui_gateway/_compute_host_frames.py
Normal file
@@ -0,0 +1,59 @@
|
||||
"""Causal, thread-safe compute-host frame capture for protocol tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
|
||||
class FrameSink:
|
||||
"""Capture complete JSON-line frames and wake waiters when one is published."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._condition = threading.Condition()
|
||||
self._pending = ""
|
||||
self._frames: list[dict] = []
|
||||
|
||||
def write(self, data: str) -> int:
|
||||
with self._condition:
|
||||
self._pending += data
|
||||
published = False
|
||||
while "\n" in self._pending:
|
||||
line, self._pending = self._pending.split("\n", 1)
|
||||
if line.strip():
|
||||
self._frames.append(json.loads(line))
|
||||
published = True
|
||||
if published:
|
||||
self._condition.notify_all()
|
||||
return len(data)
|
||||
|
||||
def flush(self) -> None:
|
||||
pass
|
||||
|
||||
def frames(self) -> list[dict]:
|
||||
with self._condition:
|
||||
return list(self._frames)
|
||||
|
||||
def wait_for(self, predicate: Callable[[dict], bool], timeout: float = 20.0) -> dict:
|
||||
"""Wait on frame publication; ``timeout`` is only a deadlock guard."""
|
||||
deadline = time.monotonic() + timeout
|
||||
with self._condition:
|
||||
while True:
|
||||
for frame in self._frames:
|
||||
if predicate(frame):
|
||||
return frame
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
raise AssertionError(f"timed out; saw={self._frames}")
|
||||
self._condition.wait(remaining)
|
||||
|
||||
|
||||
def start_test_work(target, *, name: str, session: dict | None = None) -> threading.Thread:
|
||||
"""Start the real turn thread without process-retirement/import side paths."""
|
||||
thread = threading.Thread(target=target, name=name)
|
||||
if session is not None:
|
||||
session["_run_thread"] = thread
|
||||
thread.start()
|
||||
return thread
|
||||
@@ -14,33 +14,21 @@ rotation and held past ``session.close`` until the child's turn settles.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import types
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.tui_gateway._compute_host_frames import FrameSink, start_test_work
|
||||
from tui_gateway import server
|
||||
from tui_gateway.compute_host import ComputeHost
|
||||
|
||||
|
||||
def _frames(out: io.StringIO) -> list[dict]:
|
||||
return [json.loads(line) for line in out.getvalue().splitlines() if line.strip()]
|
||||
|
||||
|
||||
def _wait(out: io.StringIO, predicate, timeout: float = 5.0) -> dict:
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
for frame in _frames(out):
|
||||
if predicate(frame):
|
||||
return frame
|
||||
time.sleep(0.01)
|
||||
raise AssertionError(f"timed out; saw={_frames(out)}")
|
||||
|
||||
|
||||
def _stub_agent(deltas: list[str]) -> types.SimpleNamespace:
|
||||
def run_conversation(prompt, *, conversation_history=None, stream_callback=None, **_kw):
|
||||
final = "".join(deltas)
|
||||
@@ -108,11 +96,19 @@ def isolated_env(monkeypatch, tmp_path):
|
||||
monkeypatch.setattr(server, "_session_cwd", lambda session: str(tmp_path))
|
||||
monkeypatch.setattr(server, "_register_session_cwd", lambda session: None)
|
||||
monkeypatch.setattr(server, "_tts_stream_begin", lambda: None)
|
||||
monkeypatch.setattr(server, "_ensure_session_db_row", lambda session: True)
|
||||
monkeypatch.setattr(server, "_persist_branch_seed", lambda session: None)
|
||||
monkeypatch.setattr(server, "_routing_provenance_db", lambda session: contextlib.nullcontext(None))
|
||||
monkeypatch.setattr(server, "_reopen_routed_session_row", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_record_turn_marker", lambda *a, **k: "marker")
|
||||
monkeypatch.setattr(server, "_retire_turn_marker", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_start_session_work", start_test_work)
|
||||
monkeypatch.setattr(server, "_get_usage", lambda agent_: {})
|
||||
monkeypatch.setattr(server, "_hydrate_session_cwd", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_wire_session_agent", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_start_session_services", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_schedule_mcp_late_refresh", lambda *a, **k: None)
|
||||
monkeypatch.setitem(sys.modules, "hermes_undo", types.SimpleNamespace(on_user_message_appended=lambda key: None))
|
||||
import tui_gateway.prompt_turn as prompt_turn
|
||||
|
||||
for mod in (server, prompt_turn):
|
||||
@@ -123,16 +119,16 @@ def isolated_env(monkeypatch, tmp_path):
|
||||
server._sessions.pop(sid, None)
|
||||
|
||||
|
||||
def _run_turn(frame: dict, timeout: float = 5.0) -> tuple[list[dict], dict | None]:
|
||||
def _run_turn(frame: dict, timeout: float = 20.0) -> tuple[list[dict], dict | None]:
|
||||
"""Run one turn.start through the real child path; return (all frames, turn.end frame)."""
|
||||
out = io.StringIO()
|
||||
out = FrameSink()
|
||||
host = ComputeHost(stdout=out, heartbeat_secs=0)
|
||||
try:
|
||||
host.handle_frame(frame)
|
||||
end = _wait(out, lambda f: f["type"] == "turn.end", timeout=timeout)
|
||||
end = out.wait_for(lambda f: f["type"] == "turn.end", timeout=timeout)
|
||||
finally:
|
||||
host.close()
|
||||
return _frames(out), end
|
||||
return out.frames(), end
|
||||
|
||||
|
||||
# ── Fix 1: the child borrows instead of re-claiming ─────────────────────────
|
||||
|
||||
@@ -9,44 +9,46 @@ a ``server._sessions`` entry whose agent runs on the turn worker.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import json
|
||||
import contextlib
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import types
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.tui_gateway._compute_host_frames import FrameSink, start_test_work
|
||||
from tui_gateway import server
|
||||
from tui_gateway.compute_host import ComputeHost
|
||||
|
||||
|
||||
def _frames(out: io.StringIO) -> list[dict]:
|
||||
return [json.loads(line) for line in out.getvalue().splitlines() if line.strip()]
|
||||
|
||||
|
||||
def _wait(out: io.StringIO, predicate, timeout: float = 5.0) -> dict:
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
for frame in _frames(out):
|
||||
if predicate(frame):
|
||||
return frame
|
||||
time.sleep(0.01)
|
||||
raise AssertionError(f"timed out; saw={_frames(out)}")
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def turn_env(monkeypatch, tmp_path):
|
||||
"""Neutralize the turn pipeline's environment-heavy side paths (same set as the
|
||||
prompt.submit tests) so the frame protocol is what's under test. Threads stay REAL:
|
||||
``_run_real_turn`` joins ``session["_run_thread"]`` itself before emitting ``turn.end``."""
|
||||
monkeypatch.setattr(server, "_wire_callbacks", lambda sid: None)
|
||||
monkeypatch.setattr(server, "_apply_pending_model_switch", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_sync_agent_model_with_config", lambda sid, session: None)
|
||||
monkeypatch.setattr(server, "_sync_agent_compression_with_config", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_sync_agent_fallback_with_config", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_sync_bot_capabilities", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_adopt_out_of_band_turns", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_session_cwd", lambda session: str(tmp_path))
|
||||
monkeypatch.setattr(server, "_register_session_cwd", lambda session: None)
|
||||
monkeypatch.setattr(server, "_tts_stream_begin", lambda: None)
|
||||
monkeypatch.setattr(server, "_ensure_session_db_row", lambda session: True)
|
||||
monkeypatch.setattr(server, "_persist_branch_seed", lambda session: None)
|
||||
monkeypatch.setattr(server, "_routing_provenance_db", lambda session: contextlib.nullcontext(None))
|
||||
monkeypatch.setattr(server, "_reopen_routed_session_row", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_record_turn_marker", lambda *a, **k: "marker")
|
||||
monkeypatch.setattr(server, "_retire_turn_marker", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_start_session_work", start_test_work)
|
||||
monkeypatch.setattr(server, "_sync_session_key_after_compress", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_get_usage", lambda agent: {})
|
||||
monkeypatch.setattr(server, "_start_turn_voice", lambda: (None, False))
|
||||
monkeypatch.setattr(server, "_start_usage_ticker", lambda *a, **k: (threading.Event(), start_test_work(lambda: None, name="usage")))
|
||||
monkeypatch.setitem(sys.modules, "hermes_undo", types.SimpleNamespace(on_user_message_appended=lambda key: None))
|
||||
|
||||
|
||||
def _agent(deltas: list[str], *, delay_s: float = 0.0, interrupt: threading.Event | None = None):
|
||||
@@ -80,18 +82,18 @@ def _session(agent) -> dict:
|
||||
|
||||
|
||||
def test_turn_start_streams_deltas_then_turn_end_with_history_identity(turn_env):
|
||||
out = io.StringIO()
|
||||
out = FrameSink()
|
||||
host = ComputeHost(stdout=out, heartbeat_secs=0)
|
||||
sid = "s1"
|
||||
server._sessions[sid] = _session(_agent(["a ", "b ", "c "]))
|
||||
try:
|
||||
host.handle_frame({"type": "turn.start", "sid": sid, "request_id": "turn", "prompt": "hello"})
|
||||
end = _wait(out, lambda f: f["type"] == "turn.end")
|
||||
end = out.wait_for(lambda f: f["type"] == "turn.end")
|
||||
finally:
|
||||
server._sessions.pop(sid, None)
|
||||
host.close()
|
||||
|
||||
frames = _frames(out)
|
||||
frames = out.frames()
|
||||
kinds = [f["type"] for f in frames]
|
||||
assert kinds[0] == "turn.started"
|
||||
assert kinds[-1] == "turn.end"
|
||||
@@ -117,18 +119,18 @@ def test_turn_start_streams_deltas_then_turn_end_with_history_identity(turn_env)
|
||||
|
||||
|
||||
def test_turn_start_without_sid_is_a_turn_error(turn_env):
|
||||
out = io.StringIO()
|
||||
out = FrameSink()
|
||||
host = ComputeHost(stdout=out, heartbeat_secs=0)
|
||||
try:
|
||||
host.handle_frame({"type": "turn.start", "request_id": "nosid", "prompt": "x"})
|
||||
err = _wait(out, lambda f: f["type"] == "turn.error")
|
||||
err = out.wait_for(lambda f: f["type"] == "turn.error")
|
||||
finally:
|
||||
host.close()
|
||||
assert err["request_id"] == "nosid" and err["message"] == "sid required"
|
||||
|
||||
|
||||
def test_second_turn_start_while_running_is_session_busy(turn_env):
|
||||
out = io.StringIO()
|
||||
out = FrameSink()
|
||||
host = ComputeHost(stdout=out, heartbeat_secs=0)
|
||||
sid = "s1"
|
||||
session = _session(_agent(["x"]))
|
||||
@@ -136,7 +138,7 @@ def test_second_turn_start_while_running_is_session_busy(turn_env):
|
||||
server._sessions[sid] = session
|
||||
try:
|
||||
host.handle_frame({"type": "turn.start", "sid": sid, "request_id": "t2", "prompt": "hi"})
|
||||
err = _wait(out, lambda f: f["type"] == "turn.error")
|
||||
err = out.wait_for(lambda f: f["type"] == "turn.error")
|
||||
finally:
|
||||
server._sessions.pop(sid, None)
|
||||
host.close()
|
||||
@@ -145,7 +147,7 @@ def test_second_turn_start_while_running_is_session_busy(turn_env):
|
||||
|
||||
def test_stale_queued_prompt_generation_ends_turn_as_interrupted(turn_env):
|
||||
"""A queued prompt whose generation was bumped by an interrupt must not run."""
|
||||
out = io.StringIO()
|
||||
out = FrameSink()
|
||||
host = ComputeHost(stdout=out, heartbeat_secs=0)
|
||||
sid = "s1"
|
||||
session = _session(_agent(["x"]))
|
||||
@@ -154,27 +156,27 @@ def test_stale_queued_prompt_generation_ends_turn_as_interrupted(turn_env):
|
||||
try:
|
||||
host.handle_frame({"type": "turn.start", "sid": sid, "request_id": "q", "prompt": "hi",
|
||||
"queued_prompt_generation": 2})
|
||||
end = _wait(out, lambda f: f["type"] == "turn.end")
|
||||
end = out.wait_for(lambda f: f["type"] == "turn.end")
|
||||
finally:
|
||||
server._sessions.pop(sid, None)
|
||||
host.close()
|
||||
assert end["interrupted"] is True and end["request_id"] == "q"
|
||||
assert not any(f["type"] == "turn.started" for f in _frames(out))
|
||||
assert not any(f["type"] == "turn.started" for f in out.frames())
|
||||
|
||||
|
||||
def test_interrupt_frame_acks_and_marks_turn_interrupted(turn_env):
|
||||
"""The turn runs on the host worker while ``interrupt`` arrives on the control path."""
|
||||
out = io.StringIO()
|
||||
out = FrameSink()
|
||||
host = ComputeHost(stdout=out, heartbeat_secs=0)
|
||||
sid = "s1"
|
||||
stop = threading.Event()
|
||||
server._sessions[sid] = _session(_agent([f"{i:03d} " for i in range(200)], delay_s=0.01, interrupt=stop))
|
||||
try:
|
||||
host.handle_frame({"type": "turn.start", "sid": sid, "request_id": "turn", "prompt": "go"})
|
||||
_wait(out, lambda f: f["type"] == "rpc" and f["message"]["params"]["type"] == "message.delta")
|
||||
out.wait_for(lambda f: f["type"] == "rpc" and f["message"]["params"]["type"] == "message.delta")
|
||||
host.handle_frame({"type": "interrupt", "sid": sid, "request_id": "stop"})
|
||||
ack = _wait(out, lambda f: f["type"] == "interrupt.ack")
|
||||
end = _wait(out, lambda f: f["type"] == "turn.end")
|
||||
ack = out.wait_for(lambda f: f["type"] == "interrupt.ack")
|
||||
end = out.wait_for(lambda f: f["type"] == "turn.end")
|
||||
finally:
|
||||
stop.set()
|
||||
server._sessions.pop(sid, None)
|
||||
@@ -182,19 +184,19 @@ def test_interrupt_frame_acks_and_marks_turn_interrupted(turn_env):
|
||||
|
||||
assert ack["applied"] is True and ack["request_id"] == "stop" and "applied_ns" in ack
|
||||
assert end["interrupted"] is True
|
||||
deltas = sum(1 for f in _frames(out) if f["type"] == "rpc" and f["message"]["params"]["type"] == "message.delta")
|
||||
deltas = sum(1 for f in out.frames() if f["type"] == "rpc" and f["message"]["params"]["type"] == "message.delta")
|
||||
assert 0 < deltas < 200
|
||||
|
||||
|
||||
def test_unknown_frame_type_is_an_error():
|
||||
out = io.StringIO()
|
||||
out = FrameSink()
|
||||
host = ComputeHost(stdout=out, heartbeat_secs=0)
|
||||
try:
|
||||
host.handle_frame({"type": "bogus", "request_id": "b"})
|
||||
finally:
|
||||
host.close()
|
||||
assert _frames(out) == [{"type": "error", "request_id": "b", "message": "unknown frame type: bogus",
|
||||
"host_ns": _frames(out)[0]["host_ns"]}]
|
||||
assert out.frames() == [{"type": "error", "request_id": "b", "message": "unknown frame type: bogus",
|
||||
"host_ns": out.frames()[0]["host_ns"]}]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("kind", ["legacy", "hard-only", "dynamic-getattr"])
|
||||
@@ -224,7 +226,7 @@ def test_compute_host_interrupt_uses_explicit_stop_compatibility(monkeypatch, ki
|
||||
agent = {"legacy": _Legacy(), "hard-only": _HardOnly(), "dynamic-getattr": _Dynamic()}[kind]
|
||||
# The child never routes back to a supervisor (HERMES_COMPUTE_HOST_CHILD=1 in production).
|
||||
monkeypatch.setenv("HERMES_COMPUTE_HOST_CHILD", "1")
|
||||
out = io.StringIO()
|
||||
out = FrameSink()
|
||||
host = ComputeHost(stdout=out, heartbeat_secs=0)
|
||||
sid = "s1"
|
||||
session = _session(agent)
|
||||
@@ -237,7 +239,7 @@ def test_compute_host_interrupt_uses_explicit_stop_compatibility(monkeypatch, ki
|
||||
host.close()
|
||||
|
||||
assert calls == ["hard" if kind == "hard-only" else "legacy"]
|
||||
ack = _frames(out)[-1]
|
||||
ack = out.frames()[-1]
|
||||
assert ack["type"] == "interrupt.ack" and ack["applied"] is True
|
||||
assert session["_turn_cancel_requested"] is True
|
||||
|
||||
@@ -245,7 +247,7 @@ def test_compute_host_interrupt_uses_explicit_stop_compatibility(monkeypatch, ki
|
||||
def test_host_builds_the_session_agent_with_the_frame_login(monkeypatch):
|
||||
"""The host process has no record for a first turn, and its pipe names no login, so the agent is built
|
||||
from the login the frame carries and the new record keeps it for later rebuilds."""
|
||||
host = ComputeHost(stdout=io.StringIO(), heartbeat_secs=0)
|
||||
host = ComputeHost(stdout=FrameSink(), heartbeat_secs=0)
|
||||
captured = {}
|
||||
|
||||
def fake_make_agent(sid, key, **kwargs):
|
||||
|
||||
Reference in New Issue
Block a user