TradeFlowAI / src /ai /graph.py
github-actions[bot]
Automated deployment from GitHub Actions: c077d743be852092402bf29515950ab5874e2735
1cf88ff
Raw
History Blame Contribute Delete
4.92 kB
"""
TradeFlow AI — LangGraph Extraction Graph (Step 2 Assembly)
PRD §10 — Full LangGraph pipeline:
preprocess → llm_extraction → [fallback if needed] → validate
→ risk_assessment → [interrupt if review needed] → DONE
The graph is compiled with a Redis checkpointer for persistence
and resumability across server restarts.
"""
from __future__ import annotations
import redis
import structlog
from langgraph.checkpoint.memory import MemorySaver
from langgraph.graph import END, StateGraph
from ..config import settings
from .nodes.extract import llm_extraction_node
from .nodes.fallback_ocr import fallback_ocr_node
from .nodes.human_review import human_review_node
from .nodes.preprocess import preprocess_documents_node
from .nodes.risk import risk_assessment_node
from .nodes.validate import validation_node
from .state import ExtractionGraphState
log = structlog.get_logger()
def _needs_fallback(state: ExtractionGraphState) -> str:
"""Route to OCR ensemble fallback when quality, confidence, or data is weak."""
if settings.CLOUD_LLM_ONLY:
log.info("CLOUD_LLM_ONLY is active — bypassing heavy OCR ensemble fallback")
return "validate"
for doc in state.get("documents", []):
if doc.get("error") or not doc.get("extracted_data"):
return "fallback"
if doc.get("quality_score", 1.0) < settings.OCR_FALLBACK_TRIGGER_QUALITY:
return "fallback"
if doc.get("document_mode") == "digital_pdf_text" and doc.get("ocr_method") == "digital_text_parser":
continue
confidences = doc.get("field_confidences") or {}
if confidences and min(confidences.values()) < settings.OCR_FALLBACK_TRIGGER_CONFIDENCE:
return "fallback"
if doc.get("ocr_conflicts"):
return "fallback"
if len(doc.get("ocr_candidates") or {}) > 1 and doc.get("document_mode") != "digital_pdf_text":
return "fallback"
return "validate"
def _needs_review(state: ExtractionGraphState) -> str:
"""Conditional edge: route to human review if flagged."""
if state.get("needs_human_review", False):
return "human_review"
return END
def build_extraction_graph() -> StateGraph:
"""Build and compile the LangGraph extraction pipeline."""
workflow = StateGraph(ExtractionGraphState)
# ── Add nodes ────────────────────────────────────────────────
workflow.add_node("preprocess", preprocess_documents_node)
workflow.add_node("llm_extraction", llm_extraction_node)
workflow.add_node("fallback_ocr", fallback_ocr_node)
workflow.add_node("validate", validation_node)
workflow.add_node("risk_assessment", risk_assessment_node)
workflow.add_node("human_review", human_review_node)
# ── Entry point ───────────────────────────────────────────────
workflow.set_entry_point("preprocess")
# ── Edges ─────────────────────────────────────────────────────
workflow.add_edge("preprocess", "llm_extraction")
# After extraction: check if fallback needed
workflow.add_conditional_edges(
"llm_extraction",
_needs_fallback,
{
"fallback": "fallback_ocr",
"validate": "validate",
},
)
# Fallback always proceeds to validate
workflow.add_edge("fallback_ocr", "validate")
# After validation: compute risk
workflow.add_edge("validate", "risk_assessment")
# After risk: check if human review needed
workflow.add_conditional_edges(
"risk_assessment",
_needs_review,
{
"human_review": "human_review",
END: END,
},
)
# After human review: graph ends (operator approved)
workflow.add_edge("human_review", END)
return workflow
def get_compiled_graph():
"""
Returns the compiled graph with Redis checkpointer.
The checkpointer enables:
- State persistence across Celery task restarts
- interrupt() resumability for human review
- LangSmith tracing integration
"""
workflow = build_extraction_graph()
# Using MemorySaver to support async ainvoke correctly
checkpointer = MemorySaver()
graph = workflow.compile(checkpointer=checkpointer, interrupt_before=["human_review"])
log.info("LangGraph extraction graph compiled", nodes=list(workflow.nodes.keys()))
return graph
# ── Singleton (imported by Celery tasks) ──────────────────────────────────────
extraction_graph = get_compiled_graph()