Fix compacted tool history for resumed sessions
This commit is contained in:
@@ -37,7 +37,7 @@ from .agent_prompting import (
|
||||
build_system_prompt_parts,
|
||||
render_system_prompt,
|
||||
)
|
||||
from .agent_session import AgentSessionState
|
||||
from .agent_session import AgentSessionState, sanitize_model_message_sequence
|
||||
from .agent_slash_commands import preprocess_slash_command
|
||||
from .agent_tools import (
|
||||
AgentTool,
|
||||
@@ -2468,7 +2468,9 @@ class LocalCodingAgent:
|
||||
if compact_end <= prefix_count:
|
||||
return False
|
||||
candidates = session.messages[prefix_count:compact_end]
|
||||
preserved_tail = list(session.messages[compact_end:])
|
||||
preserved_tail = sanitize_model_message_sequence(
|
||||
list(session.messages[compact_end:])
|
||||
)
|
||||
if not candidates:
|
||||
return False
|
||||
compacted_tokens = sum(
|
||||
@@ -2492,13 +2494,13 @@ class LocalCodingAgent:
|
||||
turn_index=turn_index,
|
||||
estimated_tokens_before=usage_total,
|
||||
estimated_tokens_removed=compacted_tokens,
|
||||
preserved_tail_count=tail_count,
|
||||
preserved_tail_count=len(preserved_tail),
|
||||
preserved_tail=preserved_tail,
|
||||
)
|
||||
session.messages = (
|
||||
session.messages[:prefix_count]
|
||||
+ [compact_message]
|
||||
+ session.messages[compact_end:]
|
||||
+ preserved_tail
|
||||
)
|
||||
stream_events.append(
|
||||
{
|
||||
@@ -2507,7 +2509,7 @@ class LocalCodingAgent:
|
||||
'compacted_message_count': len(candidates),
|
||||
'estimated_tokens_before': usage_total,
|
||||
'estimated_tokens_removed': compacted_tokens,
|
||||
'preserved_tail_count': tail_count,
|
||||
'preserved_tail_count': len(preserved_tail),
|
||||
'preserved_tail_ids': [
|
||||
message.message_id for message in preserved_tail if message.message_id
|
||||
],
|
||||
|
||||
@@ -96,6 +96,84 @@ class AgentMessage:
|
||||
)
|
||||
|
||||
|
||||
def _tool_call_id(tool_call: JSONDict) -> str | None:
|
||||
raw_id = tool_call.get('id')
|
||||
return raw_id if isinstance(raw_id, str) and raw_id else None
|
||||
|
||||
|
||||
def _assistant_tool_call_ids(message: AgentMessage) -> list[str]:
|
||||
if message.role != 'assistant' or not message.tool_calls:
|
||||
return []
|
||||
ids: list[str] = []
|
||||
for tool_call in message.tool_calls:
|
||||
call_id = _tool_call_id(tool_call)
|
||||
if call_id:
|
||||
ids.append(call_id)
|
||||
return ids
|
||||
|
||||
|
||||
def sanitize_model_message_sequence(
|
||||
messages: list[AgentMessage],
|
||||
) -> list[AgentMessage]:
|
||||
"""Remove invalid tool-call fragments from model-facing history.
|
||||
|
||||
Chat backends require a tool result to appear immediately after the
|
||||
assistant message that requested it. Compaction can otherwise preserve a
|
||||
tail starting with a ``tool`` message after the matching assistant tool call
|
||||
has been summarized away. Display history is append-only and should not use
|
||||
this helper; it is only for model-facing context.
|
||||
"""
|
||||
|
||||
cleaned: list[AgentMessage] = []
|
||||
index = 0
|
||||
while index < len(messages):
|
||||
message = messages[index]
|
||||
if message.role == 'tool':
|
||||
index += 1
|
||||
continue
|
||||
|
||||
tool_call_ids = _assistant_tool_call_ids(message)
|
||||
if not tool_call_ids:
|
||||
cleaned.append(message)
|
||||
index += 1
|
||||
continue
|
||||
|
||||
group: list[AgentMessage] = []
|
||||
seen: set[str] = set()
|
||||
cursor = index + 1
|
||||
while cursor < len(messages) and messages[cursor].role == 'tool':
|
||||
tool_message = messages[cursor]
|
||||
call_id = tool_message.tool_call_id
|
||||
if (
|
||||
call_id is None
|
||||
or call_id not in tool_call_ids
|
||||
or call_id in seen
|
||||
):
|
||||
break
|
||||
group.append(tool_message)
|
||||
seen.add(call_id)
|
||||
cursor += 1
|
||||
|
||||
if len(seen) == len(tool_call_ids):
|
||||
cleaned.append(message)
|
||||
cleaned.extend(group)
|
||||
index = cursor
|
||||
continue
|
||||
|
||||
# Drop the incomplete assistant tool-call request and any immediately
|
||||
# adjacent partial tool results tied to it. Keeping either side would
|
||||
# make the next model request invalid.
|
||||
index += 1
|
||||
while (
|
||||
index < len(messages)
|
||||
and messages[index].role == 'tool'
|
||||
and messages[index].tool_call_id in set(tool_call_ids)
|
||||
):
|
||||
index += 1
|
||||
|
||||
return cleaned
|
||||
|
||||
|
||||
@dataclass
|
||||
class AgentSessionState:
|
||||
system_prompt_parts: tuple[str, ...]
|
||||
@@ -592,6 +670,7 @@ class AgentSessionState:
|
||||
for message in messages
|
||||
if isinstance(message, dict)
|
||||
]
|
||||
model_messages = sanitize_model_message_sequence(model_messages)
|
||||
display_source = messages if display_messages is None else display_messages
|
||||
visible_messages = [
|
||||
AgentMessage.from_openai_message(message)
|
||||
|
||||
+2
-2
@@ -22,7 +22,7 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
from .agent_context_usage import estimate_tokens
|
||||
from .agent_types import UsageStats
|
||||
from .agent_session import AgentMessage
|
||||
from .agent_session import AgentMessage, sanitize_model_message_sequence
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .agent_runtime import LocalCodingAgent
|
||||
@@ -522,7 +522,7 @@ def compact_conversation(
|
||||
)
|
||||
|
||||
candidates = list(session.messages[prefix_count:compact_end])
|
||||
preserved_tail = list(session.messages[compact_end:])
|
||||
preserved_tail = sanitize_model_message_sequence(list(session.messages[compact_end:]))
|
||||
|
||||
if not candidates:
|
||||
return CompactionResult(
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from src.agent_session import AgentMessage, sanitize_model_message_sequence
|
||||
|
||||
|
||||
class TestSanitizeModelMessageSequence(unittest.TestCase):
|
||||
def test_removes_orphan_tool_message(self) -> None:
|
||||
messages = [
|
||||
AgentMessage(role='user', content='summary', message_id='summary'),
|
||||
AgentMessage(
|
||||
role='tool',
|
||||
content='orphan result',
|
||||
tool_call_id='call_1',
|
||||
message_id='tool_1',
|
||||
),
|
||||
AgentMessage(role='assistant', content='review', message_id='assistant_1'),
|
||||
AgentMessage(role='user', content='continue', message_id='user_1'),
|
||||
]
|
||||
|
||||
cleaned = sanitize_model_message_sequence(messages)
|
||||
|
||||
self.assertEqual(
|
||||
[message.message_id for message in cleaned],
|
||||
['summary', 'assistant_1', 'user_1'],
|
||||
)
|
||||
|
||||
def test_preserves_complete_assistant_tool_group(self) -> None:
|
||||
assistant = AgentMessage(
|
||||
role='assistant',
|
||||
content='',
|
||||
tool_calls=(
|
||||
{
|
||||
'id': 'call_1',
|
||||
'type': 'function',
|
||||
'function': {'name': 'read_file', 'arguments': '{}'},
|
||||
},
|
||||
),
|
||||
message_id='assistant_1',
|
||||
)
|
||||
tool = AgentMessage(
|
||||
role='tool',
|
||||
content='ok',
|
||||
tool_call_id='call_1',
|
||||
message_id='tool_1',
|
||||
)
|
||||
final = AgentMessage(role='assistant', content='done', message_id='assistant_2')
|
||||
|
||||
cleaned = sanitize_model_message_sequence([assistant, tool, final])
|
||||
|
||||
self.assertEqual(
|
||||
[message.message_id for message in cleaned],
|
||||
['assistant_1', 'tool_1', 'assistant_2'],
|
||||
)
|
||||
|
||||
def test_removes_incomplete_assistant_tool_call_group(self) -> None:
|
||||
messages = [
|
||||
AgentMessage(role='user', content='question', message_id='user_1'),
|
||||
AgentMessage(
|
||||
role='assistant',
|
||||
content='',
|
||||
tool_calls=(
|
||||
{
|
||||
'id': 'call_1',
|
||||
'type': 'function',
|
||||
'function': {'name': 'read_file', 'arguments': '{}'},
|
||||
},
|
||||
),
|
||||
message_id='assistant_1',
|
||||
),
|
||||
AgentMessage(role='user', content='next', message_id='user_2'),
|
||||
]
|
||||
|
||||
cleaned = sanitize_model_message_sequence(messages)
|
||||
|
||||
self.assertEqual(
|
||||
[message.message_id for message in cleaned],
|
||||
['user_1', 'user_2'],
|
||||
)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user