fix(python/omp-rpc): accepted string or array systemPrompt values in session state parsing

- Added `_optional_str_list` to normalize `systemPrompt` payloads into a tuple when absent, a bare string, or a string array, and to reject invalid shapes.
- Updated `SessionState` to hold `system_prompt` as a tuple and to parse daemon payloads with the new helper in `parse_session_state`.
- Extended protocol tests to verify array and legacy string parsing, default empty tuple behavior, and invalid `systemPrompt` values.
This commit is contained in:
can1357
2026-05-15 00:55:52 +02:00
parent bddf9989b5
commit 3e9ca3e279
2 changed files with 74 additions and 2 deletions
+25 -2
View File
@@ -164,6 +164,29 @@ def _optional_str(payload: JsonObject, field: str) -> str | None:
return value
def _optional_str_list(payload: JsonObject, field: str) -> tuple[str, ...]:
"""Parse an optional string-or-array-of-strings field.
The agent's `systemPrompt` (and similar) became `string[]` server-side
when multi-prompt support landed. Older daemons still emit a bare string,
so we accept either shape. Returns an empty tuple when the field is
absent or null.
"""
value = payload.get(field)
if value is None:
return ()
if isinstance(value, str):
return (value,)
if isinstance(value, list):
items: list[str] = []
for index, item in enumerate(value):
if not isinstance(item, str):
raise ValueError(f"{field}[{index}] must be a string")
items.append(item)
return tuple(items)
raise ValueError(f"{field} must be a string or an array of strings")
def _optional_bool(payload: JsonObject, field: str) -> bool | None:
value = payload.get(field)
if value is None:
@@ -669,7 +692,7 @@ class SessionState:
message_count: int
queued_message_count: int
todo_phases: tuple[TodoPhase, ...] = ()
system_prompt: str | None = None
system_prompt: tuple[str, ...] = ()
dump_tools: tuple[ToolDescriptor, ...] = ()
@@ -1133,7 +1156,7 @@ def parse_session_state(payload: JsonObject) -> SessionState:
message_count=int(payload.get("messageCount", 0)),
queued_message_count=int(payload.get("queuedMessageCount", 0)),
todo_phases=parse_todo_phases(cast(JsonValue | None, payload.get("todoPhases"))),
system_prompt=_optional_str(payload, "systemPrompt"),
system_prompt=_optional_str_list(payload, "systemPrompt"),
dump_tools=dump_tools,
)
+49
View File
@@ -82,6 +82,8 @@ class ProtocolParsingTests(unittest.TestCase):
self.assertEqual(state.follow_up_mode, "all")
self.assertEqual(state.model.id if state.model else None, "claude-sonnet-4-5")
self.assertEqual(state.todo_phases[0].tasks[0].status, "in_progress")
# Legacy bare-string systemPrompt is accepted and wrapped to a tuple.
self.assertEqual(state.system_prompt, ("You are useful.",))
self.assertEqual(state.dump_tools[0].name, "read")
def test_parse_agent_end_notification(self) -> None:
@@ -182,6 +184,53 @@ class ProtocolParsingTests(unittest.TestCase):
}
)
def test_parse_session_state_accepts_system_prompt_array(self) -> None:
state = parse_session_state(
{
"sessionId": "session-abc",
"steeringMode": "one-at-a-time",
"followUpMode": "one-at-a-time",
"interruptMode": "immediate",
"systemPrompt": ["base instructions", "extra policy"],
}
)
self.assertEqual(state.system_prompt, ("base instructions", "extra policy"))
def test_parse_session_state_defaults_system_prompt_to_empty_tuple(self) -> None:
state = parse_session_state(
{
"sessionId": "session-abc",
"steeringMode": "one-at-a-time",
"followUpMode": "one-at-a-time",
"interruptMode": "immediate",
}
)
self.assertEqual(state.system_prompt, ())
def test_parse_session_state_rejects_non_string_in_system_prompt_array(self) -> None:
with self.assertRaises(ValueError):
parse_session_state(
{
"sessionId": "session-abc",
"steeringMode": "one-at-a-time",
"followUpMode": "one-at-a-time",
"interruptMode": "immediate",
"systemPrompt": ["ok", 42],
}
)
def test_parse_session_state_rejects_invalid_system_prompt_shape(self) -> None:
with self.assertRaises(ValueError):
parse_session_state(
{
"sessionId": "session-abc",
"steeringMode": "one-at-a-time",
"followUpMode": "one-at-a-time",
"interruptMode": "immediate",
"systemPrompt": {"unexpected": "object"},
}
)
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"})