chore: reformat

This commit is contained in:
can1357
2026-06-02 08:43:23 +02:00
parent 0bc9bc25b4
commit 2ecb5fd9fa
8 changed files with 716 additions and 198 deletions
+81 -24
View File
@@ -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()
+13 -6
View File
@@ -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"))
+6 -2
View File
@@ -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):
+3 -1
View File
@@ -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