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))