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