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
+247 -141
View File
@@ -37,7 +37,7 @@ from fastapi.responses import FileResponse, JSONResponse, Response, StreamingRes
from fastapi.staticfiles import StaticFiles
from pydantic import BaseModel, Field
from src.agent_session import AgentMessage, AgentSessionState
from src.agent_session import AgentSessionState
from src.agent_runtime import LocalCodingAgent
from src.agent_slash_commands import get_slash_command_specs
from src.agent_types import (
@@ -71,7 +71,9 @@ from src.run_state_store import ACTIVE_RUN_STATUSES, RunStateStore
from src.session_store import (
DEFAULT_AGENT_SESSION_DIR,
StoredAgentSession,
delete_agent_session,
deserialize_runtime_config,
list_agent_sessions,
load_agent_session,
save_agent_session,
serialize_model_config,
@@ -435,6 +437,8 @@ class RunManager:
normalized = _normalize_run_event(event)
if normalized is None:
return None
if normalized.get('type') in {'content_delta', 'tool_delta'}:
return None
normalized['recorded_at'] = time.time()
with self._lock:
record = self._runs.get(run_id)
@@ -1665,35 +1669,20 @@ def create_app(state: AgentState) -> FastAPI:
include_children: bool = False,
) -> list[dict[str, Any]]:
directory = state.account_paths(account_id)['sessions']
if not directory.exists():
return []
results: list[dict[str, Any]] = []
for path in sorted(_iter_session_files(directory), key=_session_mtime, reverse=True):
try:
data = json.loads(path.read_text(encoding='utf-8'))
except (OSError, json.JSONDecodeError):
continue
session_id = str(data.get('session_id', _session_id_from_path(path)))
session_metadata = (
data.get('session_metadata')
if isinstance(data.get('session_metadata'), dict)
else {}
)
for stored, updated_at in list_agent_sessions(directory):
session_id = stored.session_id
session_metadata = stored.session_metadata or {}
if (
not include_children
and session_metadata.get('visibility') == 'child'
):
continue
usage = data.get('usage') if isinstance(data.get('usage'), dict) else {}
model_config = (
data.get('model_config')
if isinstance(data.get('model_config'), dict)
else {}
)
messages = data.get('messages') or []
title = data.get('title')
title = session_metadata.get('title')
if not isinstance(title, str) or not title.strip():
title = None
preview = ''
for msg in messages:
for msg in _stored_display_messages(stored):
if isinstance(msg, dict) and msg.get('role') == 'user':
content = msg.get('content', '')
if isinstance(content, str) and not _is_internal_message(content):
@@ -1702,13 +1691,13 @@ def create_app(state: AgentState) -> FastAPI:
results.append(
{
'session_id': session_id,
'turns': data.get('turns', 0),
'tool_calls': data.get('tool_calls', 0),
'turns': stored.turns,
'tool_calls': stored.tool_calls,
'preview': title if isinstance(title, str) and title.strip() else preview,
'modified_at': _session_mtime(path),
'model': model_config.get('model'),
'usage': usage,
'is_training': bool(data.get('is_training', False)),
'modified_at': updated_at,
'model': stored.model_config.get('model'),
'usage': stored.usage,
'is_training': stored.is_training,
'session_metadata': session_metadata,
}
)
@@ -1730,7 +1719,7 @@ def create_app(state: AgentState) -> FastAPI:
async def delete_session(session_id: str, account_id: str | None = None) -> dict[str, Any]:
directory = state.account_paths(account_id)['sessions']
safe_id = _safe_session_id(session_id)
if safe_id is None or not _delete_session_files(directory, safe_id):
if safe_id is None or not delete_agent_session(safe_id, directory=directory):
raise HTTPException(status_code=404, detail='Session not found')
return {'deleted': True, 'session_id': safe_id}
@@ -2356,20 +2345,26 @@ def create_app(state: AgentState) -> FastAPI:
def _run_chat_payload(
request: ChatRequest,
event_sink: Any | None = None,
run_record: RunRecord | None = None,
) -> dict[str, Any]:
requested_session_id = _safe_session_id(
request.resume_session_id or request.session_id
) or uuid4().hex
account_key = state._account_key(request.account_id)
prompt = request.prompt.strip()
run_record = state.run_manager.start(account_key, requested_session_id, prompt)
state.run_state_store.start(
run_id=run_record.run_id,
account_key=account_key,
session_id=requested_session_id,
pending_prompt=prompt,
started_at=run_record.started_at,
)
if run_record is None:
run_record = state.run_manager.start(
account_key,
requested_session_id,
prompt,
)
state.run_state_store.start(
run_id=run_record.run_id,
account_key=account_key,
session_id=requested_session_id,
pending_prompt=prompt,
started_at=run_record.started_at,
)
run_lock = state.run_lock_for(request.account_id, requested_session_id)
agent = state.agent_for(request.account_id, requested_session_id)
config = state.config_for(request.account_id)
@@ -2627,6 +2622,53 @@ def create_app(state: AgentState) -> FastAPI:
)
return payload
def _start_chat_run(request: ChatRequest) -> dict[str, Any]:
requested_session_id = _safe_session_id(
request.resume_session_id or request.session_id
) or uuid4().hex
account_key = state._account_key(request.account_id)
prompt = request.prompt.strip()
run_record = state.run_manager.start(account_key, requested_session_id, prompt)
state.run_state_store.start(
run_id=run_record.run_id,
account_key=account_key,
session_id=requested_session_id,
pending_prompt=prompt,
started_at=run_record.started_at,
)
# Persist the submitted user message before the worker starts so the UI
# can immediately reload this exact session from DB without depending on
# browser-local optimistic state.
try:
agent = state.agent_for(request.account_id, requested_session_id)
_save_in_progress_session(
directory=state.account_paths(request.account_id)['sessions'],
agent=agent,
session_id=requested_session_id,
prompt=prompt,
)
except Exception:
# The worker will surface the real failure through run_state_store.
pass
def worker() -> None:
try:
_run_chat_payload(request, run_record=run_record)
except Exception as exc: # noqa: BLE001
print(
f'[chat-start] background run failed '
f'session={requested_session_id} run={run_record.run_id}: {exc}',
flush=True,
)
threading.Thread(target=worker, daemon=True).start()
return {
'session_id': requested_session_id,
'run_id': run_record.run_id,
'status': 'queued',
'started_at': run_record.started_at,
}
# ------------- chat ------------------------------------------------------
@app.post('/api/chat')
async def chat(request: ChatRequest) -> dict[str, Any]:
@@ -2651,6 +2693,24 @@ def create_app(state: AgentState) -> FastAPI:
)
return payload
@app.post('/api/chat/start')
async def chat_start(request: ChatRequest) -> dict[str, Any]:
prompt = request.prompt.strip()
if not prompt:
raise HTTPException(status_code=400, detail='Prompt is empty')
try:
return _start_chat_run(request)
except HTTPException:
raise
except Exception as exc:
return JSONResponse(
status_code=500,
content={
'error': str(exc),
'error_type': type(exc).__name__,
},
)
@app.post('/api/chat/stream')
async def chat_stream(request: ChatRequest) -> StreamingResponse:
prompt = request.prompt.strip()
@@ -2775,11 +2835,12 @@ def _truncate_api_text(text: str, limit: int) -> str:
def _serialize_stored_session(stored: StoredAgentSession) -> dict[str, Any]:
display_messages = _stored_display_messages(stored)
return {
'session_id': stored.session_id,
'turns': stored.turns,
'tool_calls': stored.tool_calls,
'messages': [_normalize_transcript_entry(dict(m)) for m in stored.messages],
'messages': [_normalize_transcript_entry(dict(m)) for m in display_messages],
'usage': stored.usage,
'total_cost_usd': stored.total_cost_usd,
'model': stored.model_config.get('model'),
@@ -2788,6 +2849,14 @@ def _serialize_stored_session(stored: StoredAgentSession) -> dict[str, Any]:
}
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)
def _stored_session_has_incomplete_tail(stored: StoredAgentSession) -> bool:
_, changed = _trim_incomplete_message_tail(stored.messages)
if changed:
@@ -2855,35 +2924,45 @@ def _mark_session_interrupted(
except (FileNotFoundError, OSError, json.JSONDecodeError):
return
trimmed, changed = _trim_incomplete_message_tail(stored.messages)
display_trimmed, display_changed = _trim_incomplete_message_tail(
_stored_display_messages(stored)
)
budget_state = (
dict(stored.budget_state)
if isinstance(stored.budget_state, dict)
else {}
)
if not changed and budget_state.get('status') not in {'queued', 'running'}:
if (
not changed
and not display_changed
and budget_state.get('status') not in {'queued', 'running'}
):
return
messages = list(trimmed)
display_messages = list(display_trimmed)
run_status_message = {
'role': 'assistant',
'content': _interrupted_session_message(status),
'state': 'final',
'stop_reason': status,
'metadata': {
'kind': 'run_status',
'status': status,
'detail': detail[:500],
'created_at_ms': int(time.time() * 1000),
},
}
if not _last_message_is_run_status(messages, status):
messages.append(
{
'role': 'assistant',
'content': _interrupted_session_message(status),
'state': 'final',
'stop_reason': status,
'metadata': {
'kind': 'run_status',
'status': status,
'detail': detail[:500],
'created_at_ms': int(time.time() * 1000),
},
}
)
messages.append(dict(run_status_message))
if not _last_message_is_run_status(display_messages, status):
display_messages.append(dict(run_status_message))
budget_state['status'] = status
budget_state['interrupted_at'] = int(time.time())
save_agent_session(
replace(
stored,
messages=tuple(messages),
display_messages=tuple(display_messages),
budget_state=budget_state,
),
directory=directory,
@@ -2900,7 +2979,16 @@ def _sanitize_stored_session_for_resume(
messages = _without_run_status_messages(trimmed)
if len(messages) != len(trimmed):
changed = True
if not changed:
display_trimmed, display_changed = _trim_incomplete_message_tail(
_stored_display_messages(stored)
)
if display_trimmed and _is_pending_user_message(display_trimmed[-1]):
display_trimmed = display_trimmed[:-1]
display_changed = True
display_messages = _without_run_status_messages(display_trimmed)
if len(display_messages) != len(display_trimmed):
display_changed = True
if not changed and not display_changed:
return stored
budget_state = (
dict(stored.budget_state)
@@ -2911,6 +2999,7 @@ def _sanitize_stored_session_for_resume(
return replace(
stored,
messages=messages,
display_messages=display_messages,
budget_state=budget_state,
)
@@ -3146,17 +3235,22 @@ def _save_in_progress_session(
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 = {
'role': 'user',
'content': prompt,
'state': 'final',
'metadata': pending_metadata,
'message_id': f'user_pending_{int(time.time() * 1000)}',
}
messages.append(
{
'role': 'user',
'content': prompt,
'state': 'final',
'metadata': pending_metadata,
'message_id': f'user_pending_{int(time.time() * 1000)}',
}
dict(pending_message)
)
display_messages.append(dict(pending_message))
budget_state = (
dict(stored.budget_state)
if isinstance(stored.budget_state, dict)
@@ -3168,6 +3262,7 @@ def _save_in_progress_session(
replace(
stored,
messages=tuple(messages),
display_messages=tuple(display_messages),
budget_state=budget_state,
),
directory=directory,
@@ -3188,6 +3283,15 @@ def _save_in_progress_session(
metadata=pending_metadata,
message_id='user_pending_0',
)
session_metadata: dict[str, Any] = {}
if initial_title:
session_metadata.update(
{
'title': initial_title,
'title_source': 'first_message',
'title_generated_at': int(time.time()),
}
)
stored = StoredAgentSession(
session_id=safe_id,
model_config=serialize_model_config(agent.model_config),
@@ -3195,7 +3299,8 @@ def _save_in_progress_session(
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=0,
tool_calls=0,
usage={},
@@ -3204,17 +3309,9 @@ def _save_in_progress_session(
budget_state={'status': 'running'},
plugin_state={},
scratchpad_directory=str(scratchpad_directory),
session_metadata=session_metadata,
)
saved_path = save_agent_session(stored, directory=directory)
if initial_title:
try:
data = json.loads(saved_path.read_text(encoding='utf-8'))
except (OSError, json.JSONDecodeError):
return initial_title
data['title'] = initial_title
data['title_source'] = 'first_message'
data['title_generated_at'] = int(time.time())
_write_session_metadata(saved_path, data)
save_agent_session(stored, directory=directory)
return initial_title
except OSError:
return None
@@ -3232,15 +3329,12 @@ def _derive_initial_session_title(prompt: str) -> str | None:
def _session_state_from_stored(stored: StoredAgentSession) -> AgentSessionState:
return AgentSessionState(
return AgentSessionState.from_persisted(
system_prompt_parts=stored.system_prompt_parts,
user_context=stored.user_context,
system_context=stored.system_context,
messages=[
AgentMessage.from_openai_message(message)
for message in stored.messages
if isinstance(message, dict)
],
messages=stored.messages,
display_messages=stored.display_messages,
)
@@ -3432,32 +3526,39 @@ def _annotate_last_assistant_elapsed(
except (OSError, json.JSONDecodeError):
return
messages = data.get('messages')
if not isinstance(messages, list):
display_messages = data.get('display_messages')
if not isinstance(display_messages, list):
display_messages = messages
if not isinstance(messages, list) and not isinstance(display_messages, list):
return
for message in reversed(messages):
if not isinstance(message, dict) or message.get('role') != 'assistant':
for message_list in (messages, display_messages):
if not isinstance(message_list, list):
continue
metadata = message.get('metadata')
if not isinstance(metadata, dict):
metadata = {}
message['metadata'] = metadata
metadata['elapsed_ms'] = elapsed_ms
try:
path.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding='utf-8')
except OSError:
return
for message in reversed(message_list):
if not isinstance(message, dict) or message.get('role') != 'assistant':
continue
metadata = message.get('metadata')
if not isinstance(metadata, dict):
metadata = {}
message['metadata'] = metadata
metadata['elapsed_ms'] = elapsed_ms
break
try:
path.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding='utf-8')
except OSError:
return
return
def _read_session_title(directory: Path, session_id: str | None) -> str | None:
if not session_id:
return None
path = _session_json_path(directory, session_id)
try:
data = json.loads(path.read_text(encoding='utf-8'))
except (OSError, json.JSONDecodeError):
stored = load_agent_session(session_id, directory=directory)
except FileNotFoundError:
return None
title = data.get('title')
metadata = stored.session_metadata or {}
title = metadata.get('title')
if isinstance(title, str) and title.strip():
return title.strip()
return None
@@ -3473,38 +3574,37 @@ def _ensure_session_title(
previous_title: str | None = None,
fallback_title: str | None = None,
) -> None:
path = _session_json_path(directory, session_id)
try:
data = json.loads(path.read_text(encoding='utf-8'))
except (OSError, json.JSONDecodeError):
stored = load_agent_session(session_id, directory=directory)
except FileNotFoundError:
return
metadata = dict(stored.session_metadata or {})
if previous_title:
data['title'] = previous_title
_write_session_metadata(path, data)
metadata['title'] = previous_title
save_agent_session(
replace(stored, session_metadata=metadata),
directory=directory,
)
return
title = data.get('title')
title_source = data.get('title_source')
title = metadata.get('title')
title_source = metadata.get('title_source')
if (
isinstance(title, str)
and title.strip()
and title_source != 'first_message'
):
return
messages = data.get('messages')
if not isinstance(messages, list):
if fallback_title and not (isinstance(title, str) and title.strip()):
data['title'] = fallback_title
data['title_source'] = 'first_message'
data['title_generated_at'] = int(time.time())
_write_session_metadata(path, data)
return
messages = [dict(message) for message in _stored_display_messages(stored)]
user_messages = _session_user_messages(messages)
if not user_messages:
if fallback_title and not (isinstance(title, str) and title.strip()):
data['title'] = fallback_title
data['title_source'] = 'first_message'
data['title_generated_at'] = int(time.time())
_write_session_metadata(path, data)
metadata['title'] = fallback_title
metadata['title_source'] = 'first_message'
metadata['title_generated_at'] = int(time.time())
save_agent_session(
replace(stored, session_metadata=metadata),
directory=directory,
)
return
generated = _generate_session_title(
user_messages,
@@ -3514,15 +3614,21 @@ def _ensure_session_title(
)
if not generated:
if fallback_title and not (isinstance(title, str) and title.strip()):
data['title'] = fallback_title
data['title_source'] = 'first_message'
data['title_generated_at'] = int(time.time())
_write_session_metadata(path, data)
metadata['title'] = fallback_title
metadata['title_source'] = 'first_message'
metadata['title_generated_at'] = int(time.time())
save_agent_session(
replace(stored, session_metadata=metadata),
directory=directory,
)
return
data['title'] = generated
data['title_source'] = 'llm'
data['title_generated_at'] = int(time.time())
_write_session_metadata(path, data)
metadata['title'] = generated
metadata['title_source'] = 'llm'
metadata['title_generated_at'] = int(time.time())
save_agent_session(
replace(stored, session_metadata=metadata),
directory=directory,
)
def _session_user_messages(messages: list[Any]) -> list[str]:
@@ -5069,7 +5175,7 @@ def _admin_session_tool_call_count(data: dict[str, Any]) -> int:
file_history = data.get('file_history')
file_history_count = len(file_history) if isinstance(file_history, list) else 0
message_count = max(
_admin_count_tool_calls_in_items(data.get('messages')),
_admin_count_tool_calls_in_items(data.get('display_messages') or data.get('messages')),
_admin_count_tool_calls_in_items(data.get('turns')),
)
derived_count = max(file_history_count, message_count)
@@ -5303,30 +5409,30 @@ def _delete_session_files(directory: Path, session_id: str) -> bool:
def _update_session_title(directory: Path, session_id: str, title: str) -> bool:
path = _session_json_path(directory, session_id)
if not path.exists():
return False
try:
data = json.loads(path.read_text(encoding='utf-8'))
except (OSError, json.JSONDecodeError):
stored = load_agent_session(session_id, directory=directory)
except FileNotFoundError:
return False
data['title'] = title
data['title_source'] = 'manual'
data['title_updated_at'] = int(time.time())
_write_session_metadata(path, data)
metadata = dict(stored.session_metadata or {})
metadata['title'] = title
metadata['title_source'] = 'manual'
metadata['title_updated_at'] = int(time.time())
save_agent_session(
replace(stored, session_metadata=metadata),
directory=directory,
)
return True
def _update_session_training(directory: Path, session_id: str, is_training: bool) -> bool:
path = _session_json_path(directory, session_id)
if not path.exists():
return False
try:
data = json.loads(path.read_text(encoding='utf-8'))
except (OSError, json.JSONDecodeError):
stored = load_agent_session(session_id, directory=directory)
except FileNotFoundError:
return False
data['is_training'] = is_training
_write_session_metadata(path, data)
save_agent_session(
replace(stored, is_training=is_training),
directory=directory,
)
return True
@@ -7022,7 +7128,7 @@ def _last_assistant_text(sessions_dir: Path, session_id: str | None) -> str | No
data = json.loads(path.read_text(encoding='utf-8'))
except (OSError, json.JSONDecodeError):
return None
messages = data.get('messages')
messages = data.get('display_messages') or data.get('messages')
if not isinstance(messages, list):
return None
for msg in reversed(messages):
@@ -7732,7 +7838,7 @@ def _extract_trigger_from_session(
data = json.loads(path.read_text(encoding='utf-8'))
except (OSError, json.JSONDecodeError):
return None
messages = data.get('messages')
messages = data.get('display_messages') or data.get('messages')
if not isinstance(messages, list):
return None
for msg in messages: