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