fix(rpc): preserve fast state in Python client

This commit is contained in:
Frederico Luz
2026-07-29 20:17:09 +01:00
parent 07f0ba8a42
commit ecbf759ec2
3 changed files with 56 additions and 0 deletions
+15
View File
@@ -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"),
)
+6
View File
@@ -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")
+35
View File
@@ -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(