fix: 修复 LangGraph Postgres checkpointer 初始化与 Windows 事件循环兼容
使用 psycopg_pool 作为 LangGraph checkpointer 连接池,并在启动时执行 setup() 确保表存在;Windows 下切换为 SelectorEventLoopPolicy 以兼容 psycopg 异步模式。
This commit is contained in:
parent
667b3bb878
commit
fd13607366
@ -1,6 +1,6 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "yuxi-know"
|
name = "yuxi-know"
|
||||||
version = "0.5.0.dev"
|
version = "0.5.2"
|
||||||
description = "基于大模型的智能知识库与知识图谱智能体开发平台,融合了 RAG 技术与知识图谱技术,基于 LangGraph v1 + Vue.js + FastAPI + LightRAG 架构构建"
|
description = "基于大模型的智能知识库与知识图谱智能体开发平台,融合了 RAG 技术与知识图谱技术,基于 LangGraph v1 + Vue.js + FastAPI + LightRAG 架构构建"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
requires-python = ">=3.12,<3.14"
|
requires-python = ">=3.12,<3.14"
|
||||||
@ -72,6 +72,7 @@ dependencies = [
|
|||||||
"redis>=5.2.0",
|
"redis>=5.2.0",
|
||||||
"aioboto3>=13.0.0",
|
"aioboto3>=13.0.0",
|
||||||
"wcmatch>=8.0.0",
|
"wcmatch>=8.0.0",
|
||||||
|
"psycopg[binary,pool]>=3.3.3",
|
||||||
]
|
]
|
||||||
[tool.ruff]
|
[tool.ruff]
|
||||||
line-length = 120 # 代码最大行宽
|
line-length = 120 # 代码最大行宽
|
||||||
|
|||||||
@ -1,5 +1,16 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import time
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# 解决 Windows 下 psycopg 异步模式不支持 ProactorEventLoop 的问题
|
||||||
|
# 注意:这段代码必须放在应用的极早期,最好在导入 FastAPI 或初始化数据库之前
|
||||||
|
# ==============================================================================
|
||||||
|
if sys.platform == "win32":
|
||||||
|
# 把当前文件 (main.py) 的上一级的上一级 (即根目录 Yuxi-Know) 加入到 sys.path
|
||||||
|
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||||
|
asyncio.set_event_loop_policy(asyncio.WindowsSelectorEventLoopPolicy())
|
||||||
|
|
||||||
from collections import defaultdict, deque
|
from collections import defaultdict, deque
|
||||||
|
|
||||||
import uvicorn
|
import uvicorn
|
||||||
@ -124,4 +135,12 @@ app.add_middleware(LoginRateLimitMiddleware)
|
|||||||
app.add_middleware(AuthMiddleware)
|
app.add_middleware(AuthMiddleware)
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
uvicorn.run(app, host="0.0.0.0", port=5050, threads=10, workers=10, reload=True)
|
# uvicorn.run(app, host="0.0.0.0", port=5050, threads=10, workers=10, reload=True)
|
||||||
|
|
||||||
|
uvicorn.run(
|
||||||
|
"server.main:app",
|
||||||
|
host="0.0.0.0",
|
||||||
|
port=5050,
|
||||||
|
reload=True,
|
||||||
|
reload_dirs=["server", "src"],
|
||||||
|
)
|
||||||
@ -1,10 +1,11 @@
|
|||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
|
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
|
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
|
||||||
|
|
||||||
from src.services.task_service import tasker
|
|
||||||
from src.services.mcp_service import init_mcp_servers
|
from src.services.mcp_service import init_mcp_servers
|
||||||
from src.services.run_queue_service import close_queue_clients, get_redis_client
|
from src.services.run_queue_service import close_queue_clients, get_redis_client
|
||||||
|
from src.services.task_service import tasker
|
||||||
from src.storage.postgres.manager import pg_manager
|
from src.storage.postgres.manager import pg_manager
|
||||||
from src.knowledge import knowledge_base
|
from src.knowledge import knowledge_base
|
||||||
from src.sandbox import init_sandbox_provider, shutdown_sandbox_provider
|
from src.sandbox import init_sandbox_provider, shutdown_sandbox_provider
|
||||||
@ -47,6 +48,13 @@ async def lifespan(app: FastAPI):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to initialize sandbox provider during startup: {e}")
|
logger.error(f"Failed to initialize sandbox provider during startup: {e}")
|
||||||
|
|
||||||
|
# =========================================================
|
||||||
|
# 2. 核心修复:在这里执行一次 setup(),建完表就拉倒
|
||||||
|
# =========================================================
|
||||||
|
checkpointer = AsyncPostgresSaver(pg_manager.langgraph_pool)
|
||||||
|
await checkpointer.setup()
|
||||||
|
print("LangGraph Checkpoint tables verified/created!")
|
||||||
|
|
||||||
await tasker.start()
|
await tasker.start()
|
||||||
yield
|
yield
|
||||||
await tasker.shutdown()
|
await tasker.shutdown()
|
||||||
|
|||||||
@ -1,5 +1,14 @@
|
|||||||
"""ARQ worker entrypoint."""
|
"""ARQ worker entrypoint."""
|
||||||
|
import asyncio
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
|
# 必须放在最顶层!
|
||||||
|
if sys.platform == "win32":
|
||||||
|
# 把当前文件 (main.py) 的上一级的上一级 (即根目录 Yuxi-Know) 加入到 sys.path
|
||||||
|
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||||
|
asyncio.set_event_loop_policy(asyncio.WindowsSelectorEventLoopPolicy())
|
||||||
|
|
||||||
from src.services.run_worker import WorkerSettings
|
from src.services.run_worker import WorkerSettings
|
||||||
|
|
||||||
__all__ = ["WorkerSettings"]
|
__all__ = ["WorkerSettings"]
|
||||||
@ -12,6 +12,7 @@ from langgraph.graph.state import CompiledStateGraph
|
|||||||
|
|
||||||
from src import config as sys_config
|
from src import config as sys_config
|
||||||
from src.agents.common.context import BaseContext
|
from src.agents.common.context import BaseContext
|
||||||
|
from src.storage.postgres.manager import pg_manager
|
||||||
from src.utils import logger
|
from src.utils import logger
|
||||||
|
|
||||||
|
|
||||||
@ -199,19 +200,9 @@ class BaseAgent:
|
|||||||
logger.warning(f"langgraph postgres checkpointer 不可用,回退 sqlite: {e}")
|
logger.warning(f"langgraph postgres checkpointer 不可用,回退 sqlite: {e}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
conn_str = postgres_url.replace("+asyncpg", "")
|
|
||||||
try:
|
try:
|
||||||
saver_factory = getattr(AsyncPostgresSaver, "from_conn_string", None)
|
saver = AsyncPostgresSaver(pg_manager.langgraph_pool)
|
||||||
if callable(saver_factory):
|
|
||||||
saver = saver_factory(conn_str)
|
|
||||||
else:
|
|
||||||
saver = AsyncPostgresSaver(conn_str) # type: ignore[call-arg]
|
|
||||||
|
|
||||||
setup_fn = getattr(saver, "setup", None)
|
|
||||||
if callable(setup_fn):
|
|
||||||
result = setup_fn()
|
|
||||||
if hasattr(result, "__await__"):
|
|
||||||
await result
|
|
||||||
logger.info(f"{self.name} 使用 postgres checkpointer")
|
logger.info(f"{self.name} 使用 postgres checkpointer")
|
||||||
return saver
|
return saver
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@ -4,6 +4,7 @@ import json
|
|||||||
import os
|
import os
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
|
|
||||||
|
from psycopg_pool import AsyncConnectionPool
|
||||||
from sqlalchemy import text
|
from sqlalchemy import text
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||||
from sqlalchemy.orm import declarative_base
|
from sqlalchemy.orm import declarative_base
|
||||||
@ -33,6 +34,7 @@ class PostgresManager(metaclass=SingletonMeta):
|
|||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.async_engine = None
|
self.async_engine = None
|
||||||
self.AsyncSession = None
|
self.AsyncSession = None
|
||||||
|
self.langgraph_pool = None
|
||||||
self._initialized = False
|
self._initialized = False
|
||||||
|
|
||||||
def initialize(self):
|
def initialize(self):
|
||||||
@ -65,6 +67,21 @@ class PostgresManager(metaclass=SingletonMeta):
|
|||||||
expire_on_commit=False,
|
expire_on_commit=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ==========================================
|
||||||
|
# 2. 为 LangGraph 专门初始化一个原生 psycopg_pool
|
||||||
|
# ==========================================
|
||||||
|
# ⚠️ 注意:psycopg 不认识 "+asyncpg" 这样的 SQLAlchemy 方言标识。
|
||||||
|
# 如果你的 db_url 是 "postgresql+asyncpg://user:pwd@host/db",
|
||||||
|
# 需要把它清洗成标准的 "postgresql://user:pwd@host/db"
|
||||||
|
langgraph_db_url = db_url.replace("+asyncpg", "").replace("+psycopg", "")
|
||||||
|
|
||||||
|
# 创建 LangGraph 专属连接池
|
||||||
|
self.langgraph_pool = AsyncConnectionPool(
|
||||||
|
conninfo=langgraph_db_url,
|
||||||
|
max_size=10, # 根据你的 Agent 并发情况设置,通常 5-10 足够了
|
||||||
|
kwargs={"autocommit": True} # LangGraph Checkpoint 强依赖 autocommit
|
||||||
|
)
|
||||||
|
|
||||||
self._initialized = True
|
self._initialized = True
|
||||||
logger.info(f"PostgreSQL manager initialized for knowledge base: {db_url.split('@')[0]}://***")
|
logger.info(f"PostgreSQL manager initialized for knowledge base: {db_url.split('@')[0]}://***")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@ -229,6 +246,9 @@ class PostgresManager(metaclass=SingletonMeta):
|
|||||||
if self.async_engine:
|
if self.async_engine:
|
||||||
await self.async_engine.dispose()
|
await self.async_engine.dispose()
|
||||||
|
|
||||||
|
if self.langgraph_pool:
|
||||||
|
await self.langgraph_pool.close()
|
||||||
|
|
||||||
async def async_check_first_run(self):
|
async def async_check_first_run(self):
|
||||||
"""检查是否首次运行(异步版本)- 检查用户表是否有数据"""
|
"""检查是否首次运行(异步版本)- 检查用户表是否有数据"""
|
||||||
from sqlalchemy import func, select
|
from sqlalchemy import func, select
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user