feat: added host-uri frame and rpc bridge for read/write/cancel routing

- Added `set_host_uri_schemes` and host-uri frame/type definitions; documented read/write/cancel behavior.
- Added `RpcHostUriBridge` in rpc mode to register schemes, dispatch read/write/cancel ops, and clear pending requests.
- Added internal URL write support with lowercased scheme matching, handler routing, and hashline-prefixed success output.
- Added Python host-uri APIs/exports, cancellable client request handling, and host-uri read/write test coverage.
This commit is contained in:
can1357
2026-05-15 00:54:09 +02:00
parent 9bbc7465ba
commit bddf9989b5
14 changed files with 1029 additions and 17 deletions
+41
View File
@@ -142,6 +142,47 @@ If you want runtime conversion into a richer Python type, pass `decode=` to
`host_tool(...)`. That lets you keep the JSON Schema contract on the wire while
parsing the incoming argument object into a dataclass or model in the handler.
## Host-Owned URI Schemes
Hosts can also expose custom URL schemes that behave like virtual files.
Registered schemes are routed through the agent's `read` (and `write`) tools
over the same RPC transport — handlers do the actual I/O on the Python side:
```python
from omp_rpc import RpcClient, host_uri
rows: dict[str, str] = {"42": "id=42\nname=Alice\n"}
def read_row(url: str, _ctx) -> str:
row_id = url.removeprefix("db://users/")
return rows[row_id]
def write_row(url: str, content: str, _ctx) -> None:
row_id = url.removeprefix("db://users/")
rows[row_id] = content
with RpcClient(
no_session=True,
host_uris=(
host_uri(
scheme="db",
description="Virtual db row files",
read=read_row,
write=write_row,
),
),
) as client:
client.prompt_and_wait("Read db://users/42 and rewrite it with name=Bob")
```
Schemes registered as read-only (no `write=`) reject `write` calls with a
clear error. The agent's `edit` tool does not target host URIs — hosts that
want mutation expose `write` and the model uses the `write` tool with the
full replacement content.
## Extension UI Requests
Extensions in RPC mode can ask the host for input. Those requests are available as
+18
View File
@@ -17,6 +17,16 @@ from .client import (
UiRequestListener,
)
from .host_tools import HostTool, HostToolContext, HostToolResultPayload, HostToolResultValue, host_tool
from .host_uris import (
HostUri,
HostUriContentType,
HostUriContext,
HostUriReadHandler,
HostUriReadResult,
HostUriReadValue,
HostUriWriteHandler,
host_uri,
)
from .protocol import (
AgentEndEvent,
AgentMessage,
@@ -106,6 +116,13 @@ __all__ = [
"HostToolContext",
"HostToolResultPayload",
"HostToolResultValue",
"HostUri",
"HostUriContentType",
"HostUriContext",
"HostUriReadHandler",
"HostUriReadResult",
"HostUriReadValue",
"HostUriWriteHandler",
"HookMessage",
"ImageContent",
"ListenerErrorEvent",
@@ -162,4 +179,5 @@ __all__ = [
"parse_session_state",
"parse_todo_phases",
"host_tool",
"host_uri",
]
+123
View File
@@ -11,6 +11,7 @@ from pathlib import Path
from typing import Any, Callable, Generic, Mapping, Sequence, TypeVar, cast
from .host_tools import HostTool, HostToolContext
from .host_uris import HostUri, HostUriContext, normalize_read_result
from .protocol import (
AgentStartEvent,
AgentEndEvent,
@@ -215,6 +216,11 @@ class _PendingHostToolCall:
cancel_event: threading.Event
@dataclass(slots=True)
class _PendingHostUriRequest:
cancel_event: threading.Event
@dataclass(slots=True)
class _BoundedHistory(Generic[THistoryItem]):
limit: int | None
@@ -277,6 +283,7 @@ class RpcClient:
provider_session_id: str | None = None,
tools: Sequence[str] | None = None,
custom_tools: Sequence[HostTool[Any, Any]] | None = None,
host_uris: Sequence[HostUri[Any]] | None = None,
no_session: bool = False,
no_skills: bool = False,
no_rules: bool = False,
@@ -300,6 +307,7 @@ class RpcClient:
self._provider_session_id = provider_session_id
self._tools = tuple(tools) if tools is not None else None
self._custom_tools = tuple(custom_tools) if custom_tools is not None else ()
self._host_uris = tuple(host_uris) if host_uris is not None else ()
self._no_session = no_session
self._no_skills = no_skills
self._no_rules = no_rules
@@ -320,6 +328,7 @@ class RpcClient:
self._event_condition = threading.Condition()
self._pending: dict[str, _PendingRequest] = {}
self._pending_host_tool_calls: dict[str, _PendingHostToolCall] = {}
self._pending_host_uri_requests: dict[str, _PendingHostUriRequest] = {}
self._request_id = 0
self._events = _BoundedHistory[JsonObject](self._max_event_history)
self._async_errors = _BoundedHistory[BaseException](_DEFAULT_ERROR_HISTORY_LIMIT)
@@ -426,6 +435,8 @@ class RpcClient:
if self._custom_tools:
self.set_custom_tools(self._custom_tools)
if self._host_uris:
self.set_host_uris(self._host_uris)
return self
def stop(self) -> None:
@@ -436,6 +447,8 @@ class RpcClient:
self._stopping = True
for pending_call in self._pending_host_tool_calls.values():
pending_call.cancel_event.set()
for pending_uri in self._pending_host_uri_requests.values():
pending_uri.cancel_event.set()
try:
if process.stdin is not None:
@@ -464,6 +477,7 @@ class RpcClient:
pass
self._fail_pending(RpcProcessExitError("RPC process stopped"))
self._pending_host_tool_calls.clear()
self._pending_host_uri_requests.clear()
self._process = None
self._ready.set()
with self._event_condition:
@@ -760,6 +774,27 @@ class RpcClient:
raise RpcError("set_host_tools response did not include toolNames")
return tuple(str(name) for name in tool_names)
def set_host_uris(self, host_uris: Sequence[HostUri[Any]]) -> tuple[str, ...]:
self._host_uris = tuple(host_uris)
if self._process is None:
return tuple(uri.scheme for uri in self._host_uris)
schemes_payload: list[JsonObject] = []
for uri in self._host_uris:
entry: JsonObject = {"scheme": uri.scheme, "writable": uri.writable, "immutable": uri.immutable}
if uri.description is not None:
entry["description"] = uri.description
schemes_payload.append(entry)
payload = self._request(
"set_host_uri_schemes",
schemes=cast(JsonValue, schemes_payload),
)
schemes = payload.get("schemes") or []
if not isinstance(schemes, list):
raise RpcError("set_host_uri_schemes response did not include schemes")
return tuple(str(entry) for entry in schemes)
def prompt(
self,
message: str,
@@ -1054,6 +1089,88 @@ class RpcClient:
if pending_call is not None:
pending_call.cancel_event.set()
def _send_host_uri_error(self, request_id: str, message: str) -> None:
self._send_notification(
{
"type": "host_uri_result",
"id": request_id,
"error": message,
"isError": True,
}
)
def _handle_host_uri_request(self, payload: JsonObject) -> None:
request_id = payload.get("id")
operation = payload.get("operation")
url = payload.get("url")
if not isinstance(request_id, str) or not isinstance(operation, str) or not isinstance(url, str):
return
if operation not in ("read", "write"):
self._send_host_uri_error(request_id, f"Unsupported host URI operation: {operation}")
return
try:
from urllib.parse import urlparse
parsed = urlparse(url)
except ValueError:
self._send_host_uri_error(request_id, f"Could not parse host URI: {url}")
return
scheme = (parsed.scheme or "").lower()
uri = next((candidate for candidate in self._host_uris if candidate.scheme == scheme), None)
if uri is None:
self._send_host_uri_error(request_id, f'Host URI scheme "{scheme}://" is not registered')
return
if operation == "write" and uri.write is None:
self._send_host_uri_error(
request_id, f'Host URI scheme "{scheme}://" was not registered with a write handler'
)
return
pending = _PendingHostUriRequest(cancel_event=threading.Event())
self._pending_host_uri_requests[request_id] = pending
def run() -> None:
try:
context = HostUriContext(url=url, operation=cast(Any, operation), _cancel_event=pending.cancel_event)
if operation == "read":
value = uri.read(url, context)
if pending.cancel_event.is_set():
return
result_fields = normalize_read_result(value)
self._send_notification(
{
"type": "host_uri_result",
"id": request_id,
**result_fields,
}
)
else:
raw_content = payload.get("content")
content = str(raw_content) if raw_content is not None else ""
assert uri.write is not None
uri.write(url, content, context)
if pending.cancel_event.is_set():
return
self._send_notification({"type": "host_uri_result", "id": request_id})
except Exception as exc:
if pending.cancel_event.is_set():
return
self._send_host_uri_error(request_id, str(exc))
finally:
self._pending_host_uri_requests.pop(request_id, None)
threading.Thread(target=run, name=f"omp-rpc-host-uri:{scheme}:{operation}", daemon=True).start()
def _handle_host_uri_cancel(self, payload: JsonObject) -> None:
target_id = payload.get("targetId")
if not isinstance(target_id, str):
return
pending = self._pending_host_uri_requests.get(target_id)
if pending is not None:
pending.cancel_event.set()
def _add_typed_event_listener(self, event_type: str, listener: TEventListener) -> Callable[[], None]:
listeners = self._typed_event_listeners.setdefault(event_type, [])
typed_listener = cast(AgentEventListener, listener)
@@ -1232,6 +1349,12 @@ class RpcClient:
if payload.get("type") == "host_tool_cancel":
self._handle_host_tool_cancel(payload)
continue
if payload.get("type") == "host_uri_request":
self._handle_host_uri_request(payload)
continue
if payload.get("type") == "host_uri_cancel":
self._handle_host_uri_cancel(payload)
continue
notification = parse_notification(payload)
listener_notification = parse_notification(payload)
+122
View File
@@ -0,0 +1,122 @@
from __future__ import annotations
import threading
from dataclasses import dataclass, field
from typing import Callable, Generic, Literal, TypeAlias, TypeVar, TypedDict
from .protocol import JsonObject
TPayload = TypeVar("TPayload")
HostUriContentType: TypeAlias = Literal["text/markdown", "application/json", "text/plain"]
class HostUriReadResult(TypedDict, total=False):
"""Structured response a `read` handler may return.
Plain strings are also accepted; they are normalized to `{"content": <str>}`.
"""
content: str
content_type: HostUriContentType
notes: list[str]
immutable: bool
HostUriReadValue: TypeAlias = HostUriReadResult | str
@dataclass(slots=True)
class HostUriContext:
"""Per-request context passed to host URI handlers.
Mirrors the cancellation hooks `HostToolContext` exposes for parity, so
handlers can poll for cancellation when serving long-running reads/writes.
"""
url: str
operation: Literal["read", "write"]
_cancel_event: threading.Event = field(default_factory=threading.Event)
@property
def cancelled(self) -> bool:
return self._cancel_event.is_set()
HostUriReadHandler: TypeAlias = Callable[[str, HostUriContext], HostUriReadValue]
HostUriWriteHandler: TypeAlias = Callable[[str, str, HostUriContext], None]
@dataclass(slots=True, frozen=True)
class HostUri(Generic[TPayload]):
"""Definition of a custom URI scheme served by the Python host.
Hosts register a `HostUri` per scheme. The bridge dispatches `<scheme>://`
URLs the agent reads (and, when `write` is provided, writes) to the
matching callbacks. The agent's `edit` tool is not supported for virtual
URIs — hosts that want to mutate virtual files expose a `write` handler
and let the model use the `write` tool with the full replacement content.
"""
scheme: str
read: HostUriReadHandler
write: HostUriWriteHandler | None = None
description: str | None = None
immutable: bool = False
@property
def writable(self) -> bool:
return self.write is not None
def host_uri(
*,
scheme: str,
read: HostUriReadHandler,
write: HostUriWriteHandler | None = None,
description: str | None = None,
immutable: bool = False,
) -> HostUri[None]:
cleaned = (scheme or "").strip().lower()
if not cleaned:
raise ValueError("scheme must be a non-empty string")
return HostUri(
scheme=cleaned,
read=read,
write=write,
description=description,
immutable=immutable,
)
def normalize_read_result(value: HostUriReadValue) -> JsonObject:
"""Convert a handler's `read` return into the wire-frame fields.
Returns a dict suitable for spreading into a `host_uri_result` payload.
"""
if isinstance(value, str):
return {"content": value}
if not isinstance(value, dict):
raise TypeError("Host URI read handlers must return a string or a HostUriReadResult mapping")
payload: JsonObject = {}
if "content" not in value:
raise ValueError("HostUriReadResult requires a 'content' field")
payload["content"] = str(value["content"])
content_type = value.get("content_type")
if content_type is not None:
if content_type not in ("text/markdown", "application/json", "text/plain"):
raise ValueError(f"Unsupported content_type: {content_type!r}")
payload["contentType"] = content_type
notes = value.get("notes")
if notes is not None:
payload["notes"] = [str(item) for item in notes]
if "immutable" in value:
payload["immutable"] = bool(value["immutable"])
return payload
+235
View File
@@ -0,0 +1,235 @@
from __future__ import annotations
import sys
import textwrap
import threading
import time
import unittest
from omp_rpc import RpcClient, host_uri
from omp_rpc.host_uris import HostUri, normalize_read_result
URI_SERVER = textwrap.dedent(
"""
import json
import sys
print(json.dumps({"type": "ready"}), flush=True)
pending_uri_id = 1
def respond(request_id, command, data=None, success=True, error=None):
frame = {"id": request_id, "type": "response", "command": command, "success": success}
if success:
if data is not None:
frame["data"] = data
else:
frame["error"] = error or "error"
print(json.dumps(frame), flush=True)
for raw_line in sys.stdin:
raw_line = raw_line.strip()
if not raw_line:
continue
command = json.loads(raw_line)
command_type = command.get("type")
request_id = command.get("id")
if command_type == "set_host_uri_schemes":
schemes = command.get("schemes", [])
respond(
request_id,
"set_host_uri_schemes",
{"schemes": [entry.get("scheme", "") for entry in schemes]},
)
elif command_type == "trigger_read":
print(
json.dumps(
{
"type": "host_uri_request",
"id": f"uri-req-{pending_uri_id}",
"operation": "read",
"url": command["url"],
}
),
flush=True,
)
pending_uri_id += 1
respond(request_id, "trigger_read", {})
elif command_type == "trigger_write":
print(
json.dumps(
{
"type": "host_uri_request",
"id": f"uri-req-{pending_uri_id}",
"operation": "write",
"url": command["url"],
"content": command["content"],
}
),
flush=True,
)
pending_uri_id += 1
respond(request_id, "trigger_write", {})
elif command_type == "host_uri_result":
# Echo back as response so the test can assert on the wire frame
print(
json.dumps(
{
"type": "response",
"command": "uri_echo",
"success": True,
"data": {"frame": command},
}
),
flush=True,
)
else:
respond(request_id, command_type, success=False, error=f"unsupported: {command_type}")
"""
)
class HostUriHelperTests(unittest.TestCase):
def test_normalize_read_result_accepts_string(self) -> None:
self.assertEqual(normalize_read_result("hello"), {"content": "hello"})
def test_normalize_read_result_accepts_full_mapping(self) -> None:
result = normalize_read_result(
{
"content": "body",
"content_type": "application/json",
"notes": ["fresh"],
"immutable": True,
}
)
self.assertEqual(result["content"], "body")
self.assertEqual(result["contentType"], "application/json")
self.assertEqual(result["notes"], ["fresh"])
self.assertTrue(result["immutable"])
def test_normalize_read_result_requires_content(self) -> None:
with self.assertRaises(ValueError):
normalize_read_result({"content_type": "text/plain"}) # type: ignore[arg-type]
def test_normalize_read_result_rejects_invalid_content_type(self) -> None:
with self.assertRaises(ValueError):
normalize_read_result({"content": "x", "content_type": "application/octet-stream"}) # type: ignore[arg-type]
def test_host_uri_helper_normalizes_scheme(self) -> None:
uri = host_uri(scheme=" DB ", read=lambda url, ctx: "x")
self.assertEqual(uri.scheme, "db")
self.assertFalse(uri.writable)
with self.assertRaises(ValueError):
host_uri(scheme="", read=lambda url, ctx: "x")
def test_host_uri_writable_when_write_supplied(self) -> None:
uri = host_uri(scheme="db", read=lambda url, ctx: "x", write=lambda url, content, ctx: None)
self.assertTrue(uri.writable)
class RpcHostUriBridgeTests(unittest.TestCase):
def _make_client(self, **kwargs: object) -> RpcClient:
return RpcClient(
command=[sys.executable, "-u", "-c", URI_SERVER],
startup_timeout=2.0,
request_timeout=2.0,
**kwargs,
)
def test_set_host_uris_registers_schemes_on_start(self) -> None:
captured: list[tuple[str, str]] = []
def read_db(url: str, _ctx) -> str:
captured.append(("read", url))
return "id=42"
with self._make_client(
host_uris=(host_uri(scheme="db", read=read_db, description="test rows"),),
) as client:
# No public list — we exercise the on-start side effect by hitting the wire.
payload = client._request("trigger_read", url="db://users/42") # type: ignore[attr-defined]
self.assertEqual(payload, {})
frame = self._await_echo(client)
self.assertEqual(frame["type"], "host_uri_result")
self.assertEqual(frame["content"], "id=42")
self.assertEqual(captured, [("read", "db://users/42")])
def test_read_handler_can_return_structured_result(self) -> None:
def read_db(_url: str, _ctx):
return {
"content": '{"name":"Alice"}',
"content_type": "application/json",
"notes": ["row fresh"],
"immutable": True,
}
with self._make_client(host_uris=(host_uri(scheme="db", read=read_db),)) as client:
client._request("trigger_read", url="db://users/42") # type: ignore[attr-defined]
frame = self._await_echo(client)
self.assertEqual(frame["content"], '{"name":"Alice"}')
self.assertEqual(frame["contentType"], "application/json")
self.assertEqual(frame["notes"], ["row fresh"])
self.assertTrue(frame["immutable"])
def test_write_handler_receives_content_and_succeeds(self) -> None:
seen: dict[str, str] = {}
def write_db(url: str, content: str, _ctx) -> None:
seen[url] = content
uri = host_uri(scheme="db", read=lambda url, ctx: "ignored", write=write_db)
with self._make_client(host_uris=(uri,)) as client:
client._request("trigger_write", url="db://users/42", content="name=Bob") # type: ignore[attr-defined]
frame = self._await_echo(client)
self.assertEqual(frame["type"], "host_uri_result")
self.assertNotIn("isError", frame)
self.assertEqual(seen, {"db://users/42": "name=Bob"})
def test_write_rejected_for_read_only_scheme(self) -> None:
with self._make_client(
host_uris=(host_uri(scheme="db", read=lambda url, ctx: "x"),),
) as client:
client._request("trigger_write", url="db://users/42", content="ignored") # type: ignore[attr-defined]
frame = self._await_echo(client)
self.assertTrue(frame.get("isError"))
self.assertIn("write handler", frame["error"])
def test_unknown_scheme_is_rejected_with_error(self) -> None:
with self._make_client(
host_uris=(host_uri(scheme="db", read=lambda url, ctx: "x"),),
) as client:
client._request("trigger_read", url="other://stuff") # type: ignore[attr-defined]
frame = self._await_echo(client)
self.assertTrue(frame.get("isError"))
self.assertIn("not registered", frame["error"])
def test_handler_exception_is_surfaced_as_error(self) -> None:
def read_db(_url: str, _ctx) -> str:
raise RuntimeError("boom")
with self._make_client(host_uris=(host_uri(scheme="db", read=read_db),)) as client:
client._request("trigger_read", url="db://users/42") # type: ignore[attr-defined]
frame = self._await_echo(client)
self.assertTrue(frame.get("isError"))
self.assertEqual(frame["error"], "boom")
def _await_echo(self, client: RpcClient) -> dict:
# The fake server echoes the host_uri_result frame back as an
# `uri_echo` response. We poll the events history to surface it.
deadline = time.time() + 2.0
while time.time() < deadline:
with client._state_lock: # type: ignore[attr-defined]
events = client._events.snapshot() # type: ignore[attr-defined]
for event in events:
if event.get("command") == "uri_echo" and event.get("data"):
return event["data"]["frame"]
time.sleep(0.02)
self.fail("Timed out waiting for host_uri_result echo")
if __name__ == "__main__":
unittest.main()