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