From 879d6a4c78e9f9df7ef4e2946c8e222cdfda1e2b Mon Sep 17 00:00:00 2001 From: ethernet Date: Tue, 18 Aug 2026 16:00:19 -0400 Subject: [PATCH] feat(tui_gateway): batch clarify bridge with per-question locks One clarify.request carries the question list (qid, question, choices, multi_select per entry). clarify.respond gains an optional question_id: each respond locks one answer, a repeat respond overwrites it, and the batch resolves when every question is locked. A respond without question_id keeps its existing meaning (cancel the whole prompt). Locked answers survive the deadline: a timed-out batch returns the partial answer map with a timed_out flag instead of an empty string. The reconnect replay snapshot also carries the locked answers, so a reattached client restores its per-question state. Both agent-side clarify dispatch sites forward the questions arg. --- agent/agent_runtime_helpers.py | 1 + agent/tool_executor.py | 1 + tests/tui_gateway/test_protocol.py | 204 +++++++++++++++++++++++++++++ tui_gateway/server.py | 124 +++++++++++++++--- 4 files changed, 314 insertions(+), 16 deletions(-) diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index 96f7bb821a..5d87894a65 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -3200,6 +3200,7 @@ def invoke_tool(agent, function_name: str, function_args: dict, effective_task_i question=next_args.get("question", ""), choices=next_args.get("choices"), multi_select=next_args.get("multi_select", False), + questions=next_args.get("questions"), callback=agent.clarify_callback, ), next_args, diff --git a/agent/tool_executor.py b/agent/tool_executor.py index 381f1000e9..4fedb6dc7d 100644 --- a/agent/tool_executor.py +++ b/agent/tool_executor.py @@ -2120,6 +2120,7 @@ def execute_tool_calls_sequential(agent, assistant_message, messages: list, effe question=next_args.get("question", ""), choices=next_args.get("choices"), multi_select=next_args.get("multi_select", False), + questions=next_args.get("questions"), callback=agent.clarify_callback, ) function_result, function_args, middleware_trace, _execution_blocked, _execution_dispatched = _managed_values(_run_agent_tool_execution_middleware( diff --git a/tests/tui_gateway/test_protocol.py b/tests/tui_gateway/test_protocol.py index 7b27ae6253..c2cc809591 100644 --- a/tests/tui_gateway/test_protocol.py +++ b/tests/tui_gateway/test_protocol.py @@ -354,6 +354,210 @@ def test_late_prompt_response_is_idempotent(server, method, value_key): assert response["result"] == {"status": "expired"} +# ── clarify batch (multi-question) bridge ──────────────────────────── + + +def _drain_batch_block(server, qids, timeout=5, payload=None): + """Run a batch _block on a worker thread and return (thread, result box, + emitted request payload). The caller resolves questions via + handle_request and then joins.""" + box = {} + + def run(): + box["answer"] = server._block( + "clarify.request", + "s1", + dict(payload or {"questions": [{"qid": q, "question": q} for q in qids]}), + timeout=timeout, + batch_qids=list(qids), + ) + + thread = threading.Thread(target=run, daemon=True) + thread.start() + # Wait for the request to be registered so respond calls can find it. + deadline = time.monotonic() + 2 + while time.monotonic() < deadline: + with server._prompt_lock: + if server._batch_clarify: + rid = next(iter(server._batch_clarify)) + return thread, box, rid + time.sleep(0.01) + raise AssertionError("batch clarify request never registered") + + +def test_clarify_batch_resolves_when_all_questions_locked(capture): + server, buf = capture + thread, box, rid = _drain_batch_block(server, ["q0", "q1"]) + + first = server.handle_request({ + "id": "a1", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q1", "answer": "beta"}, + }) + assert first["result"]["status"] == "ok" + assert first["result"]["remaining"] == ["q0"] + assert thread.is_alive() # one question left — still blocking + + second = server.handle_request({ + "id": "a2", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q0", "answer": "alpha"}, + }) + assert second["result"]["status"] == "ok" + assert second["result"]["remaining"] == [] + + thread.join(timeout=5) + assert not thread.is_alive() + assert json.loads(box["answer"]) == {"answers": {"q0": "alpha", "q1": "beta"}} + + +def test_clarify_batch_answer_update_overwrites_before_completion(server): + thread, box, rid = _drain_batch_block(server, ["q0", "q1"]) + + server.handle_request({ + "id": "a1", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q0", "answer": "first"}, + }) + server.handle_request({ + "id": "a2", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q0", "answer": "changed"}, + }) + server.handle_request({ + "id": "a3", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q1", "answer": "done"}, + }) + + thread.join(timeout=5) + assert json.loads(box["answer"])["answers"]["q0"] == "changed" + + +def test_clarify_batch_empty_answer_is_a_locked_skip(server): + """Skipping one question locks an empty answer — it counts toward + completion instead of leaving the batch waiting.""" + thread, box, rid = _drain_batch_block(server, ["q0", "q1"]) + + server.handle_request({ + "id": "a1", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q0", "answer": ""}, + }) + server.handle_request({ + "id": "a2", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q1", "answer": "kept"}, + }) + + thread.join(timeout=5) + assert json.loads(box["answer"]) == {"answers": {"q0": "", "q1": "kept"}} + + +def test_clarify_batch_unknown_question_id_rejected(server): + thread, box, rid = _drain_batch_block(server, ["q0"]) + + response = server.handle_request({ + "id": "bad", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q9", "answer": "x"}, + }) + assert response["error"]["code"] == 4002 + + server.handle_request({ + "id": "ok", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q0", "answer": "fine"}, + }) + thread.join(timeout=5) + + +def test_clarify_batch_timeout_keeps_locked_answers(capture): + """Locked answers survive the deadline: the tool sees the partials plus + timed_out instead of an empty string.""" + server, buf = capture + thread, box, rid = _drain_batch_block(server, ["q0", "q1"], timeout=1) + + server.handle_request({ + "id": "a1", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q0", "answer": "kept"}, + }) + + thread.join(timeout=10) + assert not thread.is_alive() + result = json.loads(box["answer"]) + assert result == {"answers": {"q0": "kept"}, "timed_out": True} + # The expire notification still fires for the un-finished batch. + messages = [json.loads(line) for line in buf.getvalue().splitlines()] + assert any(m["params"]["type"] == "clarify.expire" for m in messages) + + +def test_clarify_batch_cancel_all_returns_empty(server): + """A respond without question_id cancels the whole batch (Esc path).""" + thread, box, rid = _drain_batch_block(server, ["q0", "q1"]) + + server.handle_request({ + "id": "cancel", "method": "clarify.respond", + "params": {"request_id": rid, "answer": ""}, + }) + + thread.join(timeout=5) + assert box["answer"] == "" + + +def test_clarify_batch_late_question_respond_is_idempotent(server): + response = server.handle_request({ + "id": "late", "method": "clarify.respond", + "params": {"request_id": "gone", "question_id": "q0", "answer": "x"}, + }) + assert response["result"] == {"status": "expired"} + + +def test_clarify_batch_state_cleared_after_resolution(server): + thread, box, rid = _drain_batch_block(server, ["q0"]) + server.handle_request({ + "id": "a", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q0", "answer": "x"}, + }) + thread.join(timeout=5) + with server._prompt_lock: + assert rid not in server._batch_clarify + assert rid not in server._pending + + +def test_clarify_block_helper_builds_batch_payload(capture): + """_clarify_block forwards only wire fields (qid/question/choices/ + multi_select) — the tool-side normalized entries carry extra keys the + renderer must not see.""" + server, buf = capture + normalized = [ + { + "qid": "q0", "id": "approach", "question": "Which?", + "choices": ["a (Recommended)", "b"], "choices_offered": ["a", "b"], + "multi_select": False, + }, + ] + + box = {} + + def run(): + box["answer"] = server._clarify_block("s1", "", None, questions=normalized) + + thread = threading.Thread(target=run, daemon=True) + thread.start() + deadline = time.monotonic() + 2 + rid = None + while time.monotonic() < deadline and rid is None: + with server._prompt_lock: + rid = next(iter(server._batch_clarify), None) + time.sleep(0.01) + assert rid + + server.handle_request({ + "id": "a", "method": "clarify.respond", + "params": {"request_id": rid, "question_id": "q0", "answer": "a"}, + }) + thread.join(timeout=5) + + messages = [json.loads(line) for line in buf.getvalue().splitlines()] + request = messages[0]["params"] + assert request["type"] == "clarify.request" + sent = request["payload"]["questions"][0] + assert set(sent) == {"qid", "question", "choices", "multi_select"} + assert "id" not in sent and "choices_offered" not in sent + + def test_approval_pending_replays_unresolved_requests(server, monkeypatch): from tools import approval diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 70213607c7..d024c7e411 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -145,6 +145,10 @@ _methods: dict[str, callable] = {} _pending: dict[str, tuple[str, threading.Event]] = {} _pending_prompt_payloads: dict[str, tuple[str, dict]] = {} _answers: dict[str, str] = {} +# Batch clarify accumulators: rid → {"qids": [...], "answers": {qid: answer}}. +# Written by clarify.respond (per-question lock, update-in-place), read out by +# _block on resolution/timeout so locked answers survive the deadline. +_batch_clarify: dict[str, dict] = {} _db = None _db_error: str | None = None _stdout_lock = threading.Lock() @@ -1952,7 +1956,14 @@ def _pending_clarify_request_payload(sid: str) -> dict | None: continue event, prompt_payload = _pending_prompt_payloads.get(rid, ("", {})) if event == "clarify.request": - return dict(prompt_payload) + snapshot = dict(prompt_payload) + # Batch clarify: replay the answers locked so far, so a + # reconnecting client restores its per-question ✓ state + # instead of presenting every question as unanswered. + batch = _batch_clarify.get(rid) + if batch is not None and batch["answers"]: + snapshot["answers"] = dict(batch["answers"]) + return snapshot return None @@ -3469,16 +3480,28 @@ def _enable_gateway_prompts() -> None: # ── Blocking prompt factory ────────────────────────────────────────── -def _block(event: str, sid: str, payload: dict, timeout: float | None = 300) -> str: +def _block( + event: str, + sid: str, + payload: dict, + timeout: float | None = 300, + batch_qids: list[str] | None = None, +) -> str: rid = uuid.uuid4().hex[:8] ev = threading.Event() with _prompt_lock: _pending[rid] = (sid, ev) payload["request_id"] = rid _pending_prompt_payloads[rid] = (event, dict(payload)) + if batch_qids: + # Multi-question clarify: per-question answers accumulate here + # (update-in-place until every qid is locked). Locked answers + # survive a timeout — see the batch read-out below. + _batch_clarify[rid] = {"qids": list(batch_qids), "answers": {}} answered = False answer = "" answer_present = False + batch_answers: dict | None = None try: _emit(event, sid, payload) # Natural Event semantics: None → wait forever (clarify configured with @@ -3491,6 +3514,27 @@ def _block(event: str, sid: str, payload: dict, timeout: float | None = 300) -> _pending_prompt_payloads.pop(rid, None) answer_present = rid in _answers answer = _answers.pop(rid, "") + batch_state = _batch_clarify.pop(rid, None) + if batch_state is not None: + batch_answers = dict(batch_state["answers"]) + + if batch_qids is not None: + # Cancel-all (respond with no question_id) resolves via _answers with + # an empty string — that stays a plain cancel, not a partial result. + if answer_present: + return answer + result: dict[str, object] = {"answers": batch_answers or {}} + if not answered: + # Deadline hit: keep whatever was locked, tell the tool the rest + # are absences (not skips), and still fire the expire + # notification so live cards tear down. + result["timed_out"] = True + _emit( + f"{event.removesuffix('.request')}.expire", + sid, + {"request_id": rid}, + ) + return json.dumps(result, ensure_ascii=False) # Emit an `.expire` notification on timeout for every blocking request type # whose `*.respond` handler tolerates a late reply (allow_expired=True). @@ -3529,6 +3573,50 @@ def _clarify_timeout_seconds() -> float | None: return 300 +def _clarify_block(sid: str, q, c, multi_select=False, questions=None) -> str: + """Bridge the clarify tool callback onto _block. + + Single-question calls keep the exact historical payload shape (older + renderers never see a new field). Batch calls emit one clarify.request + carrying the question list — only wire fields (qid/question/choices/ + multi_select) are forwarded; the tool-side normalized entries also carry + result-assembly keys (id, choices_offered) the renderer must not see. + The tool decodes the JSON reply via its batch answer parser. + """ + if questions: + wire = [ + { + "qid": entry["qid"], + "question": entry["question"], + "choices": entry["choices"], + "multi_select": bool(entry["multi_select"]), + } + for entry in questions + ] + return _block( + "clarify.request", + sid, + {"questions": wire}, + timeout=_clarify_timeout_seconds(), + batch_qids=[entry["qid"] for entry in questions], + ) + # multi_select is a pass-through hint: renderers with checkbox + # support can honor it; older renderers ignore the extra field + # and stay single-select (a single answer still parses as a + # one-element list on the tool side). Only emitted when True so + # single-select payloads keep the exact pre-multi-select shape. + return _block( + "clarify.request", + sid, + ( + {"question": q, "choices": c, "multi_select": True} + if multi_select + else {"question": q, "choices": c} + ), + timeout=_clarify_timeout_seconds(), + ) + + def _clear_pending(sid: str | None = None) -> None: """Release pending prompts with an empty answer. @@ -6176,20 +6264,8 @@ def _agent_cbs(sid: str) -> dict: "notice_clear_callback": lambda key: _emit( "notification.clear", sid, {"key": key} ), - "clarify_callback": lambda q, c, multi_select=False: _block( - "clarify.request", - sid, - # multi_select is a pass-through hint: renderers with checkbox - # support can honor it; older renderers ignore the extra field - # and stay single-select (a single answer still parses as a - # one-element list on the tool side). Only emitted when True so - # single-select payloads keep the exact pre-multi-select shape. - ( - {"question": q, "choices": c, "multi_select": True} - if multi_select - else {"question": q, "choices": c} - ), - timeout=_clarify_timeout_seconds(), + "clarify_callback": lambda q, c, multi_select=False, questions=None: ( + _clarify_block(sid, q, c, multi_select=multi_select, questions=questions) ), # read_terminal tool (desktop GUI): same blocking bridge as clarify — the # renderer answers terminal.read.respond with the serialized buffer. @@ -11711,6 +11787,7 @@ def _stage_session_file_attachment( def _respond(rid, params, key, *, allow_expired=False): r = params.get("request_id", "") + question_id = str(params.get("question_id") or "") with _prompt_lock: entry = _pending.get(r) if not entry: @@ -11718,6 +11795,21 @@ def _respond(rid, params, key, *, allow_expired=False): return _ok(rid, {"status": "expired"}) return _err(rid, 4009, f"no pending {key} request") _, ev = entry + batch = _batch_clarify.get(r) + if batch is not None and question_id: + # Per-question lock (multi-question clarify). Update-in-place is + # deliberate: a locked answer stays editable until the batch + # completes, and completion is exactly "every qid locked" — the + # final lock is the Confirm-and-continue click. + if question_id not in batch["qids"]: + return _err(rid, 4002, f"unknown question_id {question_id!r}") + batch["answers"][question_id] = params.get(key, "") + remaining = [ + qid for qid in batch["qids"] if qid not in batch["answers"] + ] + if not remaining: + ev.set() + return _ok(rid, {"status": "ok", "remaining": remaining}) _answers[r] = params.get(key, "") ev.set() return _ok(rid, {"status": "ok"})