Coverage for agentos/storage/base.py: 54%

28 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 21:26 +0800

1""" 

2AgentOS v0.20 持久化存储层。 

3Base + SQLite实现,支持Checkpoint持久化。 

4""" 

5 

6from __future__ import annotations 

7 

8import json 

9import sqlite3 

10import time 

11from abc import ABC, abstractmethod 

12from dataclasses import dataclass 

13 

14# ── 抽象基类 ──────────────────────────────────── 

15 

16 

17class CheckpointStore(ABC): 

18 """检查点存储基类。""" 

19 

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]: ... 

28 

29 

30@dataclass 

31class SqliteStore(CheckpointStore): 

32 """SQLite 持久化存储。""" 

33 

34 path: str = ":memory:" 

35 

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() 

46 

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() 

57 

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 

63 

64 async def delete(self, session_id: str): 

65 self._conn.execute("DELETE FROM checkpoints WHERE session_id=?", (session_id,)) 

66 self._conn.commit() 

67 

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]