fix(rpc): preserve fast state in Python client
This commit is contained in:
@@ -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"),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user