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.
This commit is contained in:
can1357
2026-07-27 07:11:18 +02:00
parent 0388946e85
commit a9b7df8821
6 changed files with 199 additions and 12 deletions
+33
View File
@@ -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(
+92 -4
View File
@@ -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://<name>` 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")
+14 -2
View File
@@ -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]
+9 -2
View File
@@ -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,
},
)
@@ -975,5 +975,3 @@ def test_route_non_directive_comment_carries_no_pragmas() -> None:
)
assert decision.directive is False
assert decision.directive_pragmas == ()
+51 -2
View File
@@ -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
# ---------------------------------------------------------------------------