File size: 5,532 Bytes
cdc87cb 9ddb890 d1ac4a8 cdc87cb 9ddb890 cdc87cb 9ddb890 cdc87cb 9ddb890 cdc87cb 9ddb890 cdc87cb d1ac4a8 cdc87cb 9ddb890 cdc87cb | 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 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 | """
StateGraph topology for the ERP procurement assistant.
Nodes (in src/graph/nodes/):
classifier β retriever β validator β [tool_caller] β response_generator β memory
Edges:
classifier β(general_chat)βββββββββββββββββββββββΊ response_generator
classifier β(factual / workflow)βββββββββββββββββΊ retriever
retriever βββββββββββββββββββββββββββββββββββββββΊ validator
validator β(Β¬relevant β§ attempts<MAX)ββββββββββββΊ retriever # retry loop
validator β(workflow OR ID pattern)ββββββββββββββΊ tool_caller
validator β(otherwise)βββββββββββββββββββββββββββΊ response_generator
tool_caller βββββββββββββββββββββββββββββββββββββΊ response_generator
response_generator ββββββββββββββββββββββββββββββΊ memory ββΊ END
"""
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
from typing import List, Optional
from langgraph.graph import END, START, StateGraph
from src.graph.config import GRAPH_RECURSION_LIMIT, MAX_RETRIEVAL_ATTEMPTS
from src.graph.nodes import (
call_tools,
classify_query,
generate,
grade_answer,
retrieve_chunks,
update_memory,
validate_chunks,
)
from src.graph.state import ProcurementState
from src.graph.tools import has_id_pattern
# βββββββββββββββββββ Conditional routing functions ββββββββββββββββββββββββββββ
def route_after_classifier(state: dict) -> str:
qtype = state.get("query_type")
if qtype == "general_chat":
return "response_generator"
return "retriever"
def route_after_validator(state: dict) -> str:
if not state.get("is_relevant", False) and state.get("retrieval_attempt", 0) < MAX_RETRIEVAL_ATTEMPTS:
return "retriever"
qtype = state.get("query_type")
query = state.get("query", "")
if qtype == "workflow_guidance" or has_id_pattern(query):
return "tool_caller"
return "response_generator"
def route_after_grader(state: dict) -> str:
if state.get("grader_decision") == "fail" and state.get("retrieval_attempt", 0) < MAX_RETRIEVAL_ATTEMPTS:
return "retriever"
return "memory"
# βββββββββββββββββββ Graph construction βββββββββββββββββββββββββββββββββββββββ
def _build_graph():
builder = StateGraph(ProcurementState)
builder.add_node("classifier", classify_query)
builder.add_node("retriever", retrieve_chunks)
builder.add_node("validator", validate_chunks)
builder.add_node("tool_caller", call_tools)
builder.add_node("response_generator", generate)
builder.add_node("grader", grade_answer)
builder.add_node("memory", update_memory)
builder.add_edge(START, "classifier")
builder.add_conditional_edges(
"classifier",
route_after_classifier,
{
"retriever": "retriever",
"response_generator": "response_generator",
},
)
builder.add_edge("retriever", "validator")
builder.add_conditional_edges(
"validator",
route_after_validator,
{
"retriever": "retriever",
"tool_caller": "tool_caller",
"response_generator": "response_generator",
},
)
builder.add_edge("tool_caller", "response_generator")
builder.add_edge("response_generator", "grader")
builder.add_conditional_edges(
"grader",
route_after_grader,
{
"retriever": "retriever",
"memory": "memory",
},
)
builder.add_edge("memory", END)
return builder.compile()
compiled_graph = _build_graph().with_config(recursion_limit=GRAPH_RECURSION_LIMIT)
# βββββββββββββββββββ Convenience entry point ββββββββββββββββββββββββββββββββββ
def run_graph(query: str, history: Optional[List[dict]] = None,
memory_summary: str = "", session_id: str = "default") -> dict:
"""Invoke the graph with a clean initial state."""
initial_state: dict = {
"query": query,
"original_query": query,
"history": history or [],
"memory_summary": memory_summary or "",
"session_id": session_id,
"retrieval_attempt": 0,
"trace": [],
}
return compiled_graph.invoke(initial_state)
if __name__ == "__main__":
# Smoke test β requires HF_TOKEN set and FAISS index built.
import json
result = run_graph("What is a Purchase Requisition?")
print(json.dumps({
"query_type": result.get("query_type"),
"confidence": result.get("confidence"),
"answer": (result.get("answer") or "")[:200],
"sources": result.get("sources"),
"trace": [
{"node": t["node"], "status": t["status"],
"ms": t["duration_ms"], "summary": t["summary"]}
for t in result.get("trace", [])
],
}, indent=2))
|