Spaces:
Sleeping
Sleeping
| """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() | |