Refactor session runtime persistence

This commit is contained in:
wuyang6
2026-06-12 16:17:58 +08:00
parent 77d360c1e8
commit 953126e1d3
15 changed files with 717 additions and 253 deletions
+9 -1
View File
@@ -569,6 +569,7 @@ class LocalCodingAgent:
user_context=stored_session.user_context,
system_context=stored_session.system_context,
messages=stored_session.messages,
display_messages=stored_session.display_messages,
)
self._append_file_history_replay_if_needed(
session,
@@ -916,6 +917,7 @@ class LocalCodingAgent:
'continuation_index': continuation_count,
},
message_id=f'continuation_{turn_index}',
display=False,
)
stream_events.append(
{
@@ -1073,6 +1075,7 @@ class LocalCodingAgent:
'continuation_index': continuation_count,
},
message_id=f'continuation_{turn_index}',
display=False,
)
stream_events.append(
{
@@ -1444,6 +1447,7 @@ class LocalCodingAgent:
'plugin_preflight_count': len(plugin_preflight_messages),
},
message_id=f'plugin_tool_runtime_{tool_call.id}',
display=False,
)
stream_events.append(
{
@@ -3862,6 +3866,7 @@ class LocalCodingAgent:
'file_history_snapshot_count': snapshot_count,
},
message_id=f'file_history_replay_{replay_count}',
display=False,
)
def _render_file_history_replay(
@@ -3977,6 +3982,7 @@ class LocalCodingAgent:
message_id=(
f'compaction_replay_{len(compact_messages)}_{len(snipped_messages)}'
),
display=False,
)
def _render_compaction_replay(
@@ -4194,6 +4200,7 @@ class LocalCodingAgent:
'message_count': len(persist_messages),
},
message_id=f'plugin_persist_{result.session_id}',
display=False,
)
persist_events.append(
{
@@ -4239,7 +4246,8 @@ class LocalCodingAgent:
system_prompt_parts=session.system_prompt_parts,
user_context=dict(session.user_context),
system_context=dict(session.system_context),
messages=session.transcript(),
messages=session.model_transcript(),
display_messages=session.display_transcript(),
turns=previous_turns + result.turns,
tool_calls=previous_tool_calls + result.tool_calls,
usage=result.usage.to_dict(),
+128 -33
View File
@@ -102,8 +102,16 @@ class AgentSessionState:
user_context: dict[str, str] = field(default_factory=dict)
system_context: dict[str, str] = field(default_factory=dict)
messages: list[AgentMessage] = field(default_factory=list)
display_messages: list[AgentMessage] = field(default_factory=list)
mutation_serial: int = 0
def __post_init__(self) -> None:
# Backward compatibility for tests and older callers that construct an
# 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)
@classmethod
def create(
cls,
@@ -118,42 +126,48 @@ class AgentSessionState:
user_context=dict(user_context or {}),
system_context=dict(system_context or {}),
)
state.messages.append(
state._append_message(
AgentMessage(
role='system',
content='\n\n'.join(
_append_system_context(system_prompt_parts, state.system_context)
),
blocks=_text_blocks('\n\n'.join(_append_system_context(system_prompt_parts, state.system_context))),
message_id='system_0',
metadata=_initialize_message_metadata(
role='system',
message_id='system_0',
),
)
),
display=False,
)
if state.user_context:
state.messages.append(
state._append_message(
AgentMessage(
role='user',
content=_render_user_context_reminder(state.user_context),
blocks=_text_blocks(_render_user_context_reminder(state.user_context)),
message_id='user_context_0',
metadata=_initialize_message_metadata(
role='user',
message_id='user_context_0',
),
)
),
display=False,
)
if user_prompt is not None:
state.messages.append(
state._append_message(
AgentMessage(
role='user',
content=user_prompt,
blocks=_text_blocks(user_prompt),
message_id='user_0',
metadata=_initialize_message_metadata(
role='user',
message_id='user_0',
),
)
),
display=True,
)
return state
@@ -165,41 +179,47 @@ class AgentSessionState:
message_id: str | None = None,
stop_reason: str | None = None,
usage: UsageStats | None = None,
display: bool = True,
) -> None:
self.messages.append(
actual_message_id = message_id or f'assistant_{len(self.messages)}'
self._append_message(
AgentMessage(
role='assistant',
content=content,
tool_calls=tool_calls,
blocks=_assistant_blocks(content, tool_calls),
message_id=message_id,
message_id=actual_message_id,
stop_reason=stop_reason,
usage=usage or UsageStats(),
metadata=_initialize_message_metadata(
role='assistant',
message_id=message_id or f'assistant_{len(self.messages)}',
message_id=actual_message_id,
),
)
),
display=display,
)
def start_assistant(
self,
*,
message_id: str | None = None,
display: bool = True,
) -> int:
self.messages.append(
actual_message_id = message_id or f'assistant_{len(self.messages)}'
self._append_message(
AgentMessage(
role='assistant',
content='',
tool_calls=(),
blocks=(),
message_id=message_id,
message_id=actual_message_id,
state='streaming',
metadata=_initialize_message_metadata(
role='assistant',
message_id=message_id or f'assistant_{len(self.messages)}',
message_id=actual_message_id,
),
)
),
display=display,
)
return len(self.messages) - 1
@@ -214,12 +234,14 @@ class AgentSessionState:
mutation_serial=self._next_mutation_serial(),
)
merged_metadata = _advance_lineage_revision(merged_metadata)
self.messages[index] = replace(
updated = replace(
message,
content=message.content + delta,
blocks=_assistant_blocks(message.content + delta, message.tool_calls),
metadata=merged_metadata,
)
self.messages[index] = updated
self._replace_display_message(message, updated)
def merge_assistant_tool_call_delta(
self,
@@ -261,12 +283,14 @@ class AgentSessionState:
mutation_serial=self._next_mutation_serial(),
)
merged_metadata = _advance_lineage_revision(merged_metadata)
self.messages[index] = replace(
updated = replace(
message,
tool_calls=tuple(tool_calls),
blocks=_assistant_blocks(message.content, tuple(tool_calls)),
metadata=merged_metadata,
)
self.messages[index] = updated
self._replace_display_message(message, updated)
def finalize_assistant(
self,
@@ -285,7 +309,7 @@ class AgentSessionState:
mutation_serial=self._next_mutation_serial(),
)
merged_metadata = _advance_lineage_revision(merged_metadata)
self.messages[index] = replace(
updated = replace(
message,
state='final',
stop_reason=finish_reason,
@@ -293,6 +317,8 @@ class AgentSessionState:
blocks=_assistant_blocks(message.content, message.tool_calls),
metadata=merged_metadata,
)
self.messages[index] = updated
self._replace_display_message(message, updated)
def append_user(
self,
@@ -301,8 +327,10 @@ class AgentSessionState:
model_content: str | None = None,
metadata: dict[str, Any] | None = None,
message_id: str | None = None,
display: bool = True,
) -> None:
self.messages.append(
actual_message_id = message_id or f'user_{len(self.messages)}'
self._append_message(
AgentMessage(
role='user',
content=content,
@@ -310,27 +338,38 @@ class AgentSessionState:
blocks=_text_blocks(content),
metadata=_initialize_message_metadata(
role='user',
message_id=message_id or f'user_{len(self.messages)}',
message_id=actual_message_id,
metadata=dict(metadata or {}),
),
message_id=message_id,
)
message_id=actual_message_id,
),
display=display,
)
def append_tool(self, name: str, tool_call_id: str, content: str) -> None:
self.messages.append(
def append_tool(
self,
name: str,
tool_call_id: str,
content: str,
*,
display: bool = True,
) -> None:
actual_message_id = f'tool_{len(self.messages)}'
self._append_message(
AgentMessage(
role='tool',
content=content,
name=name,
tool_call_id=tool_call_id,
blocks=_tool_blocks(name, tool_call_id, content),
message_id=actual_message_id,
metadata=_initialize_message_metadata(
role='tool',
message_id=f'tool_{len(self.messages)}',
message_id=actual_message_id,
metadata={'tool_name': name, 'tool_call_id': tool_call_id},
),
)
),
display=display,
)
def start_tool(
@@ -340,26 +379,29 @@ class AgentSessionState:
tool_call_id: str,
message_id: str | None = None,
metadata: dict[str, Any] | None = None,
display: bool = True,
) -> int:
self.messages.append(
actual_message_id = message_id or f'tool_{len(self.messages)}'
self._append_message(
AgentMessage(
role='tool',
content='',
name=name,
tool_call_id=tool_call_id,
blocks=(),
message_id=message_id,
message_id=actual_message_id,
state='streaming',
metadata=_initialize_message_metadata(
role='tool',
message_id=message_id or f'tool_{len(self.messages)}',
message_id=actual_message_id,
metadata={
'tool_name': name,
'tool_call_id': tool_call_id,
**dict(metadata or {}),
},
),
)
),
display=display,
)
return len(self.messages) - 1
@@ -383,12 +425,14 @@ class AgentSessionState:
merged_metadata = _advance_lineage_revision(merged_metadata)
if metadata:
merged_metadata.update(metadata)
self.messages[index] = replace(
updated = replace(
message,
content=message.content + delta,
blocks=_tool_blocks(message.name, message.tool_call_id, message.content + delta),
metadata=merged_metadata,
)
self.messages[index] = updated
self._replace_display_message(message, updated)
def finalize_tool(
self,
@@ -413,7 +457,7 @@ class AgentSessionState:
merged_metadata = _advance_lineage_revision(merged_metadata)
if metadata:
merged_metadata.update(metadata)
self.messages[index] = replace(
updated = replace(
message,
content=content,
blocks=_tool_blocks(message.name, message.tool_call_id, content),
@@ -421,6 +465,8 @@ class AgentSessionState:
stop_reason=stop_reason,
metadata=merged_metadata,
)
self.messages[index] = updated
self._replace_display_message(message, updated)
def update_message(
self,
@@ -453,7 +499,7 @@ class AgentSessionState:
merged_metadata = _advance_lineage_revision(merged_metadata)
if metadata:
merged_metadata.update(metadata)
self.messages[index] = replace(
updated = replace(
message,
content=new_content,
blocks=_derive_blocks(
@@ -468,6 +514,8 @@ class AgentSessionState:
stop_reason=new_stop_reason,
metadata=merged_metadata,
)
self.messages[index] = updated
self._replace_display_message(message, updated)
def tombstone_message(
self,
@@ -491,8 +539,40 @@ class AgentSessionState:
return [message.to_openai_message() for message in self.messages]
def transcript(self) -> tuple[JSONDict, ...]:
return self.display_transcript()
def model_transcript(self) -> tuple[JSONDict, ...]:
return tuple(message.to_transcript_entry() for message in self.messages)
def display_transcript(self) -> tuple[JSONDict, ...]:
return tuple(message.to_transcript_entry() for message in self.display_messages)
def _append_message(self, message: AgentMessage, *, display: bool) -> None:
self.messages.append(message)
if display:
self.display_messages.append(message)
def _replace_display_message(
self,
previous_message: AgentMessage,
updated_message: AgentMessage,
) -> None:
index = self._find_display_message_index(previous_message)
if index is not None:
self.display_messages[index] = updated_message
def _find_display_message_index(self, message: AgentMessage) -> int | None:
if message.message_id:
for index in range(len(self.display_messages) - 1, -1, -1):
if self.display_messages[index].message_id == message.message_id:
return index
lineage_id = message.metadata.get('lineage_id')
if isinstance(lineage_id, str) and lineage_id:
for index in range(len(self.display_messages) - 1, -1, -1):
if self.display_messages[index].metadata.get('lineage_id') == lineage_id:
return index
return None
def _next_mutation_serial(self) -> int:
self.mutation_serial += 1
return self.mutation_serial
@@ -505,12 +585,27 @@ class AgentSessionState:
user_context: dict[str, str] | None,
system_context: dict[str, str] | None,
messages: tuple[JSONDict, ...] | list[JSONDict],
display_messages: tuple[JSONDict, ...] | list[JSONDict] | None = None,
) -> 'AgentSessionState':
model_messages = [
AgentMessage.from_openai_message(message)
for message in messages
if isinstance(message, dict)
]
display_source = messages if display_messages is None else display_messages
visible_messages = [
AgentMessage.from_openai_message(message)
for message in display_source
if isinstance(message, dict)
]
if display_messages is not None and not visible_messages:
visible_messages = list(model_messages)
return cls(
system_prompt_parts=tuple(system_prompt_parts),
user_context=dict(user_context or {}),
system_context=dict(system_context or {}),
messages=[AgentMessage.from_openai_message(message) for message in messages],
messages=model_messages,
display_messages=visible_messages,
mutation_serial=max(
(
int(message.get('metadata', {}).get('last_mutation_serial', 0))
+6 -1
View File
@@ -275,6 +275,11 @@ class RunStateStore:
continue
if isinstance(event, dict):
events.append(event)
elapsed_ms = row['elapsed_ms']
if row['status'] in ACTIVE_RUN_STATUSES:
started_at = row['started_at']
if isinstance(started_at, (int, float)):
elapsed_ms = max(0, int((time.time() - float(started_at)) * 1000))
return {
'run_id': row['run_id'],
'session_id': row['session_id'],
@@ -283,7 +288,7 @@ class RunStateStore:
'started_at': row['started_at'],
'updated_at': row['updated_at'],
'finished_at': row['finished_at'],
'elapsed_ms': row['elapsed_ms'],
'elapsed_ms': elapsed_ms,
'pending_prompt': row['pending_prompt'] or '',
'error': row['error'] or '',
'cancellable': bool(row['cancellable']),
+192 -9
View File
@@ -1,6 +1,8 @@
from __future__ import annotations
import json
import sqlite3
import time
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any
@@ -27,6 +29,7 @@ class StoredSession:
DEFAULT_SESSION_DIR = Path('.port_sessions')
DEFAULT_AGENT_SESSION_DIR = DEFAULT_SESSION_DIR / 'agent'
AGENT_SESSION_DB_FILENAME = 'sessions.db'
def save_session(session: StoredSession, directory: Path | None = None) -> Path:
@@ -70,6 +73,7 @@ class StoredAgentSession:
scratchpad_directory: str | None = None
is_training: bool = False
session_metadata: JSONDict | None = None
display_messages: tuple[JSONDict, ...] = ()
def save_agent_session(session: StoredAgentSession, directory: Path | None = None) -> Path:
@@ -78,16 +82,96 @@ def save_agent_session(session: StoredAgentSession, directory: Path | None = Non
session_dir = target_dir / session.session_id
session_dir.mkdir(parents=True, exist_ok=True)
path = session_dir / 'session.json'
path.write_text(json.dumps(asdict(session), indent=2), encoding='utf-8')
payload = asdict(session)
_write_agent_session_db_payload(target_dir, session.session_id, payload, path)
path.write_text(json.dumps(payload, indent=2, ensure_ascii=False), encoding='utf-8')
return path
def load_agent_session(session_id: str, directory: Path | None = None) -> StoredAgentSession:
target_dir = directory or DEFAULT_AGENT_SESSION_DIR
data = _read_agent_session_db_payload(target_dir, session_id)
if data is not None:
return _stored_agent_session_from_payload(data)
path = target_dir / session_id / 'session.json'
if not path.exists():
path = target_dir / f'{session_id}.json'
data = json.loads(path.read_text(encoding='utf-8'))
# Backfill SQLite on first read so future live UI reads do not depend on
# repeatedly parsing JSON snapshots.
_write_agent_session_db_payload(target_dir, session_id, data, path)
return _stored_agent_session_from_payload(data)
def list_agent_sessions(
directory: Path | None = None,
) -> list[tuple[StoredAgentSession, float]]:
target_dir = directory or DEFAULT_AGENT_SESSION_DIR
sessions: dict[str, tuple[StoredAgentSession, float]] = {}
for data, updated_at in _iter_agent_session_db_payloads(target_dir):
try:
stored = _stored_agent_session_from_payload(data)
except (KeyError, TypeError, ValueError):
continue
sessions[stored.session_id] = (stored, updated_at)
for path in _iter_agent_session_snapshot_files(target_dir):
session_id = _session_id_from_snapshot_path(path)
if session_id in sessions:
continue
try:
stored = load_agent_session(session_id, directory=target_dir)
except (FileNotFoundError, json.JSONDecodeError, KeyError, TypeError, ValueError):
continue
try:
updated_at = path.stat().st_mtime
except OSError:
updated_at = 0.0
sessions[stored.session_id] = (stored, updated_at)
return sorted(sessions.values(), key=lambda item: item[1], reverse=True)
def delete_agent_session(session_id: str, directory: Path | None = None) -> bool:
target_dir = directory or DEFAULT_AGENT_SESSION_DIR
deleted = False
with _connect_agent_session_db(target_dir) as conn:
cursor = conn.execute(
'delete from agent_sessions where session_id = ?',
(session_id,),
)
deleted = cursor.rowcount > 0
nested = target_dir / session_id
if nested.exists():
import shutil
shutil.rmtree(nested)
deleted = True
legacy = target_dir / f'{session_id}.json'
if legacy.exists():
legacy.unlink()
deleted = True
return deleted
def _stored_agent_session_from_payload(data: JSONDict) -> StoredAgentSession:
messages = tuple(
message for message in data['messages'] if isinstance(message, dict)
)
display_messages = tuple(
message
for message in data.get('display_messages', messages)
if isinstance(message, dict)
)
if not display_messages:
display_messages = messages
session_metadata = (
dict(data.get('session_metadata', {}))
if isinstance(data.get('session_metadata'), dict)
else {}
)
for key in ('title', 'title_source', 'title_updated_at'):
if key in data and key not in session_metadata:
session_metadata[key] = data[key]
return StoredAgentSession(
session_id=data['session_id'],
model_config=dict(data['model_config']),
@@ -95,9 +179,7 @@ def load_agent_session(session_id: str, directory: Path | None = None) -> Stored
system_prompt_parts=tuple(data['system_prompt_parts']),
user_context=dict(data['user_context']),
system_context=dict(data['system_context']),
messages=tuple(
message for message in data['messages'] if isinstance(message, dict)
),
messages=messages,
turns=int(data['turns']),
tool_calls=int(data['tool_calls']),
usage=dict(data.get('usage', {})),
@@ -121,14 +203,115 @@ def load_agent_session(session_id: str, directory: Path | None = None) -> Stored
else None
),
is_training=bool(data.get('is_training', False)),
session_metadata=(
dict(data.get('session_metadata', {}))
if isinstance(data.get('session_metadata'), dict)
else {}
),
session_metadata=session_metadata,
display_messages=display_messages,
)
def _write_agent_session_db_payload(
directory: Path,
session_id: str,
payload: JSONDict,
snapshot_path: Path,
) -> None:
updated_at = time.time()
payload_json = json.dumps(payload, ensure_ascii=False, sort_keys=True)
with _connect_agent_session_db(directory) as conn:
conn.execute(
"""
insert into agent_sessions (
session_id, updated_at, snapshot_path, payload_json
)
values (?, ?, ?, ?)
on conflict(session_id) do update set
updated_at=excluded.updated_at,
snapshot_path=excluded.snapshot_path,
payload_json=excluded.payload_json
""",
(session_id, updated_at, str(snapshot_path), payload_json),
)
def _read_agent_session_db_payload(
directory: Path,
session_id: str,
) -> JSONDict | None:
with _connect_agent_session_db(directory) as conn:
row = conn.execute(
'select payload_json from agent_sessions where session_id = ?',
(session_id,),
).fetchone()
if row is None:
return None
try:
payload = json.loads(row['payload_json'])
except (TypeError, json.JSONDecodeError):
return None
return payload if isinstance(payload, dict) else None
def _iter_agent_session_db_payloads(directory: Path) -> list[tuple[JSONDict, float]]:
with _connect_agent_session_db(directory) as conn:
rows = conn.execute(
"""
select payload_json, updated_at
from agent_sessions
order by updated_at desc
"""
).fetchall()
result: list[tuple[JSONDict, float]] = []
for row in rows:
try:
payload = json.loads(row['payload_json'])
except (TypeError, json.JSONDecodeError):
continue
if isinstance(payload, dict):
result.append((payload, float(row['updated_at'] or 0.0)))
return result
def _connect_agent_session_db(directory: Path) -> sqlite3.Connection:
directory.mkdir(parents=True, exist_ok=True)
conn = sqlite3.connect(directory.parent / AGENT_SESSION_DB_FILENAME, timeout=10)
conn.row_factory = sqlite3.Row
conn.execute('pragma journal_mode=wal')
conn.execute('pragma busy_timeout=5000')
conn.execute(
"""
create table if not exists agent_sessions (
session_id text primary key,
updated_at real not null,
snapshot_path text default '',
payload_json text not null
)
"""
)
conn.execute(
"""
create index if not exists idx_agent_sessions_updated_at
on agent_sessions(updated_at)
"""
)
return conn
def _iter_agent_session_snapshot_files(directory: Path) -> list[Path]:
if not directory.exists():
return []
by_session_id: dict[str, Path] = {}
for path in directory.glob('*.json'):
by_session_id[_session_id_from_snapshot_path(path)] = path
for path in directory.glob('*/session.json'):
by_session_id[_session_id_from_snapshot_path(path)] = path
return list(by_session_id.values())
def _session_id_from_snapshot_path(path: Path) -> str:
if path.name == 'session.json':
return path.parent.name
return path.stem
def serialize_model_config(model_config: ModelConfig) -> JSONDict:
return {
'model': model_config.model,