fix session display replay consistency

This commit is contained in:
wuyang6
2026-06-30 14:40:20 +08:00
parent 84ff22fc65
commit 0007d99763
4 changed files with 291 additions and 19 deletions
+36 -9
View File
@@ -3169,8 +3169,42 @@ def _stored_display_messages(
stored: StoredAgentSession,
) -> tuple[dict[str, Any], ...]:
if stored.display_messages:
return tuple(dict(message) for message in stored.display_messages)
return tuple(dict(message) for message in stored.messages)
messages = tuple(dict(message) for message in stored.display_messages)
else:
messages = tuple(dict(message) for message in stored.messages)
return tuple(message for message in messages if _is_visible_display_message(message))
_INTERNAL_DISPLAY_KINDS = {
'compact_boundary',
'compact_summary',
'continuation_request',
'file_history_replay',
'plugin_tool_runtime',
'runtime_context',
'snipped_message',
'system_context',
}
def _is_visible_display_message(message: dict[str, Any]) -> bool:
if message.get('role') == 'system':
return False
metadata = message.get('metadata')
if isinstance(metadata, dict):
kind = metadata.get('kind')
if isinstance(kind, str) and kind in _INTERNAL_DISPLAY_KINDS:
return False
content = message.get('content')
if isinstance(content, str):
stripped = content.strip()
if stripped.startswith('<system-reminder>'):
return False
if message.get('role') == 'user' and stripped.startswith(
'This session is being continued from a previous conversation'
):
return False
return True
def _stored_session_has_incomplete_tail(stored: StoredAgentSession) -> bool:
@@ -3572,10 +3606,7 @@ def _save_in_progress_session(
stored = load_agent_session(safe_id, directory=directory)
except (FileNotFoundError, OSError, json.JSONDecodeError):
return None
messages = list(stored.messages)
display_messages = list(_stored_display_messages(stored))
if messages and _is_pending_user_message(dict(messages[-1])):
messages.pop()
if display_messages and _is_pending_user_message(dict(display_messages[-1])):
display_messages.pop()
pending_message = {
@@ -3585,9 +3616,6 @@ def _save_in_progress_session(
'metadata': pending_metadata,
'message_id': f'user_pending_{int(time.time() * 1000)}',
}
messages.append(
dict(pending_message)
)
display_messages.append(dict(pending_message))
budget_state = (
dict(stored.budget_state)
@@ -3599,7 +3627,6 @@ def _save_in_progress_session(
save_agent_session(
replace(
stored,
messages=tuple(messages),
display_messages=tuple(display_messages),
budget_state=budget_state,
),
+56 -5
View File
@@ -724,7 +724,11 @@ class LocalCodingAgent:
if stored_resume_state is not None:
starting_usage = usage_from_payload(stored_resume_state.usage)
starting_cost_usd = stored_resume_state.total_cost_usd
starting_tool_calls = stored_resume_state.tool_calls
starting_tool_calls = self._sanitize_persisted_tool_call_count(
stored_resume_state.tool_calls,
stored_resume_state.messages,
stored_resume_state.display_messages,
)
starting_session_turns = stored_resume_state.turns
budget_state = (
stored_resume_state.budget_state
@@ -1830,6 +1834,7 @@ class LocalCodingAgent:
)
return turn, ()
request_messages = session.to_openai_messages()
assistant_index = session.start_assistant(
message_id=f'assistant_{len(session.messages)}'
)
@@ -1837,7 +1842,7 @@ class LocalCodingAgent:
finish_reason: str | None = None
events: list[StreamEvent] = []
for event in self.client.stream(
session.to_openai_messages(),
request_messages,
tool_specs,
output_schema=self.runtime_config.output_schema,
):
@@ -4416,14 +4421,22 @@ class LocalCodingAgent:
previous = None
if previous is not None:
previous_turns = previous.turns
previous_tool_calls = previous.tool_calls
previous_tool_calls = self._sanitize_persisted_tool_call_count(
previous.tool_calls,
previous.messages,
previous.display_messages,
)
if isinstance(previous.budget_state, dict):
previous_budget_state = dict(previous.budget_state)
total_tool_calls = self._merge_tool_call_count(
previous_tool_calls,
result.tool_calls,
)
budget_state = {
'model_calls': int(previous_budget_state.get('model_calls', 0))
+ max(result.turns, 0),
'session_turns': previous_turns + result.turns,
'tool_calls': previous_tool_calls + result.tool_calls,
'tool_calls': total_tool_calls,
'delegated_tasks': sum(
1 for entry in result.file_history if entry.get('action') in ('delegate_agent', 'Agent')
),
@@ -4438,7 +4451,7 @@ class LocalCodingAgent:
messages=session.model_transcript(),
display_messages=session.display_transcript(),
turns=previous_turns + result.turns,
tool_calls=previous_tool_calls + result.tool_calls,
tool_calls=total_tool_calls,
usage=result.usage.to_dict(),
total_cost_usd=result.total_cost_usd,
file_history=result.file_history,
@@ -4463,6 +4476,44 @@ class LocalCodingAgent:
transcript=session.transcript(),
)
@staticmethod
def _merge_tool_call_count(previous_tool_calls: int, result_tool_calls: int) -> int:
previous = max(0, int(previous_tool_calls or 0))
current = max(0, int(result_tool_calls or 0))
# Resumed runs initialize the runtime counter from persisted state, so
# result.tool_calls is usually already session-cumulative. Do not add it
# to the previous total again.
if current >= previous:
return current
# Some early-return paths still return a per-run delta.
return previous + current
@staticmethod
def _sanitize_persisted_tool_call_count(
persisted_tool_calls: int,
model_messages: tuple[dict[str, object], ...],
display_messages: tuple[dict[str, object], ...],
) -> int:
persisted = max(0, int(persisted_tool_calls or 0))
counted = max(
LocalCodingAgent._count_assistant_tool_calls(model_messages),
LocalCodingAgent._count_assistant_tool_calls(display_messages),
)
if counted and persisted > counted * 100:
return counted
return persisted
@staticmethod
def _count_assistant_tool_calls(messages: tuple[dict[str, object], ...]) -> int:
total = 0
for message in messages:
if not isinstance(message, dict) or message.get('role') != 'assistant':
continue
tool_calls = message.get('tool_calls')
if isinstance(tool_calls, (list, tuple)):
total += len(tool_calls)
return total
def _inject_runtime_guidance(
self,
session: AgentSessionState,
+48 -2
View File
@@ -174,6 +174,42 @@ def sanitize_model_message_sequence(
return cleaned
_INTERNAL_DISPLAY_KINDS = {
'compact_boundary',
'compact_summary',
'continuation_request',
'file_history_replay',
'plugin_tool_runtime',
'runtime_context',
'snipped_message',
'system_context',
}
def is_display_visible_message(message: AgentMessage) -> bool:
"""Return whether a persisted message belongs in the user-visible transcript.
Model-facing context contains system prompts, compact summaries, replay
reminders, and other runtime-only messages. Those must survive in
``messages`` but must not be backfilled into ``display_messages`` after
compaction or old-session migration.
"""
if message.role == 'system':
return False
metadata = message.metadata if isinstance(message.metadata, dict) else {}
kind = metadata.get('kind')
if isinstance(kind, str) and kind in _INTERNAL_DISPLAY_KINDS:
return False
content = message.content.strip()
if content.startswith('<system-reminder>'):
return False
if message.role == 'user' and content.startswith(
'This session is being continued from a previous conversation'
):
return False
return True
@dataclass
class AgentSessionState:
system_prompt_parts: tuple[str, ...]
@@ -188,7 +224,10 @@ class AgentSessionState:
# AgentSessionState directly. Runtime-created sessions append messages
# through helpers and explicitly decide whether each message is visible.
if self.messages and not self.display_messages:
self.display_messages = list(self.messages)
self.display_messages = [
message for message in self.messages
if is_display_visible_message(message)
]
@classmethod
def create(
@@ -677,8 +716,15 @@ class AgentSessionState:
for message in display_source
if isinstance(message, dict)
]
visible_messages = [
message for message in visible_messages
if is_display_visible_message(message)
]
if display_messages is not None and not visible_messages:
visible_messages = list(model_messages)
visible_messages = [
message for message in model_messages
if is_display_visible_message(message)
]
return cls(
system_prompt_parts=tuple(system_prompt_parts),
user_context=dict(user_context or {}),
+151 -3
View File
@@ -32,6 +32,41 @@ DEFAULT_SESSION_DIR = Path('.port_sessions')
DEFAULT_AGENT_SESSION_DIR = DEFAULT_SESSION_DIR / 'agent'
AGENT_SESSION_DB_FILENAME = 'sessions.db'
_INTERNAL_DISPLAY_KINDS = {
'compact_boundary',
'compact_summary',
'continuation_request',
'file_history_replay',
'plugin_tool_runtime',
'runtime_context',
'snipped_message',
'system_context',
}
def _is_display_message_visible(message: JSONDict) -> bool:
if message.get('role') == 'system':
return False
metadata = message.get('metadata')
if isinstance(metadata, dict):
kind = metadata.get('kind')
if isinstance(kind, str) and kind in _INTERNAL_DISPLAY_KINDS:
return False
content = message.get('content')
if isinstance(content, str):
stripped = content.strip()
if stripped.startswith('<system-reminder>'):
return False
if message.get('role') == 'user' and stripped.startswith(
'This session is being continued from a previous conversation'
):
return False
return True
def _filter_display_messages(messages: tuple[JSONDict, ...]) -> tuple[JSONDict, ...]:
return tuple(message for message in messages if _is_display_message_visible(message))
def save_session(session: StoredSession, directory: Path | None = None) -> Path:
target_dir = directory or DEFAULT_SESSION_DIR
@@ -84,12 +119,20 @@ def save_agent_session(session: StoredAgentSession, directory: Path | None = Non
session_dir.mkdir(parents=True, exist_ok=True)
path = session_dir / 'session.json'
payload = asdict(session)
display_messages = (
_filter_display_messages(
tuple(message for message in session.display_messages if isinstance(message, dict))
)
or _filter_display_messages(
tuple(message for message in session.messages if isinstance(message, dict))
)
)
payload['display_messages'] = list(display_messages)
_write_agent_session_db_payload(target_dir, session.session_id, payload, path)
_sync_agent_display_messages(
target_dir,
session.session_id,
tuple(message for message in session.display_messages if isinstance(message, dict))
or tuple(message for message in session.messages if isinstance(message, dict)),
display_messages,
)
path.write_text(json.dumps(payload, indent=2, ensure_ascii=False), encoding='utf-8')
return path
@@ -550,8 +593,9 @@ def _stored_agent_session_from_payload(data: JSONDict) -> StoredAgentSession:
for message in data.get('display_messages', messages)
if isinstance(message, dict)
)
display_messages = _filter_display_messages(display_messages)
if not display_messages:
display_messages = messages
display_messages = _filter_display_messages(messages)
session_metadata = (
dict(data.get('session_metadata', {}))
if isinstance(data.get('session_metadata'), dict)
@@ -693,6 +737,15 @@ def _sync_agent_display_messages(
""",
(session_id,),
).fetchall()
if not _display_rows_match_snapshot(existing_rows, display_messages):
_rebuild_agent_display_messages(
conn,
session_id,
display_messages,
existing_rows=existing_rows,
now=now,
)
return
existing: dict[int, sqlite3.Row] = {
int(row['seq']): row for row in existing_rows
}
@@ -886,6 +939,97 @@ def _sync_agent_display_messages(
)
def _display_rows_match_snapshot(
rows: list[sqlite3.Row],
display_messages: tuple[JSONDict, ...],
) -> bool:
stored_messages: list[JSONDict] = []
for row in sorted(rows, key=lambda item: int(item['seq'])):
try:
message = json.loads(row['message_json'])
except (TypeError, json.JSONDecodeError, KeyError):
return False
if not isinstance(message, dict):
return False
if _is_persistent_display_message(message):
continue
stored_messages.append(message)
expected_messages = [
message for message in display_messages if isinstance(message, dict)
]
if len(stored_messages) != len(expected_messages):
return False
for stored, expected in zip(stored_messages, expected_messages):
if _display_message_key(stored) != _display_message_key(expected):
return False
if _display_message_json(stored) != _display_message_json(expected):
return False
return True
def _rebuild_agent_display_messages(
conn: sqlite3.Connection,
session_id: str,
display_messages: tuple[JSONDict, ...],
*,
existing_rows: list[sqlite3.Row],
now: float,
) -> None:
snapshot_messages = [
message for message in display_messages if isinstance(message, dict)
]
snapshot_keys = {_display_message_key(message) for message in snapshot_messages}
persistent_by_anchor: dict[int, list[JSONDict]] = {}
nonpersistent_seen = 0
for row in sorted(existing_rows, key=lambda item: int(item['seq'])):
try:
message = json.loads(row['message_json'])
except (TypeError, json.JSONDecodeError, KeyError):
continue
if not isinstance(message, dict):
continue
if _is_persistent_display_message(message):
if _display_message_key(message) in snapshot_keys:
continue
persistent_by_anchor.setdefault(nonpersistent_seen, []).append(message)
else:
nonpersistent_seen += 1
conn.execute(
'delete from agent_display_messages where session_id = ?',
(session_id,),
)
seq = 1
def insert_message(message: JSONDict) -> None:
nonlocal seq
conn.execute(
"""
insert into agent_display_messages (
session_id, seq, message_key, role, updated_at, message_json
)
values (?, ?, ?, ?, ?, ?)
""",
(
session_id,
seq,
_display_message_key(message),
str(message.get('role') or ''),
now,
_display_message_json(message),
),
)
seq += 1
for message in persistent_by_anchor.get(0, []):
insert_message(message)
for index, message in enumerate(snapshot_messages, start=1):
insert_message(message)
for persistent in persistent_by_anchor.get(index, []):
insert_message(persistent)
def _display_message_key(message: JSONDict) -> str:
message_id = message.get('message_id')
if isinstance(message_id, str) and message_id:
@@ -899,6 +1043,10 @@ def _display_message_key(message: JSONDict) -> str:
return f'hash:{digest}'
def _display_message_json(message: JSONDict) -> str:
return json.dumps(message, ensure_ascii=False, sort_keys=True)
def _is_replaceable_display_message(message: JSONDict) -> bool:
metadata = message.get('metadata')
if not isinstance(metadata, dict):