fix(gateway): relay compute-host clarify state
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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 "")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user