bhatnagarvikalp24/reporting_framework
0
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 