appQQQ commited on
Commit
e89c535
ยท
verified ยท
1 Parent(s): a61dbb4

chore: upload app/agents/graph.py

Browse files
Files changed (1) hide show
  1. app/agents/graph.py +124 -0
app/agents/graph.py ADDED
@@ -0,0 +1,124 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """LangGraph StateGraph ่ฃ…้… + ็ผ–่ฏ‘.
2
+
3
+ ่ฎพ่ฎก: ๅ›พๅช่ดŸ่ดฃ"ๆฃ€็ดขๅพช็Žฏ" (route โ†’ query_rewrite โ†’ retrieve โ†’ rerank โ†’ evaluate).
4
+ answer ๆตๅผ็”Ÿๆˆ็”ฑ chat ็ซฏ็‚น็›ดๆŽฅ้ฉฑๅŠจ (answer_node_stream + asyncio.Queue),
5
+ ่ฟ™ๆ ท SSE token ๆตไธไพ่ต– astream_events ็š„ๅคๆ‚ๆ€ง.
6
+
7
+ ๆต็จ‹ๅ›พ:
8
+
9
+ START
10
+ โ”‚
11
+ โ–ผ
12
+ route โ”€โ”€โ”€ direct โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ–บ END
13
+ โ”‚
14
+ โ–ผ
15
+ query_rewrite โ”€โ”€โ–บ retrieve โ”€โ”€โ–บ rerank โ”€โ”€โ–บ evaluate
16
+ โ”‚
17
+ โ”Œโ”€โ”€โ”€โ”€ (needs_more, iter<max) โ”€โ”€โ”€โ”
18
+ โ”‚ โ”‚
19
+ โ””โ”€โ”€โ”€โ”€ (relevant/irrelevant/max) โ”˜
20
+ โ”‚
21
+ โ–ผ
22
+ END
23
+ """
24
+ from __future__ import annotations
25
+
26
+ import logging
27
+
28
+ from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
29
+ from langgraph.graph import END, START, StateGraph
30
+
31
+ from app.agents.nodes import (
32
+ evaluate_node,
33
+ query_rewrite_node,
34
+ rerank_node,
35
+ retrieve_node,
36
+ route_node,
37
+ )
38
+ from app.agents.state import AgentState
39
+ from app.config import settings
40
+
41
+ logger = logging.getLogger(__name__)
42
+
43
+
44
+ # ========== ่พนๅ†ณ็ญ– ==========
45
+ def _after_route(state: AgentState) -> str:
46
+ decision = state.get("route_decision", "retrieve")
47
+ if decision == "direct":
48
+ return END
49
+ return "query_rewrite"
50
+
51
+
52
+ def _after_evaluate(state: AgentState) -> str:
53
+ """evaluate ๅŽ: needs_more + iter<max -> ๅ›ž retrieve; ๅฆๅˆ™ END."""
54
+ if state.get("needs_more_retrieval", False):
55
+ return "query_rewrite"
56
+ return END
57
+
58
+
59
+ # ========== ็ผ–่ฏ‘ ==========
60
+ def _build_graph() -> StateGraph:
61
+ g = StateGraph(AgentState)
62
+
63
+ g.add_node("route", route_node)
64
+ g.add_node("query_rewrite", query_rewrite_node)
65
+ g.add_node("retrieve", retrieve_node)
66
+ g.add_node("rerank", rerank_node)
67
+ g.add_node("evaluate", evaluate_node)
68
+
69
+ g.add_edge(START, "route")
70
+ g.add_conditional_edges(
71
+ "route", _after_route,
72
+ {END: END, "query_rewrite": "query_rewrite"},
73
+ )
74
+ g.add_edge("query_rewrite", "retrieve")
75
+ g.add_edge("retrieve", "rerank")
76
+ g.add_edge("rerank", "evaluate")
77
+ g.add_conditional_edges(
78
+ "evaluate", _after_evaluate,
79
+ {END: END, "query_rewrite": "query_rewrite"},
80
+ )
81
+
82
+ return g
83
+
84
+
85
+ # ========== ็ผ–่ฏ‘ไบง็‰ฉ (ๅธฆ checkpointer) ==========
86
+ _compiled = None
87
+ _saver_cm = None
88
+ _compiled_loop = None
89
+
90
+
91
+ async def get_compiled_graph():
92
+ """ๆ‡’ๅŠ ่ฝฝ + ๅ•ไพ‹ (ๆŒ‰ event loop ้š”็ฆป)."""
93
+ global _compiled, _saver_cm, _compiled_loop
94
+ import asyncio
95
+ current_loop = asyncio.get_running_loop()
96
+
97
+ if _compiled is not None and _compiled_loop is current_loop:
98
+ return _compiled
99
+
100
+ if _compiled is not None:
101
+ await close_checkpointer()
102
+
103
+ g = _build_graph()
104
+ _saver_cm = AsyncSqliteSaver.from_conn_string(str(settings.langgraph_db_path))
105
+ saver = await _saver_cm.__aenter__()
106
+ _compiled = g.compile(checkpointer=saver)
107
+ _compiled_loop = current_loop
108
+ logger.info(
109
+ "LangGraph compiled with AsyncSqliteSaver at %s",
110
+ settings.langgraph_db_path,
111
+ )
112
+ return _compiled
113
+
114
+
115
+ async def close_checkpointer() -> None:
116
+ global _compiled, _saver_cm, _compiled_loop
117
+ _compiled = None
118
+ _compiled_loop = None
119
+ if _saver_cm is not None:
120
+ try:
121
+ await _saver_cm.__aexit__(None, None, None)
122
+ except Exception as e: # noqa: BLE001
123
+ logger.warning("AsyncSqliteSaver close failed: %s", e)
124
+ _saver_cm = None