From a9b7df88218f05a2b2816f26b85dcb1d94e4b448 Mon Sep 17 00:00:00 2001 From: can1357 Date: Mon, 27 Jul 2026 07:11:18 +0200 Subject: [PATCH] feat(python): implemented host tool event normalization and error checks - Added `_normalize_host_tool_event` to `RpcClient` to remap transport tool events to executed host tools. - Updated worker tool-end handling to verify `is_error` is false before marking terminal actions complete. - Added test coverage in `test_client.py` and `test_worker.py` for host tool event normalization and errored review submission behavior. --- python/omp-rpc/src/omp_rpc/client.py | 33 ++++++++ python/omp-rpc/tests/test_client.py | 96 ++++++++++++++++++++++- python/robomp/src/github_client.py | 16 +++- python/robomp/src/worker.py | 11 ++- python/robomp/tests/test_github_events.py | 2 - python/robomp/tests/test_worker.py | 53 ++++++++++++- 6 files changed, 199 insertions(+), 12 deletions(-) diff --git a/python/omp-rpc/src/omp_rpc/client.py b/python/omp-rpc/src/omp_rpc/client.py index beed7ad95..538e6c4ff 100644 --- a/python/omp-rpc/src/omp_rpc/client.py +++ b/python/omp-rpc/src/omp_rpc/client.py @@ -508,6 +508,7 @@ class RpcClient: self._event_condition = threading.Condition() self._pending: dict[str, _PendingRequest] = {} self._pending_host_tool_calls: dict[str, _PendingHostToolCall] = {} + self._host_tool_dispatch_names: dict[str, str] = {} self._pending_host_uri_requests: dict[str, _PendingHostUriRequest] = {} self._request_id = 0 self._events = _BoundedHistory[JsonObject](self._max_event_history) @@ -706,6 +707,7 @@ class RpcClient: # reader's exception path) returns early. self._mark_closed(RpcProcessExitError("RPC process stopped")) self._pending_host_tool_calls.clear() + self._host_tool_dispatch_names.clear() self._pending_host_uri_requests.clear() self._process = None self._pgid = None @@ -1418,6 +1420,30 @@ class RpcClient: return cast(JsonObject, dict(result)) raise RpcError("Host tool handlers must return a string or a result mapping") + def _normalize_host_tool_event(self, payload: JsonObject) -> None: + """Rename transport tool events for in-flight host-tool dispatches. + + With `tools.xdev` enabled, omp mounts custom tools as `xd://` devices + and the agent invokes them through the `write` tool, so + `tool_execution_update`/`tool_execution_end` events report the + transport tool (`write`) rather than the host tool that actually ran. + The `host_tool_call` frame carries the outer call's `toolCallId` (the + device dispatch forwards it verbatim), which lets events for that call + be renamed to the executed host tool — consumers observe the same tool + names regardless of transport. A top-level call (xdev off) maps the + name onto itself. `tool_execution_start` precedes the `host_tool_call` + frame on the wire, so start events keep the transport name. + """ + tool_call_id = payload.get("toolCallId") + if not isinstance(tool_call_id, str): + return + if payload.get("type") == "tool_execution_end": + tool_name = self._host_tool_dispatch_names.pop(tool_call_id, None) + else: + tool_name = self._host_tool_dispatch_names.get(tool_call_id) + if tool_name is not None: + payload["toolName"] = tool_name + def _handle_host_tool_call(self, payload: JsonObject) -> None: request_id = payload.get("id") tool_name = payload.get("toolName") @@ -1429,6 +1455,10 @@ class RpcClient: or not isinstance(tool_call_id, str) ): return + # Remember the dispatch so tool_execution_* events for this call id can + # be renamed from the transport tool to the host tool that ran; see + # _normalize_host_tool_event. + self._host_tool_dispatch_names[tool_call_id] = tool_name if not isinstance(raw_arguments, Mapping): self._send_notification( { @@ -1860,6 +1890,9 @@ class RpcClient: self._handle_host_uri_cancel(payload) continue + payload_type = payload.get("type") + if payload_type in ("tool_execution_update", "tool_execution_end"): + self._normalize_host_tool_event(payload) notification = parse_notification(payload) listener_notification = parse_notification(payload) self._dispatch_listeners( diff --git a/python/omp-rpc/tests/test_client.py b/python/omp-rpc/tests/test_client.py index fb803b66f..c9129ac8f 100644 --- a/python/omp-rpc/tests/test_client.py +++ b/python/omp-rpc/tests/test_client.py @@ -70,6 +70,8 @@ FAKE_SERVER = textwrap.dedent( } registered_host_tools = [] + host_event_tool_call_id = "toolu_host_1" + host_event_tool_name = "echo_host" def current_state(): return { @@ -391,6 +393,8 @@ FAKE_SERVER = textwrap.dedent( continue if message == "needs host tool": print(json.dumps({"type": "agent_start"}), flush=True) + host_event_tool_call_id = "toolu_host_1" + host_event_tool_name = "echo_host" print( json.dumps( { @@ -404,6 +408,34 @@ FAKE_SERVER = textwrap.dedent( flush=True, ) continue + if message == "needs xd host tool": + print(json.dumps({"type": "agent_start"}), flush=True) + host_event_tool_call_id = "toolu_write_1" + host_event_tool_name = "write" + print( + json.dumps( + { + "type": "tool_execution_start", + "toolCallId": "toolu_write_1", + "toolName": "write", + "args": {"path": "xd://echo_host", "content": '{"message": "hello"}'}, + } + ), + flush=True, + ) + print( + json.dumps( + { + "type": "host_tool_call", + "id": "host-call-2", + "toolCallId": "toolu_write_1", + "toolName": "echo_host", + "arguments": {"message": "hello"}, + } + ), + flush=True, + ) + continue if message == "notifications": print(json.dumps({"type": "extension_error", "extensionPath": "/tmp/ext.py", "event": "run", "error": "boom"}), flush=True) print(json.dumps({"type": "unknown_future_event", "value": 1}), flush=True) @@ -418,8 +450,8 @@ FAKE_SERVER = textwrap.dedent( json.dumps( { "type": "tool_execution_update", - "toolCallId": "toolu_host_1", - "toolName": "echo_host", + "toolCallId": host_event_tool_call_id, + "toolName": host_event_tool_name, "args": {"message": "hello"}, "partialResult": command["partialResult"], } @@ -431,8 +463,8 @@ FAKE_SERVER = textwrap.dedent( json.dumps( { "type": "tool_execution_end", - "toolCallId": "toolu_host_1", - "toolName": "echo_host", + "toolCallId": host_event_tool_call_id, + "toolName": host_event_tool_name, "result": command["result"], "isError": command.get("isError", False), } @@ -908,6 +940,62 @@ class RpcClientTests(unittest.TestCase): self.assertEqual(len(end_events), 1) self.assertEqual(end_events[0].result["content"][0]["text"], "host:hello") + def test_xd_dispatched_custom_tool_events_carry_host_tool_name(self) -> None: + """Events for an xd:// device dispatch are renamed to the executed host tool. + + With `tools.xdev` on, omp invokes a custom tool through `write + xd://` and the wire events carry the transport tool (`write`). + Consumers must observe the host-tool name on update/end events + regardless of transport — roboomp's terminal-action detection + triple-posted PR reviews when end events only said `write` + (oh-my-pi#6696). `tool_execution_start` precedes the `host_tool_call` + frame on the wire and keeps the transport name. + """ + + def echo_host(args: dict[str, str], context) -> str: + context.send_update(f"working:{args['message']}") + return f"host:{args['message']}" + + with self.make_client( + custom_tools=( + host_tool( + name="echo_host", + description="Echo from the Python host process", + parameters={ + "type": "object", + "properties": {"message": {"type": "string"}}, + "required": ["message"], + "additionalProperties": False, + }, + execute=echo_host, + ), + ) + ) as client: + turn = client.prompt_and_wait("needs xd host tool", timeout=2.0) + start_names = [ + event.tool_name + for event in turn.events + if getattr(event, "type", None) == "tool_execution_start" + ] + update_events = [ + event + for event in turn.events + if getattr(event, "type", None) == "tool_execution_update" + ] + end_events = [ + event + for event in turn.events + if getattr(event, "type", None) == "tool_execution_end" + ] + + self.assertEqual(start_names, ["write"]) + self.assertEqual( + [event.tool_name for event in update_events], ["echo_host"] + ) + self.assertEqual([event.tool_name for event in end_events], ["echo_host"]) + self.assertEqual(end_events[0].tool_call_id, "toolu_write_1") + self.assertEqual(end_events[0].result["content"][0]["text"], "host:hello") + def test_extension_ui_round_trip(self) -> None: with self.make_client() as client: client.prompt("needs ui") diff --git a/python/robomp/src/github_client.py b/python/robomp/src/github_client.py index b03b73297..2e9c76646 100644 --- a/python/robomp/src/github_client.py +++ b/python/robomp/src/github_client.py @@ -265,7 +265,13 @@ class GitHubClient: last_exc = exc log.warning( "transient github 5xx, retrying", - extra={"method": method, "path": path, "attempt": attempt + 1, "delay": delay, "status": exc.status}, + extra={ + "method": method, + "path": path, + "attempt": attempt + 1, + "delay": delay, + "status": exc.status, + }, ) time.sleep(delay) raise last_exc # type: ignore[misc] @@ -294,7 +300,13 @@ class GitHubClient: last_exc = exc log.warning( "transient github 5xx, retrying", - extra={"method": method, "path": path, "attempt": attempt + 1, "delay": delay, "status": exc.status}, + extra={ + "method": method, + "path": path, + "attempt": attempt + 1, + "delay": delay, + "status": exc.status, + }, ) await asyncio.sleep(delay) raise last_exc # type: ignore[misc] diff --git a/python/robomp/src/worker.py b/python/robomp/src/worker.py index dddc67abd..79e731eb8 100644 --- a/python/robomp/src/worker.py +++ b/python/robomp/src/worker.py @@ -500,14 +500,21 @@ def _run_rpc_blocking( def _on_tool_end(event: ToolExecutionEndEvent) -> None: tool_name = event.tool_name - if event.result is not None: + # `tool_name` is transport-normalized by omp_rpc: an xd:// device + # dispatch (`write xd://submit_pr_review`) reports the host tool that + # ran, so terminal-action detection can match on host-tool names. A + # failed execution (`is_error`) does not count as reaching the + # terminal action — a rejected submit must still trigger the + # completion reminder. + ok = event.result is not None and not event.is_error + if ok: tools_called.add(tool_name) log.info( "tool_end", extra={ "issue": bindings.issue_key, "tool": tool_name, - "ok": event.result is not None, + "ok": ok, }, ) diff --git a/python/robomp/tests/test_github_events.py b/python/robomp/tests/test_github_events.py index 35aa71dbf..6bf8c6b3a 100644 --- a/python/robomp/tests/test_github_events.py +++ b/python/robomp/tests/test_github_events.py @@ -975,5 +975,3 @@ def test_route_non_directive_comment_carries_no_pragmas() -> None: ) assert decision.directive is False assert decision.directive_pragmas == () - - diff --git a/python/robomp/tests/test_worker.py b/python/robomp/tests/test_worker.py index 3094a5a7d..2d0fb3b94 100644 --- a/python/robomp/tests/test_worker.py +++ b/python/robomp/tests/test_worker.py @@ -650,7 +650,7 @@ async def test_run_rpc_stops_reminding_after_terminal_tool(tmp_path: Path, setti # driver registers the callback before prompt_and_wait, so we # replay it here. for cb in client._tool_end_callbacks: - cb(SimpleNamespace(tool_name="gh_open_pr", result={})) + cb(SimpleNamespace(tool_name="gh_open_pr", result={}, is_error=None)) # Capture the registered tool_end callback on the fake. original_on_tool_end = _FakeRpcClient.on_tool_execution_end @@ -759,7 +759,7 @@ async def test_run_rpc_review_pr_stops_after_submit_without_dirty_probe( def _on_prompt(client: _FakeRpcClient, _prompt: str) -> None: for cb in client._tool_end_callbacks: - cb(SimpleNamespace(tool_name="submit_pr_review", result={})) + cb(SimpleNamespace(tool_name="submit_pr_review", result={}, is_error=None)) _FakeRpcClient.on_tool_execution_end = _record_tool_end # type: ignore[assignment] try: @@ -783,6 +783,55 @@ async def test_run_rpc_review_pr_stops_after_submit_without_dirty_probe( assert fake.prompts == ["kickoff"] +@pytest.mark.asyncio +async def test_run_rpc_review_pr_still_reminds_when_submit_fails(tmp_path: Path, settings: Settings) -> None: + """An errored terminal-tool end event does not count as the terminal action. + + omp_rpc normalizes xd:// device dispatches to the host-tool name, so a + rejected `submit_pr_review` surfaces as an end event with `is_error=True` + under its real name. Counting it would end the review task silently with + no review submitted — the completion reminder must still fire. + """ + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False) + original_on_tool_end = _FakeRpcClient.on_tool_execution_end + + def _record_tool_end(self, cb) -> None: + self._tool_end_callbacks = getattr(self, "_tool_end_callbacks", []) + self._tool_end_callbacks.append(cb) + + def _on_prompt(client: _FakeRpcClient, _prompt: str) -> None: + for cb in client._tool_end_callbacks: + cb( + SimpleNamespace( + tool_name="submit_pr_review", + result={"content": [{"type": "text", "text": "GitHub rejected PR review: 422"}]}, + is_error=True, + ) + ) + + _FakeRpcClient.on_tool_execution_end = _record_tool_end # type: ignore[assignment] + try: + _FakeRpcClient.on_prompt = staticmethod(_on_prompt) # type: ignore[attr-defined] + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="review_pr", + prompt="kickoff", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + finally: + _FakeRpcClient.on_tool_execution_end = original_on_tool_end # type: ignore[assignment] + delattr(_FakeRpcClient, "on_prompt") + + fake = _FakeRpcClient.instances[0] + assert len(fake.prompts) == 1 + settings.task_completion_max_reminders + assert all("submit_pr_review" in p for p in fake.prompts[1:]) + + # --------------------------------------------------------------------------- # Dirty-state watchdog # ---------------------------------------------------------------------------