From 1dd2129017c288227453640df3861857ee13bd52 Mon Sep 17 00:00:00 2001 From: GeneralJerel <85066839+GeneralJerel@users.noreply.github.com> Date: Mon, 3 Aug 2026 15:19:43 -0700 Subject: [PATCH] fix(confirm-write): approve one write at a time per conversation A model turn can emit several tool calls at once, and LangGraph runs each as its own task. Two mutating calls in one turn therefore reach copilotkit_interrupt in the same super-step and leave two pending interrupts. Nothing can answer that: the Channel posts one card, and resuming without an interrupt id raises "When there are multiple pending interrupts, you must specify the interrupt id when resuming", which escapes as an ASGI error and kills the run. Reproduced against LangGraph directly -- two parallel tasks that both interrupt leave two pending, and a bare Command(resume=...) fails exactly as it did in production. A second turn arriving while a card is still unanswered does NOT stack, so the accumulation is within one turn only. Let only the first write on a thread pause for approval. The rest come back unrun with a result telling the model to re-issue them, which it does on the next turn, so the writes queue up behind one card at a time instead of racing for the same resume. The claim is keyed by the call itself so it survives the replay that resume performs, and it is released only when an approval resolves -- approved or declined. The gate is shared across interceptors because each MCP server gets its own: two writes to different servers in one turn is precisely the case a per-server gate would miss. --- agent/tests/test_write_confirmation.py | 205 +++++++++++++++++++++++++ agent/write_confirmation.py | 82 +++++++++- 2 files changed, 285 insertions(+), 2 deletions(-) diff --git a/agent/tests/test_write_confirmation.py b/agent/tests/test_write_confirmation.py index edbeaa0..982b21d 100644 --- a/agent/tests/test_write_confirmation.py +++ b/agent/tests/test_write_confirmation.py @@ -47,6 +47,14 @@ def save_project(**args): ) +class _Paused(BaseException): + """Stands in for the GraphInterrupt that suspends an interrupted task. + + A BaseException, like the real one, so `except Exception` handlers in the + interceptor cannot swallow it. + """ + + def test_summarize_args_omits_empty_values(): fields = summarize_args( { @@ -498,6 +506,203 @@ async def failing(_request): assert result is failed +def test_only_one_write_per_thread_may_pause_for_approval(monkeypatch): + """Two writes in one model turn must not raise two interrupts. + + LangGraph runs parallel tool calls as separate tasks, so both would pause; + the resume then fails with "multiple pending interrupts" and the run dies. + """ + cards = approve_and_track(monkeypatch) + capture_reports(monkeypatch) + gate = write_confirmation.ApprovalGate() + interceptor = write_confirmation.WriteConfirmationInterceptor(gate) + written = [] + + paused = [] + + def pause(**kwargs): + # Stand in for the GraphInterrupt that suspends the first task. + paused.append(kwargs["args"]) + raise _Paused + + monkeypatch.setattr(write_confirmation, "copilotkit_interrupt", pause) + + async def handler(request): + written.append(request.name) + return CallToolResult(content=[TextContent(type="text", text="ok")]) + + with pytest.raises(_Paused): + asyncio.run(interceptor(save_project(name="one"), handler)) + + # The second write in the same turn comes back unrun rather than pausing. + second = asyncio.run(interceptor(save_project(name="two"), handler)) + + assert len(paused) == 1 + assert written == [] + assert "another write" in second.content[0].text + assert "re-issue" in second.content[0].text.lower() + assert cards == [] + + +def test_the_same_write_reclaims_its_own_pause_on_resume(monkeypatch): + """Resume replays the task; it must not be mistaken for a second write.""" + capture_reports(monkeypatch) + monkeypatch.setattr(write_confirmation, "_thread_key", lambda: "thread-1") + gate = write_confirmation.ApprovalGate() + interceptor = write_confirmation.WriteConfirmationInterceptor(gate) + + calls = [] + + def pause_then_approve(**kwargs): + calls.append(kwargs["args"]) + if len(calls) == 1: + raise _Paused + return '{"confirmed": true}', {"confirmed": True} + + monkeypatch.setattr( + write_confirmation, "copilotkit_interrupt", pause_then_approve + ) + + async def handler(_request): + return CallToolResult(content=[TextContent(type="text", text="ok")]) + + request = save_project(name="one") + with pytest.raises(_Paused): + asyncio.run(interceptor(request, handler)) + + # Same call, replayed by the resume: it must reach the interrupt again. + result = asyncio.run(interceptor(request, handler)) + + assert len(calls) == 2 + assert result.content[0].text == "ok" + + +def test_a_resolved_approval_frees_the_thread_for_the_next_write(monkeypatch): + approve_and_track(monkeypatch) + capture_reports(monkeypatch) + gate = write_confirmation.ApprovalGate() + interceptor = write_confirmation.WriteConfirmationInterceptor(gate) + + async def handler(_request): + return CallToolResult(content=[TextContent(type="text", text="ok")]) + + first = asyncio.run(interceptor(save_project(name="one"), handler)) + second = asyncio.run(interceptor(save_project(name="two"), handler)) + + # The first approval resolved, so the next write is asked about normally + # rather than being deferred behind a claim nobody is holding. + assert first.content[0].text == "ok" + assert second.content[0].text == "ok" + + +def test_writes_to_different_servers_share_one_gate(monkeypatch): + """Each MCP server gets its own interceptor; the gate must span them.""" + capture_reports(monkeypatch) + monkeypatch.setattr(write_confirmation, "_thread_key", lambda: "thread-1") + gate = write_confirmation.ApprovalGate() + linear = write_confirmation.WriteConfirmationInterceptor(gate) + notion = write_confirmation.WriteConfirmationInterceptor(gate) + + def pause(**_kwargs): + raise _Paused + + monkeypatch.setattr(write_confirmation, "copilotkit_interrupt", pause) + + async def handler(_request): + return CallToolResult(content=[TextContent(type="text", text="ok")]) + + with pytest.raises(_Paused): + asyncio.run(linear(save_project(name="one"), handler)) + + deferred = asyncio.run( + notion( + MCPToolCallRequest( + name="save_document", args={"title": "x"}, server_name="notion" + ), + handler, + ) + ) + + assert "another write" in deferred.content[0].text + + +def test_a_declined_write_frees_the_thread(monkeypatch): + capture_reports(monkeypatch) + monkeypatch.setattr(write_confirmation, "_thread_key", lambda: "thread-1") + gate = write_confirmation.ApprovalGate() + interceptor = write_confirmation.WriteConfirmationInterceptor(gate) + monkeypatch.setattr( + write_confirmation, + "copilotkit_interrupt", + lambda **_k: ('{"confirmed": false}', {"confirmed": False}), + ) + + async def handler(_request): + return CallToolResult(content=[TextContent(type="text", text="ok")]) + + asyncio.run(interceptor(save_project(name="one"), handler)) + + # Cancelling must not wedge the thread against every later write. + assert gate.claim("thread-1", ("linear", "save_issue", "[]")) is True + + +def test_reads_are_never_deferred(monkeypatch): + """A read behind a pending approval still answers; only writes queue.""" + capture_reports(monkeypatch) + monkeypatch.setattr(write_confirmation, "_thread_key", lambda: "thread-1") + gate = write_confirmation.ApprovalGate() + interceptor = write_confirmation.WriteConfirmationInterceptor(gate) + gate.claim("thread-1", ("linear", "save_project", "[]")) + + async def read_issue(issue_id: str): + return issue_id + + interceptor.register_tools( + [ + StructuredTool.from_function( + coroutine=read_issue, + name="get_issue", + description="Read an issue", + metadata={"readOnlyHint": True}, + ) + ] + ) + + async def handler(_request): + return "read-result" + + result = asyncio.run( + interceptor( + MCPToolCallRequest( + name="get_issue", args={"issue_id": "CPK-9"}, server_name="linear" + ), + handler, + ) + ) + + assert result == "read-result" + + +def test_a_thread_without_an_id_still_asks_for_approval(monkeypatch): + """Failing open on the gate is right; failing open on the gate is not.""" + cards = approve_and_track(monkeypatch, thread=None) + capture_reports(monkeypatch) + gate = write_confirmation.ApprovalGate() + interceptor = write_confirmation.WriteConfirmationInterceptor(gate) + written = [] + + async def handler(request): + written.append(request.name) + return CallToolResult(content=[TextContent(type="text", text="ok")]) + + asyncio.run(interceptor(save_project(name="one"), handler)) + + # Serialization is best-effort without a thread id, but the write is still + # gated -- never silently executed. + assert len(cards) == 1 + assert written == ["save_project"] + + def test_write_confirmation_interceptor_rejects_a_malformed_resume( monkeypatch, ): diff --git a/agent/write_confirmation.py b/agent/write_confirmation.py index 7471735..574fa2c 100644 --- a/agent/write_confirmation.py +++ b/agent/write_confirmation.py @@ -145,6 +145,71 @@ def _thread_key() -> str | None: return str(thread_id) if thread_id else None +class ApprovalGate: + """Which thread currently has an approval outstanding. + + A model turn can emit several tool calls at once, and LangGraph runs each + as its own task. Two mutating calls in one turn therefore raise two + interrupts in the same super-step -- and then nothing can answer them: the + Channel posts one card, and LangGraph rejects a resume that doesn't name an + interrupt id ("When there are multiple pending interrupts, you must specify + the interrupt id when resuming"), killing the run. + + So only the first write on a thread may pause for approval. The rest come + back unrun, with a result that tells the model to re-issue them, which it + does on the next turn -- one card at a time, in order. + + Shared across interceptors because each MCP server gets its own: two writes + to different servers in one turn are exactly the case a per-server gate + would miss. + """ + + def __init__(self): + self._claims: dict[str, tuple] = {} + + def claim(self, thread: str | None, token: tuple) -> bool: + """Whether `token` may pause this thread for approval. + + The token identifies the call, so a claim survives the replay that + resume performs: the same call reclaims its own pause, while a + different call in the same turn is turned away. + """ + if thread is None: + # No thread to serialize on. Approving is still safe -- worst case + # is the multi-interrupt error this gate exists to avoid, which is + # better than skipping the gate and writing unapproved. + return True + held = self._claims.get(thread) + if held is not None and held != token: + return False + self._claims[thread] = token + return True + + def release(self, thread: str | None) -> None: + if thread is not None: + self._claims.pop(thread, None) + + +# One gate for the whole process; see ApprovalGate for why it is not per-server. +APPROVAL_GATE = ApprovalGate() + + +def _deferred_result(action: str) -> CallToolResult: + return CallToolResult( + content=[ + TextContent( + type="text", + text=( + f"Not run: another write on this conversation is waiting " + f"for the user's approval. Writes are approved one at a " + f"time, so re-issue this {action.lower()} call on its own " + f"once the current one resolves." + ), + ) + ] + ) + + class WriteConfirmationInterceptor: """Require approval for every MCP tool not marked read-only.""" @@ -154,8 +219,9 @@ class WriteConfirmationInterceptor: "API-query-data-source", } - def __init__(self): + def __init__(self, gate: ApprovalGate | None = None): self._read_only_tools = set(self._KNOWN_READ_ONLY_TOOLS) + self._gate = gate if gate is not None else APPROVAL_GATE # (thread id, tool name) -> (attempts so far, last failure text). self._failures: OrderedDict[tuple[str, str], tuple[int, str]] = ( OrderedDict() @@ -219,14 +285,26 @@ async def __call__( action = action[:1].upper() + action[1:] thread = _thread_key() key = None if thread is None else (thread, request.name) + fields = summarize_args(request.args) + + # Identifies this specific call, so the claim survives resume's replay + # but does not cover a different write issued in the same turn. + token = (request.server_name, request.name, json.dumps(fields, sort_keys=True)) + if not self._gate.claim(thread, token): + return _deferred_result(action) + + # This raises to pause the task, and the claim is deliberately not + # released on the way out: the thread stays claimed for as long as the + # card is unanswered. Only a resolved approval frees it, below. _answer, response = copilotkit_interrupt( action="confirm_write", args={ "action": action, - "fields": summarize_args(request.args), + "fields": fields, **self._retry_args(key), }, ) + self._gate.release(thread) if isinstance(response, str): try: