fix(coding-agent): reconstruct compacted RPC prompt histories

This commit is contained in:
Wolfgang Schoenberger
2026-07-18 14:42:31 -07:00
parent f8cc72ff36
commit e0712265f2
7 changed files with 159 additions and 30 deletions
+35 -2
View File
@@ -1065,9 +1065,12 @@ class RpcClient:
def _build_prompt_turn(self, events: tuple[RpcAgentEvent, ...]) -> PromptTurn:
final_messages: tuple[AgentMessage, ...] = ()
for event in reversed(events):
for event_index in range(len(events) - 1, -1, -1):
event = events[event_index]
if isinstance(event, AgentEndEvent):
final_messages = event.messages
final_messages = self._complete_agent_end_messages(
events[:event_index], event
)
break
assistant_message: AssistantMessage | None = None
@@ -1093,6 +1096,36 @@ class RpcClient:
else None,
)
@staticmethod
def _complete_agent_end_messages(
events: tuple[RpcAgentEvent, ...], terminal: AgentEndEvent
) -> tuple[AgentMessage, ...]:
if (
terminal.message_count is None
or terminal.message_count <= len(terminal.messages)
):
return terminal.messages
run_start = 0
for event_index in range(len(events) - 1, -1, -1):
if isinstance(events[event_index], AgentStartEvent):
run_start = event_index + 1
break
streamed_messages = tuple(
event.message
for event in events[run_start:]
if isinstance(event, MessageEndEvent)
)
streamed_prefix_count = terminal.message_count - len(terminal.messages)
if streamed_prefix_count > len(streamed_messages):
raise RpcError(
"Compacted agent_end references "
f"{streamed_prefix_count} streamed messages, but only "
f"{len(streamed_messages)} were retained"
)
return streamed_messages[:streamed_prefix_count] + terminal.messages
def _wait_for_agent_end(
self,
start_index: int,
+4 -2
View File
@@ -2,7 +2,7 @@ from __future__ import annotations
import base64
import mimetypes
from dataclasses import dataclass
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Final, Literal, NotRequired, TypedDict, TypeAlias, cast
@@ -905,6 +905,7 @@ class AgentStartEvent:
class AgentEndEvent:
messages: tuple[AgentMessage, ...]
type: Literal["agent_end"] = "agent_end"
message_count: int | None = field(default=None, kw_only=True)
@dataclass(slots=True, frozen=True)
@@ -1500,7 +1501,8 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
return AgentEndEvent(
messages=parse_agent_messages(
cast(JsonValue | None, payload.get("messages"))
)
),
message_count=_optional_int(payload, "messageCount"),
)
if event_type == "turn_start":
return TurnStartEvent()
+40 -5
View File
@@ -86,7 +86,12 @@ FAKE_SERVER = textwrap.dedent(
"dumpTools": [{"name": "read", "description": "Read files", "parameters": {"type": "object"}}] + registered_host_tools,
}
def emit_prompt_turn(text: str, delay: float = 0.0, include_extra_events: bool = False):
def emit_prompt_turn(
text: str,
delay: float = 0.0,
include_extra_events: bool = False,
compact_terminal: bool = False,
):
global last_assistant_text, messages
print(json.dumps({"type": "agent_start"}), flush=True)
print(json.dumps({"type": "turn_start"}), flush=True)
@@ -198,9 +203,24 @@ FAKE_SERVER = textwrap.dedent(
assistant = assistant_message(text)
print(json.dumps({"type": "message_end", "message": assistant}), flush=True)
print(json.dumps({"type": "turn_end", "message": assistant, "toolResults": []}), flush=True)
print(json.dumps({"type": "agent_end", "messages": [assistant]}), flush=True)
last_assistant_text = text
messages = [assistant]
if compact_terminal:
terminal = assistant_message("terminal")
print(
json.dumps(
{
"type": "agent_end",
"messages": [terminal],
"messageCount": 2,
}
),
flush=True,
)
last_assistant_text = "terminal"
messages = [assistant, terminal]
else:
print(json.dumps({"type": "agent_end", "messages": [assistant]}), flush=True)
last_assistant_text = text
messages = [assistant]
def respond(request_id, command, data=None, success=True, error=None):
payload = {"id": request_id, "type": "response", "command": command, "success": success}
@@ -384,7 +404,12 @@ FAKE_SERVER = textwrap.dedent(
if message == "notifications":
print(json.dumps({"type": "extension_error", "extensionPath": "/tmp/ext.py", "event": "run", "error": "boom"}), flush=True)
print(json.dumps({"type": "unknown_future_event", "value": 1}), flush=True)
emit_prompt_turn("pong", delay=0.3 if message == "slow" else 0.0, include_extra_events=message == "all events")
emit_prompt_turn(
"pong",
delay=0.3 if message == "slow" else 0.0,
include_extra_events=message == "all events",
compact_terminal=message == "compacted turn",
)
elif command_type == "host_tool_update":
print(
json.dumps(
@@ -645,6 +670,16 @@ class RpcClientTests(unittest.TestCase):
self.assertEqual(turn.require_assistant_text(), "pong")
self.assertGreaterEqual(len(turn.events), 3)
def test_prompt_and_wait_reconstructs_compacted_terminal_messages(self) -> None:
with self.make_client() as client:
turn = client.prompt_and_wait("compacted turn", timeout=2.0)
self.assertEqual(
[message["content"][0]["text"] for message in turn.messages],
["pong", "terminal"],
)
self.assertEqual(turn.require_assistant_text(), "terminal")
def test_custom_tools_are_registered_and_executed_via_rpc(self) -> None:
def echo_host(args: dict[str, str], context) -> str:
context.send_update(f"working:{args['message']}")
+6
View File
@@ -125,11 +125,17 @@ class ProtocolParsingTests(unittest.TestCase):
"timestamp": 1,
}
],
"messageCount": 1,
}
)
self.assertIsInstance(notification, AgentEndEvent)
self.assertEqual(assistant_text(notification.messages[0]), "hello")
self.assertEqual(notification.message_count, 1)
legacy = AgentEndEvent(notification.messages, "agent_end")
self.assertEqual(legacy.type, "agent_end")
self.assertIsNone(legacy.message_count)
def test_parse_extension_ui_request(self) -> None:
notification = parse_notification(