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.
This commit is contained in:
ethernet
2026-08-18 16:00:19 -04:00
parent bd8b658a63
commit 879d6a4c78
4 changed files with 314 additions and 16 deletions
+1
View File
@@ -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,
+1
View File
@@ -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(
+204
View File
@@ -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
+108 -16
View File
@@ -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"})