chore: reformat
This commit is contained in:
@@ -576,7 +576,12 @@ BROKEN_STARTUP_SERVER = textwrap.dedent(
|
||||
|
||||
class RpcClientTests(unittest.TestCase):
|
||||
def make_client(self, server: str = FAKE_SERVER, **kwargs: object) -> RpcClient:
|
||||
return RpcClient(command=[sys.executable, "-u", "-c", server], startup_timeout=2.0, request_timeout=2.0, **kwargs)
|
||||
return RpcClient(
|
||||
command=[sys.executable, "-u", "-c", server],
|
||||
startup_timeout=2.0,
|
||||
request_timeout=2.0,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def test_command_builder_supports_common_rpc_options(self) -> None:
|
||||
client = RpcClient(
|
||||
@@ -622,7 +627,9 @@ class RpcClientTests(unittest.TestCase):
|
||||
with self.make_client() as client:
|
||||
state = client.get_state()
|
||||
self.assertEqual(state.session_id, "fake-session")
|
||||
self.assertEqual(state.model.id if state.model else None, "claude-sonnet-4-5")
|
||||
self.assertEqual(
|
||||
state.model.id if state.model else None, "claude-sonnet-4-5"
|
||||
)
|
||||
|
||||
result = client.bash("echo hello")
|
||||
self.assertEqual(result.output, "hello\n")
|
||||
@@ -658,11 +665,21 @@ class RpcClientTests(unittest.TestCase):
|
||||
self.assertEqual(state.dump_tools[-1].name, "echo_host")
|
||||
|
||||
turn = client.prompt_and_wait("needs host tool", timeout=2.0)
|
||||
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"]
|
||||
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(len(update_events), 1)
|
||||
self.assertEqual(update_events[0].partial_result["content"][0]["text"], "working:hello")
|
||||
self.assertEqual(
|
||||
update_events[0].partial_result["content"][0]["text"], "working:hello"
|
||||
)
|
||||
self.assertEqual(len(end_events), 1)
|
||||
self.assertEqual(end_events[0].result["content"][0]["text"], "host:hello")
|
||||
|
||||
@@ -679,7 +696,9 @@ class RpcClientTests(unittest.TestCase):
|
||||
seen_methods: list[str] = []
|
||||
|
||||
with self.make_client() as client:
|
||||
client.install_headless_ui(on_request=lambda request: seen_methods.append(request.method))
|
||||
client.install_headless_ui(
|
||||
on_request=lambda request: seen_methods.append(request.method)
|
||||
)
|
||||
client.prompt_and_wait("needs ui", timeout=2.0)
|
||||
|
||||
self.assertEqual(seen_methods, ["input"])
|
||||
@@ -690,7 +709,9 @@ class RpcClientTests(unittest.TestCase):
|
||||
notification_types: list[str] = []
|
||||
client = self.make_client()
|
||||
client.on_ready(lambda event: ready_types.append(event.type))
|
||||
client.on_notification(lambda notification: notification_types.append(notification.type))
|
||||
client.on_notification(
|
||||
lambda notification: notification_types.append(notification.type)
|
||||
)
|
||||
client.on_turn_start(lambda event: event_types.append(event.type))
|
||||
client.on_message_update(lambda event: event_types.append(event.type))
|
||||
client.on_agent_end(lambda event: event_types.append(event.type))
|
||||
@@ -729,7 +750,10 @@ class RpcClientTests(unittest.TestCase):
|
||||
self.assertEqual(cycled.model.id, "claude-sonnet-4-5")
|
||||
|
||||
available = client.get_available_models()
|
||||
self.assertEqual([item.id for item in available], ["claude-sonnet-4-5", "claude-sonnet-4-6"])
|
||||
self.assertEqual(
|
||||
[item.id for item in available],
|
||||
["claude-sonnet-4-5", "claude-sonnet-4-6"],
|
||||
)
|
||||
|
||||
client.set_thinking_level("high")
|
||||
self.assertEqual(client.get_state().thinking_level, "high")
|
||||
@@ -853,8 +877,12 @@ class RpcClientTests(unittest.TestCase):
|
||||
seen_unknown: list[str] = []
|
||||
|
||||
with self.make_client() as client:
|
||||
client.on_extension_error(lambda event: seen_extension_errors.append(event.error))
|
||||
client.on_unknown_notification(lambda event: seen_unknown.append(str(event.payload.get("type"))))
|
||||
client.on_extension_error(
|
||||
lambda event: seen_extension_errors.append(event.error)
|
||||
)
|
||||
client.on_unknown_notification(
|
||||
lambda event: seen_unknown.append(str(event.payload.get("type")))
|
||||
)
|
||||
client.prompt_and_wait("notifications", timeout=2.0)
|
||||
|
||||
self.assertEqual(seen_extension_errors, ["boom"])
|
||||
@@ -879,20 +907,32 @@ class RpcClientTests(unittest.TestCase):
|
||||
errors: list[BaseException] = []
|
||||
|
||||
with self.make_client() as client:
|
||||
|
||||
def run_prompt() -> None:
|
||||
try:
|
||||
results.append(client.prompt_and_wait("slow", timeout=2.0).require_assistant_text())
|
||||
except BaseException as exc: # pragma: no cover - defensive thread capture
|
||||
results.append(
|
||||
client.prompt_and_wait(
|
||||
"slow", timeout=2.0
|
||||
).require_assistant_text()
|
||||
)
|
||||
except (
|
||||
BaseException
|
||||
) as exc: # pragma: no cover - defensive thread capture
|
||||
errors.append(exc)
|
||||
|
||||
thread = threading.Thread(target=run_prompt)
|
||||
thread.start()
|
||||
|
||||
deadline = time.time() + 1.0
|
||||
while client._prompt_lifecycle.active_operation != "prompt_and_wait" and time.time() < deadline:
|
||||
while (
|
||||
client._prompt_lifecycle.active_operation != "prompt_and_wait"
|
||||
and time.time() < deadline
|
||||
):
|
||||
time.sleep(0.01)
|
||||
|
||||
self.assertEqual(client._prompt_lifecycle.active_operation, "prompt_and_wait")
|
||||
self.assertEqual(
|
||||
client._prompt_lifecycle.active_operation, "prompt_and_wait"
|
||||
)
|
||||
with self.assertRaises(RpcConcurrencyError):
|
||||
client.collect_events(timeout=1.0)
|
||||
|
||||
@@ -904,7 +944,11 @@ class RpcClientTests(unittest.TestCase):
|
||||
|
||||
def test_listener_mutation_does_not_change_retained_turn(self) -> None:
|
||||
with self.make_client() as client:
|
||||
client.on_message_end(lambda event: event.message["content"].__setitem__(0, {"type": "text", "text": "mutated"}))
|
||||
client.on_message_end(
|
||||
lambda event: event.message["content"].__setitem__(
|
||||
0, {"type": "text", "text": "mutated"}
|
||||
)
|
||||
)
|
||||
turn = client.prompt_and_wait("say hello", timeout=2.0)
|
||||
messages = client.get_messages()
|
||||
|
||||
@@ -941,12 +985,16 @@ class RpcClientTests(unittest.TestCase):
|
||||
listener_errors: list[tuple[str, str | None, str]] = []
|
||||
client = self.make_client()
|
||||
client.on_notification(
|
||||
lambda notification: (_ for _ in ()).throw(RuntimeError("boom"))
|
||||
if notification.type == "turn_start"
|
||||
else None
|
||||
lambda notification: (
|
||||
(_ for _ in ()).throw(RuntimeError("boom"))
|
||||
if notification.type == "turn_start"
|
||||
else None
|
||||
)
|
||||
)
|
||||
client.on_listener_error(
|
||||
lambda event: listener_errors.append((event.listener_kind, event.source_type, str(event.error)))
|
||||
lambda event: listener_errors.append(
|
||||
(event.listener_kind, event.source_type, str(event.error))
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -986,7 +1034,6 @@ class RpcClientTests(unittest.TestCase):
|
||||
self.assertIn("max_event_history", str(ctx.exception))
|
||||
|
||||
|
||||
|
||||
HANGING_SERVER = textwrap.dedent(
|
||||
"""
|
||||
import json
|
||||
@@ -1051,22 +1098,32 @@ class StopUnblocksPromptAndWaitTests(unittest.TestCase):
|
||||
|
||||
# Wait until the prompt is in flight.
|
||||
deadline = time.time() + 2.0
|
||||
while client._prompt_lifecycle.active_operation != "prompt_and_wait" and time.time() < deadline:
|
||||
while (
|
||||
client._prompt_lifecycle.active_operation != "prompt_and_wait"
|
||||
and time.time() < deadline
|
||||
):
|
||||
time.sleep(0.01)
|
||||
self.assertEqual(client._prompt_lifecycle.active_operation, "prompt_and_wait")
|
||||
self.assertEqual(
|
||||
client._prompt_lifecycle.active_operation, "prompt_and_wait"
|
||||
)
|
||||
|
||||
t0 = time.time()
|
||||
client.stop()
|
||||
thread.join(timeout=2.0)
|
||||
elapsed = time.time() - t0
|
||||
|
||||
self.assertFalse(thread.is_alive(), "prompt_and_wait did not return after stop()")
|
||||
self.assertLess(elapsed, 2.0, f"stop() took {elapsed:.2f}s to unblock prompt_and_wait")
|
||||
self.assertFalse(
|
||||
thread.is_alive(), "prompt_and_wait did not return after stop()"
|
||||
)
|
||||
self.assertLess(
|
||||
elapsed, 2.0, f"stop() took {elapsed:.2f}s to unblock prompt_and_wait"
|
||||
)
|
||||
self.assertEqual(len(errors), 1)
|
||||
self.assertIsInstance(errors[0], RpcProcessExitError)
|
||||
finally:
|
||||
# stop() is idempotent; safe to call again on cleanup paths.
|
||||
client.stop()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -2,12 +2,11 @@ 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
|
||||
from omp_rpc.host_uris import normalize_read_result
|
||||
|
||||
|
||||
URI_SERVER = textwrap.dedent(
|
||||
@@ -114,7 +113,9 @@ class HostUriHelperTests(unittest.TestCase):
|
||||
|
||||
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]
|
||||
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")
|
||||
@@ -125,7 +126,9 @@ class HostUriHelperTests(unittest.TestCase):
|
||||
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)
|
||||
uri = host_uri(
|
||||
scheme="db", read=lambda url, ctx: "x", write=lambda url, content, ctx: None
|
||||
)
|
||||
self.assertTrue(uri.writable)
|
||||
|
||||
|
||||
@@ -168,7 +171,9 @@ class RpcHostUriBridgeTests(unittest.TestCase):
|
||||
"immutable": True,
|
||||
}
|
||||
|
||||
with self._make_client(host_uris=(host_uri(scheme="db", read=read_db),)) as client:
|
||||
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"}')
|
||||
@@ -212,7 +217,9 @@ class RpcHostUriBridgeTests(unittest.TestCase):
|
||||
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:
|
||||
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"))
|
||||
|
||||
@@ -207,7 +207,9 @@ class ProtocolParsingTests(unittest.TestCase):
|
||||
)
|
||||
self.assertEqual(state.system_prompt, ())
|
||||
|
||||
def test_parse_session_state_rejects_non_string_in_system_prompt_array(self) -> None:
|
||||
def test_parse_session_state_rejects_non_string_in_system_prompt_array(
|
||||
self,
|
||||
) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
parse_session_state(
|
||||
{
|
||||
@@ -233,7 +235,9 @@ class ProtocolParsingTests(unittest.TestCase):
|
||||
|
||||
def test_parse_extension_ui_request_rejects_invalid_method(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
parse_notification({"type": "extension_ui_request", "id": "ui-1", "method": "launch"})
|
||||
parse_notification(
|
||||
{"type": "extension_ui_request", "id": "ui-1", "method": "launch"}
|
||||
)
|
||||
|
||||
def test_parse_message_update_rejects_invalid_assistant_done_reason(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
|
||||
@@ -13,7 +13,9 @@ class _Sentinel(Exception):
|
||||
|
||||
def _start_and_capture(**kwargs):
|
||||
client = RpcClient(**kwargs)
|
||||
with patch("omp_rpc.client.subprocess.Popen", side_effect=_Sentinel("aborted")) as mock_popen:
|
||||
with patch(
|
||||
"omp_rpc.client.subprocess.Popen", side_effect=_Sentinel("aborted")
|
||||
) as mock_popen:
|
||||
with pytest.raises(_Sentinel):
|
||||
client.start()
|
||||
assert mock_popen.call_count == 1
|
||||
|
||||
Reference in New Issue
Block a user