diff --git a/backend/api/server.py b/backend/api/server.py index 5bbbf00..17ac1fb 100644 --- a/backend/api/server.py +++ b/backend/api/server.py @@ -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(''): + 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, ), diff --git a/src/agent_runtime.py b/src/agent_runtime.py index 26b6d2e..96cb1d0 100644 --- a/src/agent_runtime.py +++ b/src/agent_runtime.py @@ -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, diff --git a/src/agent_session.py b/src/agent_session.py index f236687..f6930a5 100644 --- a/src/agent_session.py +++ b/src/agent_session.py @@ -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(''): + 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 {}), diff --git a/src/session_store.py b/src/session_store.py index 8e427a7..908ada2 100644 --- a/src/session_store.py +++ b/src/session_store.py @@ -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(''): + 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):