Team Ai
Apppublic

bhatnagarvikalp24/reporting_framework

sourceHugging Facemitupdated 8mo agoView on Hugging Face
0likes
graph.py354 linesDownload Raw Back to root
1import time2import uuid3from typing import TypedDict, Optional, Annotated4from langgraph.graph import StateGraph, END5 6from knowledge_base.kb_builder import query_knowledge_base, format_kb_context7from agents.router import classify_intent8from agents.sql_agent import generate_sql_with_retry9from agents.insight_agent import generate_insight, summarize_results10from agents.definition_agent import generate_definition11from masking import DataMasker12from audit import log_query, get_conversation_history13import pandas as pd14 15 16class AgentState(TypedDict):17    """State that flows through the LangGraph workflow."""18    # Input19    user_question: str20    session_id: str21 22    # Knowledge base23    kb_context: str24 25    # Router26    intent: str  # "sql", "diagnostic", "multi_step", "insight", or "definition"27 28    # SQL Agent29    sql_query: str30    sql_result: Optional[pd.DataFrame]31    masked_result: Optional[pd.DataFrame]32 33    # Insight Agent34    insight_response: str35    summary: str36 37    # Diagnostic Agent38    diagnostic_response: str39 40    # Definition Agent41    definition_response: str42 43    # Chart44    chart_type: str45    chart_config: Optional[dict]46 47    # Metadata48    error: str49    start_time: float50    audit_id: int51 52 53# Initialize masker as module-level singleton54_masker = DataMasker()55 56 57def enrich_from_kb(state: AgentState) -> dict:58    """Node: Enrich the question with actuarial knowledge base context."""59    question = state["user_question"]60    results = query_knowledge_base(question, n_results=10)61    kb_context = format_kb_context(results)62    return {"kb_context": kb_context, "start_time": time.time()}63 64 65def route_intent(state: AgentState) -> dict:66    """Node: Classify intent as sql or insight."""67    intent = classify_intent(state["user_question"], state.get("kb_context", ""))68    return {"intent": intent}69 70 71def sql_agent_node(state: AgentState) -> dict:72    """Node: Generate, validate, execute SQL query."""73    question = state["user_question"]74    kb_context = state.get("kb_context", "")75 76    sql, df, error = generate_sql_with_retry(question, kb_context)77 78    if error:79        return {"sql_query": sql, "sql_result": None, "error": error}80 81    # Mask PII in results82    query_id = str(uuid.uuid4())83    masked_df = _masker.mask_dataframe(df, query_id)84 85    # Generate summary86    summary = summarize_results(question, sql, df, kb_context)87 88    # Auto-detect chart type89    chart_type, chart_config = _detect_chart_type(df)90 91    return {92        "sql_query": sql,93        "sql_result": df,94        "masked_result": masked_df,95        "summary": summary,96        "chart_type": chart_type,97        "chart_config": chart_config,98        "error": "",99    }100 101 102def insight_agent_node(state: AgentState) -> dict:103    """Node: Generate analytical insight."""104    question = state["user_question"]105    kb_context = state.get("kb_context", "")106    session_id = state.get("session_id", "")107 108    conversation_history = get_conversation_history(session_id) if session_id else []109 110    # If we have SQL results from a prior step, include them111    supporting_data = state.get("sql_result")112 113    insight = generate_insight(114        question=question,115        kb_context=kb_context,116        conversation_history=conversation_history,117        supporting_data=supporting_data,118    )119 120    return {"insight_response": insight, "error": ""}121 122 123def definition_agent_node(state: AgentState) -> dict:124    """Node: Generate actuarial definition or concept explanation."""125    question = state["user_question"]126    kb_context = state.get("kb_context", "")127 128    definition = generate_definition(question=question, kb_context=kb_context)129 130    return {"definition_response": definition, "error": ""}131 132 133def diagnostic_synthesizer_node(state: AgentState) -> dict:134    """Node: Synthesize diagnostic insights from SQL results using causal knowledge.135 136    Phase 2 lightweight version — uses the insight agent with a diagnostic framing.137    Will be replaced by a full Plan-and-Execute diagnostic agent in Phase 3.138    """139    from agents.diagnostic_synthesizer import generate_diagnostic_synthesis140 141    question = state["user_question"]142    kb_context = state.get("kb_context", "")143    sql_result = state.get("sql_result")144    session_id = state.get("session_id", "")145    conversation_history = get_conversation_history(session_id) if session_id else []146 147    diagnostic = generate_diagnostic_synthesis(148        question=question,149        kb_context=kb_context,150        supporting_data=sql_result,151        conversation_history=conversation_history,152    )153 154    return {"diagnostic_response": diagnostic, "error": ""}155 156 157def multi_step_insight_node(state: AgentState) -> dict:158    """Node: Generate analytical insight enriched with SQL results.159 160    Used by the multi_step path after sql_agent fetches the data.161    Provides executive-level commentary on the retrieved data.162    """163    question = state["user_question"]164    kb_context = state.get("kb_context", "")165    session_id = state.get("session_id", "")166    sql_result = state.get("sql_result")167 168    conversation_history = get_conversation_history(session_id) if session_id else []169 170    insight = generate_insight(171        question=question,172        kb_context=kb_context,173        conversation_history=conversation_history,174        supporting_data=sql_result,175    )176 177    return {"insight_response": insight, "error": ""}178 179 180def audit_and_respond(state: AgentState) -> dict:181    """Node: Log to audit and finalize response."""182    duration_ms = int((time.time() - state.get("start_time", time.time())) * 1000)183 184    findings = (185        state.get("summary", "")186        or state.get("insight_response", "")187        or state.get("diagnostic_response", "")188        or state.get("definition_response", "")189    )190 191    audit_id = log_query(192        session_id=state.get("session_id", ""),193        user_prompt=state["user_question"],194        intent=state.get("intent", ""),195        sql_query=state.get("sql_query", ""),196        row_count=len(state["sql_result"]) if state.get("sql_result") is not None else 0,197        key_findings=findings[:1000],  # Truncate for storage198        chart_type=state.get("chart_type", ""),199        agent_used=state.get("intent", ""),200        error=state.get("error", ""),201        duration_ms=duration_ms,202    )203 204    return {"audit_id": audit_id}205 206 207def _detect_chart_type(df: pd.DataFrame) -> tuple[str, Optional[dict]]:208    """Auto-detect the best chart type based on result shape."""209    if df.empty or len(df.columns) < 2:210        return "none", None211 212    numeric_cols = df.select_dtypes(include=["number"]).columns.tolist()213    non_numeric_cols = df.select_dtypes(exclude=["number"]).columns.tolist()214 215    # Single row = metric card216    if len(df) == 1:217        return "metric", {"values": df.iloc[0].to_dict()}218 219    # Time-series: has a year column + numeric220    year_cols = [c for c in df.columns if "year" in c.lower()]221    if year_cols and numeric_cols:222        return "line", {223            "x": year_cols[0],224            "y": numeric_cols[:3],  # Up to 3 metrics225            "title": f"Trend by {year_cols[0]}",226        }227 228    # Categorical + numeric = bar chart229    if non_numeric_cols and numeric_cols:230        return "bar", {231            "x": non_numeric_cols[0],232            "y": numeric_cols[0],233            "title": f"{numeric_cols[0]} by {non_numeric_cols[0]}",234        }235 236    # Two numeric columns = scatter237    if len(numeric_cols) >= 2:238        return "scatter", {239            "x": numeric_cols[0],240            "y": numeric_cols[1],241            "title": f"{numeric_cols[1]} vs {numeric_cols[0]}",242        }243 244    # Matrix/pivot = heatmap245    if len(numeric_cols) >= 3 and len(df) >= 3:246        return "heatmap", {247            "title": "Data Heatmap",248        }249 250    return "table", None251 252 253# --- Route function for conditional edges ---254def route_by_intent(state: AgentState) -> str:255    """Conditional edge: route to sql or insight agent based on intent."""256    return state.get("intent", "sql")257 258 259# --- Build the graph ---260def build_graph() -> StateGraph:261    """Construct the LangGraph workflow."""262    workflow = StateGraph(AgentState)263 264    # Add nodes265    workflow.add_node("enrich_kb", enrich_from_kb)266    workflow.add_node("router", route_intent)267    workflow.add_node("sql_agent", sql_agent_node)268    workflow.add_node("insight_agent", insight_agent_node)269    workflow.add_node("definition_agent", definition_agent_node)270    workflow.add_node("diagnostic_synthesizer", diagnostic_synthesizer_node)271    workflow.add_node("multi_step_insight", multi_step_insight_node)272    workflow.add_node("audit", audit_and_respond)273 274    # Set entry point275    workflow.set_entry_point("enrich_kb")276 277    # Edges278    workflow.add_edge("enrich_kb", "router")279 280    # Conditional routing — 5-way281    workflow.add_conditional_edges(282        "router",283        route_by_intent,284        {285            "sql": "sql_agent",286            "diagnostic": "sql_agent",      # diagnostic: SQL first, then synthesize287            "multi_step": "sql_agent",       # multi_step: SQL first, then insight288            "insight": "insight_agent",289            "definition": "definition_agent",290        },291    )292 293    # sql_agent routes based on original intent:294    #   - "sql" → audit (simple data query, done)295    #   - "diagnostic" → diagnostic_synthesizer (causal analysis)296    #   - "multi_step" → multi_step_insight (executive commentary)297    def route_after_sql(state: AgentState) -> str:298        intent = state.get("intent", "sql")299        if intent == "diagnostic":300            return "diagnostic_synthesizer"301        elif intent == "multi_step":302            return "multi_step_insight"303        return "audit"304 305    workflow.add_conditional_edges(306        "sql_agent",307        route_after_sql,308        {309            "audit": "audit",310            "diagnostic_synthesizer": "diagnostic_synthesizer",311            "multi_step_insight": "multi_step_insight",312        },313    )314 315    # Post-processing nodes → audit316    workflow.add_edge("insight_agent", "audit")317    workflow.add_edge("definition_agent", "audit")318    workflow.add_edge("diagnostic_synthesizer", "audit")319    workflow.add_edge("multi_step_insight", "audit")320 321    # Audit → END322    workflow.add_edge("audit", END)323 324    return workflow.compile()325 326 327# Module-level compiled graph328app_graph = build_graph()329 330 331def run_query(question: str, session_id: str = "") -> AgentState:332    """Run a user question through the full pipeline. Returns final state."""333    initial_state: AgentState = {334        "user_question": question,335        "session_id": session_id or str(uuid.uuid4()),336        "kb_context": "",337        "intent": "",338        "sql_query": "",339        "sql_result": None,340        "masked_result": None,341        "insight_response": "",342        "summary": "",343        "diagnostic_response": "",344        "definition_response": "",345        "chart_type": "",346        "chart_config": None,347        "error": "",348        "start_time": time.time(),349        "audit_id": 0,350    }351 352    result = app_graph.invoke(initial_state)353    return result354