"""LangGraph workflows. Standard mode — router-driven Q&A: route -> retrieve -> (qa | analyst | verify) -> confidence Advanced mode — full verification pipeline: retrieve -> analyst -> verify -> risk -> report Both graphs share one state schema; nodes read/write only the keys they own. """ from __future__ import annotations import json import time from typing import TypedDict from langgraph.graph import END, StateGraph from src.agents import analyst_agent, qa_agent, router, verifier_agent from src.analysis import confidence as confidence_mod from src.analysis import risk as risk_mod from src.reporting import report as report_mod from src.retrieval.hybrid import HybridRetriever, RetrievedChunk class AnalysisState(TypedDict, total=False): # inputs question: str history: str doc_ids: list[str] # working state route: str retrieved: list[RetrievedChunk] tool_trace: list[dict] findings: list[dict] # evidence-matrix rows from the verifier verification_text: str risk: risk_mod.RiskAssessment ratio_results: list[dict] # outputs answer: str confidence: confidence_mod.ConfidenceReport report: str timings: dict[str, float] def _timed(state: AnalysisState, key: str, started: float) -> dict: timings = dict(state.get("timings") or {}) timings[key] = round(time.perf_counter() - started, 2) return timings def _balanced_retrieve(retriever: HybridRetriever, question: str, doc_ids: list[str] | None, per_doc: int = 5) -> list[RetrievedChunk]: """Retrieve top chunks from EVERY document separately, then merge. Plain top-k retrieval can return chunks from a single document, which leaves the verifier with nothing to cross-check. Balanced retrieval guarantees each uploaded document contributes evidence. """ ids = doc_ids or retriever.store.doc_ids() merged: dict[str, RetrievedChunk] = {} for doc_id in ids: for r in retriever.search(question, k=per_doc, doc_ids=[doc_id]): merged.setdefault(r.chunk.chunk_id, r) return sorted(merged.values(), key=lambda r: r.score, reverse=True) def build_standard_graph(retriever: HybridRetriever): def route_node(state: AnalysisState) -> dict: t0 = time.perf_counter() decided = router.route(state["question"]) return {"route": decided, "timings": _timed(state, "route", t0)} def retrieve_node(state: AnalysisState) -> dict: t0 = time.perf_counter() # verification/comparison questions need evidence from every document if state.get("route") in ("verification", "comparison"): retrieved = _balanced_retrieve(retriever, state["question"], state.get("doc_ids") or None, per_doc=4) else: retrieved = retriever.search(state["question"], k=6, doc_ids=state.get("doc_ids") or None) return {"retrieved": retrieved, "timings": _timed(state, "retrieve", t0)} def qa_node(state: AnalysisState) -> dict: t0 = time.perf_counter() answer = qa_agent.answer(state["question"], state["retrieved"], state.get("history", "")) return {"answer": answer, "timings": _timed(state, "qa_agent", t0)} def analyst_node(state: AnalysisState) -> dict: t0 = time.perf_counter() answer, trace = analyst_agent.answer(state["question"], state["retrieved"], state.get("history", "")) return {"answer": answer, "tool_trace": trace, "timings": _timed(state, "analyst_agent", t0)} def verify_node(state: AnalysisState) -> dict: t0 = time.perf_counter() findings = verifier_agent.verify(state["question"], state["retrieved"]) answer = verifier_agent.narrative(findings) return {"answer": answer, "findings": findings, "timings": _timed(state, "verifier_agent", t0)} def confidence_node(state: AnalysisState) -> dict: report = confidence_mod.score(state.get("answer", ""), state.get("retrieved", []), state.get("findings")) return {"confidence": report} def pick_agent(state: AnalysisState) -> str: return {"factual": "qa", "analysis": "analyst", "comparison": "analyst", "verification": "verify"}[state["route"]] g = StateGraph(AnalysisState) g.add_node("route", route_node) g.add_node("retrieve", retrieve_node) g.add_node("qa", qa_node) g.add_node("analyst", analyst_node) g.add_node("verify", verify_node) g.add_node("confidence", confidence_node) g.set_entry_point("route") g.add_edge("route", "retrieve") g.add_conditional_edges("retrieve", pick_agent, {"qa": "qa", "analyst": "analyst", "verify": "verify"}) g.add_edge("qa", "confidence") g.add_edge("analyst", "confidence") g.add_edge("verify", "confidence") g.add_edge("confidence", END) return g.compile() def build_advanced_graph(retriever: HybridRetriever): def retrieve_node(state: AnalysisState) -> dict: t0 = time.perf_counter() retrieved = retriever.search(state["question"], k=14, doc_ids=state.get("doc_ids") or None) return {"retrieved": retrieved, "timings": _timed(state, "retrieve", t0)} def analyst_node(state: AnalysisState) -> dict: t0 = time.perf_counter() answer, trace = analyst_agent.answer(state["question"], state["retrieved"]) ratio_results = [ json.loads(t["result"]) for t in trace if t["tool"] == "calculate_ratio" and t["result"].startswith("{") ] return {"answer": answer, "tool_trace": trace, "ratio_results": ratio_results, "timings": _timed(state, "analyst_agent", t0)} def verify_node(state: AnalysisState) -> dict: t0 = time.perf_counter() # re-retrieve balanced across documents so every doc is represented evidence = _balanced_retrieve(retriever, state["question"], state.get("doc_ids") or None, per_doc=5) findings = verifier_agent.verify(state["question"], evidence) text = verifier_agent.narrative(findings) return {"findings": findings, "verification_text": text, "timings": _timed(state, "verifier_agent", t0)} def risk_node(state: AnalysisState) -> dict: assessment = risk_mod.assess(state.get("findings", []), state.get("ratio_results", [])) return {"risk": assessment} def report_node(state: AnalysisState) -> dict: t0 = time.perf_counter() text = report_mod.generate( question=state["question"], analysis_text=state.get("answer", ""), verification_text=state.get("verification_text", ""), findings=state.get("findings", []), risk=state["risk"], ratio_results=state.get("ratio_results", []), doc_ids=state.get("doc_ids", []), ) conf = confidence_mod.score(state.get("answer", ""), state.get("retrieved", []), state.get("findings")) return {"report": text, "confidence": conf, "timings": _timed(state, "report", t0)} g = StateGraph(AnalysisState) g.add_node("retrieve", retrieve_node) g.add_node("analyst", analyst_node) g.add_node("verify", verify_node) g.add_node("risk", risk_node) g.add_node("report", report_node) g.set_entry_point("retrieve") g.add_edge("retrieve", "analyst") g.add_edge("analyst", "verify") g.add_edge("verify", "risk") g.add_edge("risk", "report") g.add_edge("report", END) return g.compile()