finpy1789's picture
Download-report button; balanced per-doc retrieval for verification; robust JSON extraction for community models
025a350 verified
Raw
History Blame Contribute Delete
8.03 kB
"""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()