Spaces:
Runtime error
Runtime error
| """ | |
| LangGraph orchestrator: wires all nodes with conditional edges. | |
| Flow: | |
| START | |
| → router_node | |
| ├─ CHAT → generator_node → END | |
| └─ RAG → retriever_node → grader_node | |
| ├─ retry → retriever_node (max 2x) | |
| └─ generate → generator_node → END | |
| """ | |
| from __future__ import annotations | |
| from langgraph.graph import StateGraph, END | |
| from backend.agents.state import GraphState | |
| from backend.agents.nodes.router import router_node | |
| from backend.agents.nodes.retriever import retriever_node | |
| from backend.agents.nodes.grader import grader_node, should_retry | |
| from backend.agents.nodes.generator import generator_node | |
| def _route_after_router(state: GraphState) -> str: | |
| return state.get("route", "RAG") | |
| def _increment_retry(state: GraphState) -> dict: | |
| """Node nhỏ chỉ để tăng retry_count trước khi quay lại retriever.""" | |
| return {"retry_count": state.get("retry_count", 0) + 1} | |
| def build_graph() -> StateGraph: | |
| graph = StateGraph(GraphState) | |
| # Thêm nodes | |
| graph.add_node("router", router_node) | |
| graph.add_node("retriever", retriever_node) | |
| graph.add_node("grader", grader_node) | |
| graph.add_node("generator", generator_node) | |
| graph.add_node("increment_retry", _increment_retry) | |
| # Entry point | |
| graph.set_entry_point("router") | |
| # Router → CHAT hoặc RAG | |
| graph.add_conditional_edges( | |
| "router", | |
| _route_after_router, | |
| { | |
| "CHAT": "generator", | |
| "RAG": "retriever", | |
| }, | |
| ) | |
| # Retriever → Grader | |
| graph.add_edge("retriever", "grader") | |
| # Grader → Generate hoặc Retry | |
| graph.add_conditional_edges( | |
| "grader", | |
| should_retry, | |
| { | |
| "retry": "increment_retry", | |
| "generate": "generator", | |
| }, | |
| ) | |
| # Retry → quay lại retriever | |
| graph.add_edge("increment_retry", "retriever") | |
| # Generator → END | |
| graph.add_edge("generator", END) | |
| return graph.compile() | |
| # Singleton compiled graph | |
| _compiled_graph = None | |
| def get_graph(): | |
| global _compiled_graph | |
| if _compiled_graph is None: | |
| _compiled_graph = build_graph() | |
| return _compiled_graph | |