fix(gateway): relay compute-host clarify state

This commit is contained in:
konsisumer
2026-08-30 14:51:23 +02:00
committed by Teknium
parent 2e9a39d28e
commit 3ae74119c9
6 changed files with 303 additions and 12 deletions
+98
View File
@@ -496,6 +496,104 @@ def test_compute_host_turn_end_updates_metadata_mirror(monkeypatch):
server._sessions.pop("iso-sid", None)
def test_compute_host_clarify_snapshot_replays_and_proxies_batch_answers(monkeypatch):
"""A host-owned clarify survives activation and receives its UI answers."""
class _Supervisor:
def __init__(self):
self.responses = []
def respond(self, sid, params, *, timeout=15.0):
self.responses.append((sid, dict(params), timeout))
remaining = ["q1"] if params.get("question_id") == "q0" else []
return {"type": "respond.ack", "response": {"result": {"status": "ok", "remaining": remaining}}}
sid = "host-clarify"
supervisor = _Supervisor()
session = _session(agent=None, agent_ready=threading.Event(), _compute_host_active=True)
server._sessions[sid] = session
monkeypatch.setattr(server, "_load_cfg", lambda: {"dashboard": {"turn_isolation": True}})
monkeypatch.setattr(server, "_get_compute_host_supervisor", lambda _cfg=None: supervisor)
monkeypatch.setattr(server, "write_json", lambda _message: True)
try:
server._relay_compute_host_rpc(
{
"jsonrpc": "2.0",
"method": "event",
"params": {
"type": "clarify.request",
"session_id": sid,
"payload": {
"request_id": "host-request",
"questions": [
{"qid": "q0", "question": "First?", "choices": ["a"]},
{"qid": "q1", "question": "Second?", "choices": ["b"]},
],
},
},
}
)
activated = server._live_session_payload(sid, session)
assert activated["pending_clarify"]["request_id"] == "host-request"
response = server.handle_request(
{
"id": "clarify-q0",
"method": "clarify.respond",
"params": {"request_id": "host-request", "question_id": "q0", "answer": "a"},
}
)
assert response["result"] == {"status": "ok", "remaining": ["q1"]}
assert supervisor.responses == [
(sid, {"request_id": "host-request", "question_id": "q0", "answer": "a"}, 15.0)
]
replayed = server._live_session_payload(sid, session)["pending_clarify"]
assert replayed["answers"] == {"q0": "a"}
final_response = server.handle_request(
{
"id": "clarify-q1",
"method": "clarify.respond",
"params": {"request_id": "host-request", "question_id": "q1", "answer": "b"},
}
)
assert final_response["result"] == {"status": "ok", "remaining": []}
assert "pending_clarify" not in server._live_session_payload(sid, session)
finally:
server._sessions.pop(sid, None)
def test_compute_host_interrupt_forwards_when_parent_running_mirror_is_stale(monkeypatch):
"""The host, not the parent's mirrored running flag, owns interruption."""
interrupted = []
class _Supervisor:
def interrupt(self, sid, *, request_id=None):
interrupted.append((sid, request_id))
sid = "host-stale-running"
server._sessions[sid] = _session(
agent=None,
agent_ready=threading.Event(),
_compute_host_active=True,
running=False,
)
monkeypatch.setattr(server, "_load_cfg", lambda: {"dashboard": {"turn_isolation": True}})
monkeypatch.setattr(server, "_get_compute_host_supervisor", lambda _cfg=None: _Supervisor())
try:
response = server.handle_request(
{"id": "interrupt", "method": "session.interrupt", "params": {"session_id": sid}}
)
assert response["result"] == {"status": "interrupted", "turn_isolation": True}
assert interrupted == [(sid, "interrupt-interrupt")]
finally:
server._sessions.pop(sid, None)
def test_slash_exec_compress_flag_on_applies_host_control_mirror(monkeypatch):
class _ExplodingWorker:
def __init__(self, *args, **kwargs):
@@ -48,6 +48,40 @@ def test_compute_host_workers_inherit_tui_pool_env_or_8(monkeypatch):
assert _default_workers() == 8
def test_compute_host_routes_clarify_response_to_child_pending_registry(monkeypatch):
"""Interactive answers are handled in the process that owns `_pending`."""
out = io.StringIO()
host = ComputeHost(stdout=out, heartbeat_secs=0)
sid = "host-clarify"
server._sessions[sid] = {"history_lock": threading.Lock()}
calls = []
monkeypatch.setitem(
server._methods,
"clarify.respond",
lambda rid, params: calls.append((rid, dict(params))) or {"result": {"status": "ok"}},
)
try:
host._handle_respond(
{
"sid": sid,
"request_id": "relay-response",
"params": {"request_id": "clarify-request", "answer": "yes"},
}
)
assert calls == [("relay-response", {"request_id": "clarify-request", "answer": "yes"})]
frame = _json_lines(out)[-1]
assert frame == {
"type": "respond.ack",
"sid": sid,
"request_id": "relay-response",
"response": {"result": {"status": "ok"}},
"host_ns": frame["host_ns"],
}
finally:
server._sessions.pop(sid, None)
host.close()
def test_mutator_route_table_matches_prd_inventory():
assert MUTATOR_ROUTE_TABLE == {
"prompt.submit": "turn-path",
+54 -8
View File
@@ -272,6 +272,8 @@ class ComputeHost:
self._handle_turn_start(frame)
elif kind == "interrupt":
self._handle_interrupt(frame)
elif kind == "respond":
self._handle_respond(frame)
elif kind == "reload_mcp":
self._handle_reload_mcp(frame)
elif kind == "control":
@@ -363,18 +365,62 @@ class ComputeHost:
if session is None:
self.emit({"type": "interrupt.ack", "sid": sid, "request_id": frame.get("request_id"), "applied": False})
return
agent = session.get("agent")
if agent is not None:
request_hard_interrupt(agent)
with session.get("history_lock", threading.Lock()):
session["_turn_cancel_requested"] = True
session["queued_prompt"] = None
session.pop("queued_prompts", None)
session["_queued_prompt_generation"] = int(session.get("_queued_prompt_generation", 0)) + 1
# In the child, `_session_uses_compute_host()` is false, so the
# shared helper interrupts the local agent and releases this
# process's pending clarify Event. The parent has only a metadata
# mirror and cannot release the prompt that is blocking the turn.
server._interrupt_session_turn(sid, session)
self.emit({"type": "interrupt.ack", "sid": sid, "request_id": frame.get("request_id"), "applied": True, "applied_ns": now_ns()})
except Exception as exc:
self.emit({"type": "interrupt.ack", "sid": sid, "request_id": frame.get("request_id"), "applied": False, "message": str(exc)})
def _handle_respond(self, frame: dict[str, Any]) -> None:
"""Resolve an interactive request in the host-owned pending registry."""
sid = str(frame.get("sid") or "")
request_id = frame.get("request_id")
try:
from tui_gateway import server
if sid not in server._sessions:
self.emit(
{
"type": "respond.error",
"sid": sid,
"request_id": request_id,
"message": "session not found",
}
)
return
params = frame.get("params")
if not isinstance(params, dict):
self.emit(
{
"type": "respond.error",
"sid": sid,
"request_id": request_id,
"message": "response params must be an object",
}
)
return
response = server._methods["clarify.respond"](request_id, params)
self.emit(
{
"type": "respond.ack",
"sid": sid,
"request_id": request_id,
"response": response,
}
)
except Exception as exc:
self.emit(
{
"type": "respond.error",
"sid": sid,
"request_id": request_id,
"message": str(exc),
}
)
def _run_spike_turn(self, session: HostSession, frame: dict[str, Any]) -> None:
request_id = frame.get("request_id") or uuid.uuid4().hex
prompt = str(frame.get("prompt") or frame.get("text") or "")
+22 -1
View File
@@ -270,6 +270,27 @@ class HostSupervisor:
self.start()
self._send_frame({"type": "interrupt", "sid": sid, "request_id": request_id or uuid.uuid4().hex})
def respond(self, sid: str, params: dict[str, Any], *, timeout: float = 15.0) -> dict:
"""Deliver an interactive prompt response to the host that owns it."""
self.start()
request_id = uuid.uuid4().hex
q: queue.Queue[dict] = queue.Queue(maxsize=1)
with self._lock:
self._pending_controls[request_id] = q
try:
self._send_frame(
{
"type": "respond",
"sid": sid,
"request_id": request_id,
"params": dict(params),
}
)
return q.get(timeout=timeout)
finally:
with self._lock:
self._pending_controls.pop(request_id, None)
def reload_mcp(self, sid: str, *, request_id: str | None = None) -> dict:
return self.control(
sid,
@@ -430,7 +451,7 @@ class HostSupervisor:
if ftype in {"turn.end", "turn.error"}:
self._complete_turn(frame)
return
if ftype in {"control.ack", "control.error", "interrupt.ack", "reload_mcp.ack", "shutdown.ack"}:
if ftype in {"control.ack", "control.error", "respond.ack", "respond.error", "interrupt.ack", "reload_mcp.ack", "shutdown.ack"}:
request_id = str(frame.get("request_id") or "")
with self._lock:
q = self._pending_controls.get(request_id)
+2
View File
@@ -1704,6 +1704,8 @@ def _(rid, params: dict) -> dict:
# from _pending) while the card is still visible — common when a WebSocket
# reconnect during the wait drops tool.complete. A late answer must resolve
# gracefully instead of hitting the raw 4009 "no pending answer request".
if proxied := _respond_compute_host_clarify(rid, params):
return proxied
return _respond(rid, params, "answer", allow_expired=True)
+93 -3
View File
@@ -1245,8 +1245,10 @@ def _interrupt_session_turn(
run_thread_alive = False
if use_compute_host:
if should_interrupt:
_get_compute_host_supervisor().interrupt(sid, request_id=request_id)
# The host owns the live turn. Parent `running` is only a mirror and
# can lag behind a blocked interactive tool, so let the host determine
# whether there is work to interrupt.
_get_compute_host_supervisor().interrupt(sid, request_id=request_id)
else:
run_thread = session.get("_run_thread")
run_thread_alive = run_thread is not None and run_thread.is_alive()
@@ -2599,7 +2601,7 @@ def _get_compute_host_supervisor(cfg: dict | None = None):
from tui_gateway.host_supervisor import HostSupervisor
_compute_host_supervisor = HostSupervisor(
rpc_sink=write_json,
rpc_sink=_relay_compute_host_rpc,
heartbeat_secs=int(isolation_cfg.get("compute_host_heartbeat_secs") or 15),
respawn_max=int(isolation_cfg.get("compute_host_respawn_max") or 3),
)
@@ -2650,6 +2652,87 @@ def _metadata_mirror(session: dict | None) -> dict:
return mirror if isinstance(mirror, dict) else {}
def _relay_compute_host_rpc(message: dict) -> bool:
"""Relay host events while retaining the clarify snapshot needed on resume."""
params = message.get("params") if isinstance(message, dict) else None
if isinstance(params, dict) and params.get("type") == "clarify.request":
sid = str(params.get("session_id") or "")
payload = params.get("payload")
session = _sessions.get(sid)
if session is not None and isinstance(payload, dict) and payload.get("request_id"):
with session.get("history_lock", threading.Lock()):
session["_compute_host_pending_clarify"] = dict(payload)
elif isinstance(params, dict) and params.get("type") == "clarify.expire":
sid = str(params.get("session_id") or "")
payload = params.get("payload")
session = _sessions.get(sid)
request_id = payload.get("request_id") if isinstance(payload, dict) else None
if session is not None and request_id:
with session.get("history_lock", threading.Lock()):
pending = session.get("_compute_host_pending_clarify")
if isinstance(pending, dict) and pending.get("request_id") == request_id:
session.pop("_compute_host_pending_clarify", None)
return write_json(message)
def _compute_host_clarify_session(request_id: str) -> tuple[str, dict] | None:
"""Find the parent mirror for one host-owned clarify request."""
if not request_id:
return None
for sid, session in list(_sessions.items()):
with session.get("history_lock", threading.Lock()):
pending = session.get("_compute_host_pending_clarify")
if isinstance(pending, dict) and pending.get("request_id") == request_id:
return sid, session
return None
def _update_compute_host_clarify_snapshot(sid: str, session: dict, params: dict, result: dict) -> None:
"""Keep reconnect snapshots accurate while a batch clarify is answered."""
request_id = str(params.get("request_id") or "")
with session.get("history_lock", threading.Lock()):
pending = session.get("_compute_host_pending_clarify")
if not isinstance(pending, dict) or pending.get("request_id") != request_id:
return
if result.get("status") == "expired" or not result.get("remaining") and not params.get("question_id"):
session.pop("_compute_host_pending_clarify", None)
return
question_id = str(params.get("question_id") or "")
if question_id and isinstance(result.get("remaining"), list):
answers = dict(pending.get("answers") or {})
answers[question_id] = str(params.get("answer") or "")
pending["answers"] = answers
if not result["remaining"]:
session.pop("_compute_host_pending_clarify", None)
def _respond_compute_host_clarify(rid: str, params: dict) -> dict | None:
"""Proxy a clarify answer into the process that owns its pending Event."""
located = _compute_host_clarify_session(str(params.get("request_id") or ""))
if located is None:
return None
sid, session = located
if not _session_uses_compute_host(session):
return None
try:
ack = _get_compute_host_supervisor().respond(sid, params)
except Exception as exc:
return _err(rid, 5019, f"compute-host clarify response failed: {exc}")
if ack.get("type") == "respond.error":
return _err(rid, 5019, str(ack.get("message") or "compute-host clarify response failed"))
response = ack.get("response")
if not isinstance(response, dict):
return _err(rid, 5019, "compute-host clarify response returned an invalid response")
if "error" in response:
error = response["error"] if isinstance(response["error"], dict) else {}
return _err(rid, int(error.get("code") or 5000), str(error.get("message") or "clarify response failed"))
result = response.get("result")
if not isinstance(result, dict):
return _err(rid, 5019, "compute-host clarify response returned an invalid result")
_update_compute_host_clarify_snapshot(sid, session, params, result)
return _ok(rid, result)
def _apply_compute_host_metadata_mirror(session: dict, frame: dict | None) -> None:
"""Mirror host-owned session metadata in the serving process.
@@ -2699,6 +2782,7 @@ def _on_compute_host_turn_done(rid: str, sid: str, session: dict, frame: dict) -
session["running"] = False
session["last_active"] = time.time()
_clear_inflight_turn(session)
session.pop("_compute_host_pending_clarify", None)
if is_error:
message = str(frame.get("message") or "compute host turn failed")
_emit("message.complete", sid, {"text": f"Error: {message}", "status": "error"})
@@ -2818,6 +2902,12 @@ def _pending_clarify_request_payload(sid: str) -> dict | None:
if batch is not None and batch["answers"]:
snapshot["answers"] = dict(batch["answers"])
return snapshot
session = _sessions.get(sid)
if session is not None:
with session.get("history_lock", threading.Lock()):
pending = session.get("_compute_host_pending_clarify")
if isinstance(pending, dict):
return dict(pending)
return None