From ecbf759ec216a01e3ad42926b595e290a2752799 Mon Sep 17 00:00:00 2001 From: Frederico Luz Date: Wed, 29 Jul 2026 20:17:09 +0100 Subject: [PATCH] fix(rpc): preserve fast state in Python client --- python/omp-rpc/src/omp_rpc/protocol.py | 15 +++++++++++ python/omp-rpc/tests/test_client.py | 6 +++++ python/omp-rpc/tests/test_protocol.py | 35 ++++++++++++++++++++++++++ 3 files changed, 56 insertions(+) diff --git a/python/omp-rpc/src/omp_rpc/protocol.py b/python/omp-rpc/src/omp_rpc/protocol.py index f903964a7..1fa2cfc9a 100644 --- a/python/omp-rpc/src/omp_rpc/protocol.py +++ b/python/omp-rpc/src/omp_rpc/protocol.py @@ -233,6 +233,15 @@ def _optional_int(payload: JsonObject, field: str) -> int | None: return value +def _optional_float(payload: JsonObject, field: str) -> float | None: + value = payload.get(field) + if value is None: + return None + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ValueError(f"{field} must be a number") + return float(value) + + def _tuple_of_strings(values: object, *, field: str) -> tuple[str, ...] | None: if values is None: return None @@ -776,6 +785,9 @@ class SessionState: todo_phases: tuple[TodoPhase, ...] = () system_prompt: tuple[str, ...] = () dump_tools: tuple[ToolDescriptor, ...] = () + fast_mode_enabled: bool = False + fast_mode_active: bool = False + tokens_per_second: float | None = None @dataclass(slots=True, frozen=True) @@ -1344,6 +1356,9 @@ def parse_session_state(payload: JsonObject) -> SessionState: ), system_prompt=_optional_str_list(payload, "systemPrompt"), dump_tools=dump_tools, + fast_mode_enabled=bool(payload.get("fastModeEnabled", False)), + fast_mode_active=bool(payload.get("fastModeActive", False)), + tokens_per_second=_optional_float(payload, "tokensPerSecond"), ) diff --git a/python/omp-rpc/tests/test_client.py b/python/omp-rpc/tests/test_client.py index c9129ac8f..3ad15d703 100644 --- a/python/omp-rpc/tests/test_client.py +++ b/python/omp-rpc/tests/test_client.py @@ -84,6 +84,9 @@ FAKE_SERVER = textwrap.dedent( "interruptMode": interrupt_mode, "sessionId": "fake-session", "sessionName": session_name, + "fastModeEnabled": False, + "fastModeActive": True, + "tokensPerSecond": 7.25, "autoCompactionEnabled": auto_compaction_enabled, "messageCount": len(messages), "queuedMessageCount": 0, @@ -877,6 +880,9 @@ class RpcClientTests(unittest.TestCase): self.assertEqual( state.model.id if state.model else None, "claude-sonnet-4-5" ) + self.assertFalse(state.fast_mode_enabled) + self.assertTrue(state.fast_mode_active) + self.assertEqual(state.tokens_per_second, 7.25) result = client.bash("echo hello") self.assertEqual(result.output, "hello\n") diff --git a/python/omp-rpc/tests/test_protocol.py b/python/omp-rpc/tests/test_protocol.py index c9f7b2697..b19fa967a 100644 --- a/python/omp-rpc/tests/test_protocol.py +++ b/python/omp-rpc/tests/test_protocol.py @@ -51,6 +51,9 @@ class ProtocolParsingTests(unittest.TestCase): "sessionFile": "/tmp/test.jsonl", "sessionId": "session-123", "sessionName": "Scratchpad", + "fastModeEnabled": False, + "fastModeActive": True, + "tokensPerSecond": 12.5, "autoCompactionEnabled": True, "messageCount": 4, "queuedMessageCount": 1, @@ -95,6 +98,38 @@ class ProtocolParsingTests(unittest.TestCase): self.assertEqual(state.model.thinking.default_level, "medium") self.assertEqual(state.model.thinking.effort_map, {"high": "xhigh"}) self.assertTrue(state.model.thinking.supports_display) + self.assertFalse(state.fast_mode_enabled) + self.assertTrue(state.fast_mode_active) + self.assertEqual(state.tokens_per_second, 12.5) + + def test_parse_session_state_defaults_missing_fast_mode_and_throughput( + self, + ) -> None: + missing = object() + for tokens_per_second, expected in ( + (None, None), + (missing, None), + ): + with self.subTest(tokens_per_second=tokens_per_second): + payload = { + "sessionId": "session-123", + "steeringMode": "one-at-a-time", + "followUpMode": "all", + "interruptMode": "immediate", + } + if tokens_per_second is not missing: + payload["tokensPerSecond"] = tokens_per_second + + state = parse_session_state(payload) + + self.assertEqual( + ( + state.fast_mode_enabled, + state.fast_mode_active, + state.tokens_per_second, + ), + (False, False, expected), + ) def test_parse_agent_end_notification(self) -> None: notification = parse_notification(