fix session display replay consistency
This commit is contained in:
+36
-9
@@ -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
@@ -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
@@ -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
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user