fix(rpc): preserve v2 snapshot semantics
This commit is contained in:
@@ -117,6 +117,7 @@ _TODO_STATUS_VALUES = frozenset({"pending", "in_progress", "completed", "abandon
|
||||
_MAX_RPC_FRAME_BYTES = 1024 * 1024
|
||||
_MAX_RPC_REASSEMBLED_BYTES = 64 * 1024 * 1024
|
||||
_RPC_CHUNK_PAYLOAD_BYTES = 256 * 1024
|
||||
_RPC_MESSAGES_PAGE_BUSY_ERROR = "Cannot page messages while the session is changing"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
@@ -1007,29 +1008,39 @@ class RpcClient:
|
||||
|
||||
def get_messages(self) -> tuple[AgentMessage, ...]:
|
||||
if self._protocol_version == 2:
|
||||
messages: list[AgentMessage] = []
|
||||
seen_cursors: set[str] = set()
|
||||
total_messages: int | None = None
|
||||
cursor: str | None = None
|
||||
while True:
|
||||
page = self.get_messages_page(cursor=cursor, limit=256)
|
||||
if total_messages is not None and page.total_messages != total_messages:
|
||||
try:
|
||||
messages: list[AgentMessage] = []
|
||||
seen_cursors: set[str] = set()
|
||||
total_messages: int | None = None
|
||||
cursor: str | None = None
|
||||
while True:
|
||||
page = self.get_messages_page(cursor=cursor, limit=256)
|
||||
if (
|
||||
total_messages is not None
|
||||
and page.total_messages != total_messages
|
||||
):
|
||||
raise RpcError(
|
||||
"RPC message pagination returned an inconsistent total"
|
||||
)
|
||||
total_messages = page.total_messages
|
||||
messages.extend(page.messages)
|
||||
cursor = page.next_cursor
|
||||
if cursor is None:
|
||||
break
|
||||
if cursor in seen_cursors:
|
||||
raise RpcError("RPC message pagination repeated a cursor")
|
||||
seen_cursors.add(cursor)
|
||||
if len(messages) != total_messages:
|
||||
raise RpcError(
|
||||
"RPC message pagination returned an inconsistent total"
|
||||
"RPC message pagination ended before the advertised total"
|
||||
)
|
||||
total_messages = page.total_messages
|
||||
messages.extend(page.messages)
|
||||
cursor = page.next_cursor
|
||||
if cursor is None:
|
||||
break
|
||||
if cursor in seen_cursors:
|
||||
raise RpcError("RPC message pagination repeated a cursor")
|
||||
seen_cursors.add(cursor)
|
||||
if len(messages) != total_messages:
|
||||
raise RpcError(
|
||||
"RPC message pagination ended before the advertised total"
|
||||
)
|
||||
return tuple(messages)
|
||||
return tuple(messages)
|
||||
except RpcCommandError as error:
|
||||
if (
|
||||
error.command != "get_messages_page"
|
||||
or error.error != _RPC_MESSAGES_PAGE_BUSY_ERROR
|
||||
):
|
||||
raise
|
||||
payload = self._request("get_messages")
|
||||
return parse_agent_messages(cast(JsonValue | None, payload.get("messages")))
|
||||
|
||||
|
||||
@@ -447,6 +447,7 @@ V2_MESSAGES_SERVER = textwrap.dedent(
|
||||
"""
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
message = {
|
||||
@@ -507,6 +508,17 @@ V2_MESSAGES_SERVER = textwrap.dedent(
|
||||
}
|
||||
)
|
||||
elif command_type == "get_messages_page":
|
||||
if os.environ.get("V2_MESSAGES_BUSY") == "1":
|
||||
emit(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": command_type,
|
||||
"success": False,
|
||||
"error": "Cannot page messages while the session is changing",
|
||||
}
|
||||
)
|
||||
continue
|
||||
emit(
|
||||
{
|
||||
"id": request_id,
|
||||
@@ -520,6 +532,26 @@ V2_MESSAGES_SERVER = textwrap.dedent(
|
||||
},
|
||||
}
|
||||
)
|
||||
elif command_type == "get_messages":
|
||||
emit(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": command_type,
|
||||
"success": True,
|
||||
"data": {
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "text", "text": "streaming snapshot"}
|
||||
],
|
||||
"timestamp": 3,
|
||||
}
|
||||
]
|
||||
},
|
||||
}
|
||||
)
|
||||
else:
|
||||
emit(
|
||||
{
|
||||
@@ -956,6 +988,19 @@ class RpcClientTests(unittest.TestCase):
|
||||
self.assertEqual(len(messages), 1)
|
||||
self.assertEqual(len(messages[0]["content"][0]["text"]), 1024 * 1024)
|
||||
|
||||
def test_protocol_v2_get_messages_falls_back_to_streaming_snapshot(self) -> None:
|
||||
with self.make_client(
|
||||
server=V2_MESSAGES_SERVER, env={"V2_MESSAGES_BUSY": "1"}
|
||||
) as client:
|
||||
with self.assertRaisesRegex(
|
||||
RpcCommandError, "Cannot page messages while the session is changing"
|
||||
):
|
||||
client.get_messages_page()
|
||||
messages = client.get_messages()
|
||||
|
||||
self.assertEqual(len(messages), 1)
|
||||
self.assertEqual(messages[0]["content"][0]["text"], "streaming snapshot")
|
||||
|
||||
def test_collect_events_returns_turn_events(self) -> None:
|
||||
with self.make_client() as client:
|
||||
client.prompt("slow")
|
||||
|
||||
Reference in New Issue
Block a user