Coverage for agentos/storage/base.py: 54%
28 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 21:19 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 21:19 +0800
1"""
2AgentOS v0.20 持久化存储层。
3Base + SQLite实现,支持Checkpoint持久化。
4"""
6from __future__ import annotations
8import json
9import sqlite3
10import time
11from abc import ABC, abstractmethod
12from dataclasses import dataclass
14# ── 抽象基类 ────────────────────────────────────
17class CheckpointStore(ABC):
18 """检查点存储基类。"""
20 @abstractmethod
21 async def save(self, session_id: str, snapshot: dict): ...
22 @abstractmethod
23 async def load(self, session_id: str) -> dict | None: ...
24 @abstractmethod
25 async def delete(self, session_id: str): ...
26 @abstractmethod
27 async def list_sessions(self, limit: int = 50) -> list[str]: ...
30@dataclass
31class SqliteStore(CheckpointStore):
32 """SQLite 持久化存储。"""
34 path: str = ":memory:"
36 def __post_init__(self):
37 self._conn = sqlite3.connect(self.path, check_same_thread=False)
38 self._conn.execute("""CREATE TABLE IF NOT EXISTS checkpoints (
39 session_id TEXT PRIMARY KEY,
40 snapshot TEXT NOT NULL,
41 created_at REAL NOT NULL,
42 updated_at REAL NOT NULL
43 )""")
44 self._conn.execute("CREATE INDEX IF NOT EXISTS idx_updated ON checkpoints(updated_at DESC)")
45 self._conn.commit()
47 async def save(self, session_id: str, snapshot: dict):
48 now = time.time()
49 self._conn.execute(
50 """INSERT INTO checkpoints(session_id, snapshot, created_at, updated_at)
51 VALUES(?, ?, ?, ?)
52 ON CONFLICT(session_id) DO UPDATE SET
53 snapshot=excluded.snapshot, updated_at=excluded.updated_at""",
54 (session_id, json.dumps(snapshot, default=str), now, now),
55 )
56 self._conn.commit()
58 async def load(self, session_id: str) -> dict | None:
59 row = self._conn.execute(
60 "SELECT snapshot FROM checkpoints WHERE session_id=?", (session_id,)
61 ).fetchone()
62 return json.loads(row[0]) if row else None
64 async def delete(self, session_id: str):
65 self._conn.execute("DELETE FROM checkpoints WHERE session_id=?", (session_id,))
66 self._conn.commit()
68 async def list_sessions(self, limit: int = 50) -> list[str]:
69 rows = self._conn.execute(
70 "SELECT session_id FROM checkpoints ORDER BY updated_at DESC LIMIT ?", (limit,)
71 ).fetchall()
72 return [r[0] for r in rows]