Spaces:
Sleeping
Sleeping
Mohitcr1
Add production features: circuit breaker, exponential backoff, semantic caching, intent-based scope handling
b6ae869 | from langgraph.graph import StateGraph, START, END | |
| from langgraph.checkpoint.memory import MemorySaver | |
| from src.state import AgentState | |
| from src.nodes.sanitizer import sanitize_input | |
| from src.nodes.entity_extractor import extract_entities | |
| from src.nodes.clarification import check_clarification_needed | |
| from src.nodes.intent_classifier import classify_intent | |
| from src.nodes.sql_generator import generate_sql | |
| from src.nodes.sql_safety import check_sql_safety | |
| from src.nodes.sql_executor import execute_sql | |
| from src.nodes.rag_retriever import retrieve_rag | |
| from src.nodes.sentiment_analyzer import analyze_sentiment | |
| from src.nodes.proactive_logic import check_proactive_logic | |
| from src.nodes.compensation import offer_compensation | |
| from src.nodes.escalation import handle_escalation | |
| from src.nodes.response_generator import generate_response | |
| # Routing Functions | |
| def route_after_sanitizer(state: AgentState) -> str: | |
| if not state.get("sanitized"): | |
| return "END" | |
| # If escalation detected in sanitizer, skip to escalation | |
| if state.get("escalation_required"): | |
| return "escalation" | |
| return "entity_extractor" | |
| def route_after_clarification(state: AgentState) -> str: | |
| # If clarification is needed, skip to END with clarification message | |
| if state.get("clarification_needed"): | |
| return "END" | |
| return "intent_classifier" | |
| def route_by_intent(state: AgentState) -> str: | |
| intent = state.get("intent", "hybrid") | |
| return { | |
| "transactional": "sql_generator", | |
| "informational": "rag_retriever", | |
| "sentimental": "sentiment_analyzer", | |
| "hybrid": "sql_generator", # Simplified: hybrid starts with SQL | |
| "out_of_scope": "END", | |
| }.get(intent, "sql_generator") | |
| def route_after_sql_executor(state: AgentState) -> str: | |
| sql = state.get("sql_result", {}) | |
| retry = state.get("retry_count", 0) | |
| # Retry if error and not maxed out | |
| if sql.get("error") and sql["error"] not in ("unsafe_query", "max_retries_exceeded", "no_query") and retry < 3: | |
| return "sql_generator" | |
| return "proactive_logic" | |
| def route_after_proactive(state: AgentState) -> str: | |
| if state.get("escalation_required"): | |
| return "escalation" | |
| if state.get("is_late_delivery") and not state.get("compensation_offered"): | |
| return "compensation" | |
| # For hybrid intent, also retrieve RAG if not done yet | |
| if state.get("intent") == "hybrid" and state.get("rag_result") is None: | |
| return "rag_retriever" | |
| return "response_generator" | |
| def route_after_sentiment(state: AgentState) -> str: | |
| if state.get("escalation_required"): | |
| return "escalation" | |
| return "response_generator" | |
| def route_after_rag(state: AgentState) -> str: | |
| # Check if escalation was set by the node | |
| if state.get("escalation_required"): | |
| return "escalation" | |
| return "response_generator" | |
| def route_after_compensation(state: AgentState) -> str: | |
| # After compensation, check if we need RAG for hybrid | |
| if state.get("intent") == "hybrid" and state.get("rag_result") is None: | |
| return "rag_retriever" | |
| return "response_generator" | |
| # Graph Construction | |
| def build_graph() -> StateGraph: | |
| builder = StateGraph(AgentState) | |
| # Add all nodes | |
| builder.add_node("sanitizer", sanitize_input) | |
| builder.add_node("entity_extractor", extract_entities) | |
| builder.add_node("clarification", check_clarification_needed) | |
| builder.add_node("intent_classifier", classify_intent) | |
| builder.add_node("sql_generator", generate_sql) | |
| builder.add_node("sql_safety", check_sql_safety) | |
| builder.add_node("sql_executor", execute_sql) | |
| builder.add_node("rag_retriever", retrieve_rag) | |
| builder.add_node("sentiment_analyzer", analyze_sentiment) | |
| builder.add_node("proactive_logic", check_proactive_logic) | |
| builder.add_node("compensation", offer_compensation) | |
| builder.add_node("escalation", handle_escalation) | |
| builder.add_node("response_generator", generate_response) | |
| # Add edges | |
| builder.add_edge(START, "sanitizer") | |
| builder.add_conditional_edges("sanitizer", route_after_sanitizer, | |
| {"entity_extractor": "entity_extractor", | |
| "escalation": "escalation", | |
| "END": END}) | |
| builder.add_edge("entity_extractor", "clarification") | |
| builder.add_conditional_edges("clarification", route_after_clarification, | |
| {"intent_classifier": "intent_classifier", "END": END}) | |
| builder.add_conditional_edges("intent_classifier", route_by_intent, | |
| {"sql_generator": "sql_generator", | |
| "rag_retriever": "rag_retriever", | |
| "sentiment_analyzer": "sentiment_analyzer", | |
| "END": END}) | |
| # Transactional path | |
| builder.add_edge("sql_generator", "sql_safety") | |
| builder.add_edge("sql_safety", "sql_executor") | |
| builder.add_conditional_edges("sql_executor", route_after_sql_executor, | |
| {"sql_generator": "sql_generator", | |
| "proactive_logic": "proactive_logic"}) | |
| builder.add_conditional_edges("proactive_logic", route_after_proactive, | |
| {"escalation": "escalation", | |
| "compensation": "compensation", | |
| "rag_retriever": "rag_retriever", | |
| "response_generator": "response_generator"}) | |
| # Informational path | |
| builder.add_conditional_edges("rag_retriever", route_after_rag, | |
| {"escalation": "escalation", | |
| "response_generator": "response_generator"}) | |
| # Sentimental path | |
| builder.add_conditional_edges("sentiment_analyzer", route_after_sentiment, | |
| {"escalation": "escalation", | |
| "response_generator": "response_generator"}) | |
| # Terminal nodes | |
| builder.add_conditional_edges("compensation", route_after_compensation, | |
| {"rag_retriever": "rag_retriever", | |
| "response_generator": "response_generator"}) | |
| builder.add_edge("escalation", END) | |
| builder.add_edge("response_generator", END) | |
| checkpointer = MemorySaver() | |
| return builder.compile(checkpointer=checkpointer) | |
| # Singleton | |
| graph = build_graph() | |