checkpointer.py 2.6 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374
  1. """
  2. 共享的 LangGraph checkpointer 工厂。
  3. 原实现使用 MemorySaver(纯内存),服务重启即丢失所有对话状态。
  4. 这里改用 SQLite 持久化(AsyncSqliteSaver),支持异步调用,对话状态落盘,重启不丢。
  5. 说明:LangGraph 官方 RedisSaver 依赖 RedisSearch 模块(FT.* 命令),
  6. 当前 Redis 未安装该模块,故采用 SQLite 持久化作为务实替代。
  7. 重要:AsyncSqliteSaver 内部绑定 asyncio.Lock + 当前事件循环(get_running_loop),
  8. 必须在 FastAPI 的 running loop 内初始化(app startup 时调用 init_checkpointer),
  9. 不能在临时 loop 中创建,否则后续 ainvoke 会报 "attached to a different loop"。
  10. """
  11. import logging
  12. import os
  13. import aiosqlite
  14. from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
  15. logger = logging.getLogger(__name__)
  16. # SQLite 数据库文件路径:
  17. # - Docker 容器内 /data 为持久化卷(与 Chroma 同级),但需检查写权限
  18. # - 本地开发 fallback 到项目 data/ 目录
  19. # - 可用环境变量 CHECKPOINT_DB_PATH 覆盖
  20. _DEFAULT_DB_PATH = os.path.join(
  21. os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
  22. "data",
  23. "checkpoints.db",
  24. )
  25. # 全局单例(由 init_checkpointer 在 startup 时填充)
  26. _saver: AsyncSqliteSaver | None = None
  27. def _resolve_db_path() -> str:
  28. """解析数据库路径:优先环境变量,其次有写权限的目录。"""
  29. # 1. 环境变量覆盖
  30. env_path = os.environ.get("CHECKPOINT_DB_PATH")
  31. if env_path:
  32. return env_path
  33. # 2. 尝试 /data(容器卷),需检查写权限
  34. if os.path.isdir("/data") and os.access("/data", os.W_OK):
  35. return "/data/checkpoints.db"
  36. # 3. fallback 到项目 data 目录
  37. return _DEFAULT_DB_PATH
  38. async def init_checkpointer() -> AsyncSqliteSaver:
  39. """在运行中的事件循环内初始化 AsyncSqliteSaver(app startup 时调用)。"""
  40. global _saver
  41. if _saver is not None:
  42. return _saver
  43. db_path = _resolve_db_path()
  44. os.makedirs(os.path.dirname(db_path), exist_ok=True)
  45. conn = await aiosqlite.connect(db_path)
  46. saver = AsyncSqliteSaver(conn)
  47. await saver.setup()
  48. _saver = saver
  49. logger.info("LangGraph AsyncSqliteSaver 已初始化: %s", db_path)
  50. return _saver
  51. def get_checkpointer() -> AsyncSqliteSaver:
  52. """获取全局 AsyncSqliteSaver(须已通过 init_checkpointer 在 startup 初始化)。"""
  53. if _saver is None:
  54. raise RuntimeError(
  55. "checkpointer 未初始化,请确保 app startup 已调用 init_checkpointer()"
  56. )
  57. return _saver