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:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
@@ -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"})
|
||||
|
||||
Reference in New Issue
Block a user