Spaces:
Sleeping
Sleeping
File size: 3,675 Bytes
e89c535 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 | """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
|