"""持久化保存会话运行状态。 这个模块只保存“运行事实”:当前 run 的状态、开始/结束时间、阶段说明和 活动事件。它不负责取消进程;取消仍由内存里的 RunManager 持有进程句柄来做。 前端刷新、切换会话或服务短暂重启后,可以通过这里恢复可靠的展示状态。 """ from __future__ import annotations import json import sqlite3 import time from pathlib import Path from threading import RLock from typing import Any ACTIVE_RUN_STATUSES = {'queued', 'running'} TERMINAL_RUN_STATUSES = {'completed', 'failed', 'cancelled', 'interrupted'} class RunStateStore: """SQLite backed run-state store. 数据库放在 `.port_sessions/run_state.db`,用 account_key/session_id 做隔离。 SQLite 足够支撑这个场景:写入量小,读多写少,并且可以避免前端回放依赖 后端进程内存。 """ def __init__(self, db_path: Path, *, max_events_per_run: int = 240) -> None: self.db_path = db_path self.max_events_per_run = max_events_per_run self._lock = RLock() def start( self, *, run_id: str, account_key: str, session_id: str, pending_prompt: str, started_at: float, status: str = 'queued', ) -> None: now = time.time() with self._connect() as conn: conn.execute( """ insert into run_states ( run_id, account_key, session_id, status, started_at, updated_at, pending_prompt, cancellable ) values (?, ?, ?, ?, ?, ?, ?, ?) on conflict(run_id) do update set account_key=excluded.account_key, session_id=excluded.session_id, status=excluded.status, started_at=excluded.started_at, updated_at=excluded.updated_at, pending_prompt=excluded.pending_prompt, cancellable=excluded.cancellable """, ( run_id, account_key, session_id, status, started_at, now, pending_prompt, 1 if status in ACTIVE_RUN_STATUSES else 0, ), ) def update( self, run_id: str, *, status: str | None = None, stage: str | None = None, error: str | None = None, cancellable: bool | None = None, elapsed_ms: int | None = None, finished_at: float | None = None, ) -> None: assignments: list[str] = ['updated_at = ?'] values: list[Any] = [time.time()] if status is not None: assignments.append('status = ?') values.append(status) if cancellable is None: cancellable = status in ACTIVE_RUN_STATUSES if stage is not None: assignments.append('current_stage = ?') values.append(stage) if error is not None: assignments.append('error = ?') values.append(error) if cancellable is not None: assignments.append('cancellable = ?') values.append(1 if cancellable else 0) if elapsed_ms is not None: assignments.append('elapsed_ms = ?') values.append(max(0, int(elapsed_ms))) if finished_at is not None: assignments.append('finished_at = ?') values.append(finished_at) values.append(run_id) with self._connect() as conn: conn.execute( f"update run_states set {', '.join(assignments)} where run_id = ?", values, ) def finish( self, run_id: str, *, status: str, elapsed_ms: int | None = None, stage: str | None = None, error: str | None = None, ) -> None: self.update( run_id, status=status, stage=stage, error=error, cancellable=False, elapsed_ms=elapsed_ms, finished_at=time.time(), ) def finish_active_for_session( self, account_key: str, session_id: str, *, status: str, stage: str | None = None, ) -> list[str]: now = time.time() with self._connect() as conn: rows = conn.execute( """ select run_id, started_at from run_states where account_key = ? and session_id = ? and status in ('queued', 'running') """, (account_key, session_id), ).fetchall() run_ids: list[str] = [] for row in rows: run_ids.append(row['run_id']) started_at = row['started_at'] elapsed_ms = ( max(0, int((now - float(started_at)) * 1000)) if isinstance(started_at, (int, float)) else None ) conn.execute( """ update run_states set status = ?, current_stage = coalesce(?, current_stage), cancellable = 0, elapsed_ms = coalesce(?, elapsed_ms), finished_at = ?, updated_at = ? where run_id = ? """, (status, stage, elapsed_ms, now, now, row['run_id']), ) return run_ids def record_event(self, run_id: str, event: dict[str, Any]) -> None: recorded_at = event.get('recorded_at') if not isinstance(recorded_at, (int, float)): recorded_at = time.time() event = {**event, 'recorded_at': recorded_at} event_json = json.dumps(event, ensure_ascii=False, sort_keys=True) with self._connect() as conn: conn.execute( """ insert into run_events (run_id, recorded_at, event_json) values (?, ?, ?) """, (run_id, float(recorded_at), event_json), ) conn.execute( """ delete from run_events where run_id = ? and id not in ( select id from run_events where run_id = ? order by id desc limit ? ) """, (run_id, run_id, self.max_events_per_run), ) def snapshot_latest(self, account_key: str, session_id: str) -> dict[str, Any] | None: with self._connect() as conn: row = conn.execute( """ select * from run_states where account_key = ? and session_id = ? order by updated_at desc, started_at desc limit 1 """, (account_key, session_id), ).fetchone() if row is None: return None return self._snapshot_from_row(conn, row) def snapshot_run(self, run_id: str) -> dict[str, Any] | None: with self._connect() as conn: row = conn.execute( "select * from run_states where run_id = ?", (run_id,), ).fetchone() if row is None: return None return self._snapshot_from_row(conn, row) def events_since( self, account_key: str, session_id: str, *, after_event_seq: int = 0, limit: int = 500, ) -> dict[str, Any]: with self._connect() as conn: rows = conn.execute( """ select e.id, e.recorded_at, e.event_json, s.run_id, s.status from run_events e join run_states s on s.run_id = e.run_id where s.account_key = ? and s.session_id = ? and e.id > ? order by e.id asc limit ? """, ( account_key, session_id, max(0, int(after_event_seq)), max(1, int(limit)), ), ).fetchall() latest_row = conn.execute( """ select max(e.id) as latest_event_seq from run_events e join run_states s on s.run_id = e.run_id where s.account_key = ? and s.session_id = ? """, (account_key, session_id), ).fetchone() events: list[dict[str, Any]] = [] for row in rows: try: event = json.loads(row['event_json']) except (TypeError, json.JSONDecodeError): continue if not isinstance(event, dict): continue events.append( { 'event_seq': int(row['id']), 'run_id': row['run_id'], 'run_status': row['status'], 'recorded_at': float(row['recorded_at'] or 0.0), 'event': event, } ) latest_event_seq = 0 if latest_row is not None and latest_row['latest_event_seq'] is not None: latest_event_seq = int(latest_row['latest_event_seq']) return { 'session_id': session_id, 'latest_event_seq': latest_event_seq, 'events': events, } def mark_interrupted_if_active( self, run_id: str, *, stage: str = '上次运行已中断,后台没有正在执行的进程', ) -> dict[str, Any] | None: snapshot = self.snapshot_run(run_id) if snapshot is None: return None if snapshot.get('status') not in ACTIVE_RUN_STATUSES: return snapshot started_at = snapshot.get('started_at') elapsed_ms = None if isinstance(started_at, (int, float)): elapsed_ms = max(0, int((time.time() - float(started_at)) * 1000)) self.finish( run_id, status='interrupted', elapsed_ms=elapsed_ms, stage=stage, ) return self.snapshot_run(run_id) def _snapshot_from_row( self, conn: sqlite3.Connection, row: sqlite3.Row, ) -> dict[str, Any]: event_rows = conn.execute( """ select event_json from run_events where run_id = ? order by id asc """, (row['run_id'],), ).fetchall() events: list[dict[str, Any]] = [] for event_row in event_rows: try: event = json.loads(event_row['event_json']) except (TypeError, json.JSONDecodeError): 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'], 'status': row['status'], 'current_stage': row['current_stage'] or '', 'started_at': row['started_at'], 'updated_at': row['updated_at'], 'finished_at': row['finished_at'], 'elapsed_ms': elapsed_ms, 'pending_prompt': row['pending_prompt'] or '', 'error': row['error'] or '', 'cancellable': bool(row['cancellable']), 'events': events, } def _connect(self) -> sqlite3.Connection: with self._lock: self.db_path.parent.mkdir(parents=True, exist_ok=True) conn = sqlite3.connect(self.db_path, timeout=10) conn.row_factory = sqlite3.Row self._init_schema(conn) return conn def _init_schema(self, conn: sqlite3.Connection) -> None: conn.execute('pragma journal_mode=wal') conn.execute( """ create table if not exists run_states ( run_id text primary key, account_key text not null, session_id text not null, status text not null, started_at real not null, updated_at real not null, finished_at real, elapsed_ms integer, pending_prompt text default '', current_stage text default '', error text default '', cancellable integer not null default 0 ) """ ) conn.execute( """ create index if not exists idx_run_states_session on run_states(account_key, session_id, updated_at) """ ) conn.execute( """ create table if not exists run_events ( id integer primary key autoincrement, run_id text not null, recorded_at real not null, event_json text not null ) """ ) conn.execute( """ create index if not exists idx_run_events_run_id on run_events(run_id, id) """ )