ai-chatbot / app /agents /graph.py
appQQQ's picture
chore: upload app/agents/graph.py
e89c535 verified
Raw
History Blame
3.68 kB
"""LangGraph StateGraph ่ฃ…้… + ็ผ–่ฏ‘.
่ฎพ่ฎก: ๅ›พๅช่ดŸ่ดฃ"ๆฃ€็ดขๅพช็Žฏ" (route โ†’ query_rewrite โ†’ retrieve โ†’ rerank โ†’ evaluate).
answer ๆตๅผ็”Ÿๆˆ็”ฑ chat ็ซฏ็‚น็›ดๆŽฅ้ฉฑๅŠจ (answer_node_stream + asyncio.Queue),
่ฟ™ๆ ท SSE token ๆตไธไพ่ต– astream_events ็š„ๅคๆ‚ๆ€ง.
ๆต็จ‹ๅ›พ:
START
โ”‚
โ–ผ
route โ”€โ”€โ”€ direct โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ–บ END
โ”‚
โ–ผ
query_rewrite โ”€โ”€โ–บ retrieve โ”€โ”€โ–บ rerank โ”€โ”€โ–บ evaluate
โ”‚
โ”Œโ”€โ”€โ”€โ”€ (needs_more, iter<max) โ”€โ”€โ”€โ”
โ”‚ โ”‚
โ””โ”€โ”€โ”€โ”€ (relevant/irrelevant/max) โ”˜
โ”‚
โ–ผ
END
"""
from __future__ import annotations
import logging
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
from langgraph.graph import END, START, StateGraph
from app.agents.nodes import (
evaluate_node,
query_rewrite_node,
rerank_node,
retrieve_node,
route_node,
)
from app.agents.state import AgentState
from app.config import settings
logger = logging.getLogger(__name__)
# ========== ่พนๅ†ณ็ญ– ==========
def _after_route(state: AgentState) -> str:
decision = state.get("route_decision", "retrieve")
if decision == "direct":
return END
return "query_rewrite"
def _after_evaluate(state: AgentState) -> str:
"""evaluate ๅŽ: needs_more + iter<max -> ๅ›ž retrieve; ๅฆๅˆ™ END."""
if state.get("needs_more_retrieval", False):
return "query_rewrite"
return END
# ========== ็ผ–่ฏ‘ ==========
def _build_graph() -> StateGraph:
g = StateGraph(AgentState)
g.add_node("route", route_node)
g.add_node("query_rewrite", query_rewrite_node)
g.add_node("retrieve", retrieve_node)
g.add_node("rerank", rerank_node)
g.add_node("evaluate", evaluate_node)
g.add_edge(START, "route")
g.add_conditional_edges(
"route", _after_route,
{END: END, "query_rewrite": "query_rewrite"},
)
g.add_edge("query_rewrite", "retrieve")
g.add_edge("retrieve", "rerank")
g.add_edge("rerank", "evaluate")
g.add_conditional_edges(
"evaluate", _after_evaluate,
{END: END, "query_rewrite": "query_rewrite"},
)
return g
# ========== ็ผ–่ฏ‘ไบง็‰ฉ (ๅธฆ checkpointer) ==========
_compiled = None
_saver_cm = None
_compiled_loop = None
async def get_compiled_graph():
"""ๆ‡’ๅŠ ่ฝฝ + ๅ•ไพ‹ (ๆŒ‰ event loop ้š”็ฆป)."""
global _compiled, _saver_cm, _compiled_loop
import asyncio
current_loop = asyncio.get_running_loop()
if _compiled is not None and _compiled_loop is current_loop:
return _compiled
if _compiled is not None:
await close_checkpointer()
g = _build_graph()
_saver_cm = AsyncSqliteSaver.from_conn_string(str(settings.langgraph_db_path))
saver = await _saver_cm.__aenter__()
_compiled = g.compile(checkpointer=saver)
_compiled_loop = current_loop
logger.info(
"LangGraph compiled with AsyncSqliteSaver at %s",
settings.langgraph_db_path,
)
return _compiled
async def close_checkpointer() -> None:
global _compiled, _saver_cm, _compiled_loop
_compiled = None
_compiled_loop = None
if _saver_cm is not None:
try:
await _saver_cm.__aexit__(None, None, None)
except Exception as e: # noqa: BLE001
logger.warning("AsyncSqliteSaver close failed: %s", e)
_saver_cm = None