Refactor session runtime persistence
This commit is contained in:
@@ -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
@@ -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))
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user