from __future__ import annotations import psycopg from langgraph.checkpoint.postgres import PostgresSaver from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer from langgraph.types import Command from sqlalchemy.orm import Session from app.core.config import DATABASE_URL from app.core.tracing.run_tracker import get_or_create_tracker, finish_tracker from app.services.workflow_service import WorkflowService from app.workflow.graph import build_workflow from app.workflow.state import WorkflowState class WorkflowExecutor: """ Executes the document workflow through LangGraph with durable PostgreSQL checkpoints. """ def __init__( self, db: Session, interrupt_before=None, ): self.workflow_service = WorkflowService(db) self.checkpoint_connection = psycopg.connect( DATABASE_URL, autocommit=True, ) serde = JsonPlusSerializer() self.checkpointer = PostgresSaver( self.checkpoint_connection, serde=serde, ) self.checkpointer.setup() self.graph = build_workflow( db, checkpointer=self.checkpointer, interrupt_before=interrupt_before, ) # ------------------------------------------------------------------ # Helpers # ------------------------------------------------------------------ @staticmethod def _is_completed(state) -> bool: """ LangGraph may return either a WorkflowState instance or a dictionary depending on the configured state schema and execution path. """ if isinstance(state, dict): return bool(state.get("completed", False)) return bool(getattr(state, "completed", False)) # ------------------------------------------------------------------ # Initial execution # ------------------------------------------------------------------ def execute( self, state: WorkflowState | None, workflow, ) -> WorkflowState: config = { "configurable": { "thread_id": str(workflow.id), } } # Start cost/time tracking tracker = get_or_create_tracker(str(workflow.id)) tracker.start_stage("total") final_state = self.graph.invoke( state, config=config, ) if self._is_completed(final_state): self.workflow_service.complete_workflow(workflow) report = finish_tracker(str(workflow.id)) self._persist_metrics(workflow, report) else: self.workflow_service.wait_for_review(workflow) # Don't finish tracker — run is paused, will resume later tracker.end_stage("total") # Persist partial metrics in case server restarts self._persist_metrics(workflow, tracker.get_report()) return final_state # ------------------------------------------------------------------ # Resume after human review # ------------------------------------------------------------------ def resume( self, workflow, ) -> WorkflowState: config = { "configurable": { "thread_id": str(workflow.id), } } # Resume tracking tracker = get_or_create_tracker(str(workflow.id)) tracker.start_stage("resume") final_state = self.graph.invoke( Command(resume=True), config=config, ) tracker.end_stage("resume") if self._is_completed(final_state): self.workflow_service.complete_workflow(workflow) report = finish_tracker(str(workflow.id)) self._persist_metrics(workflow, report) else: self.workflow_service.wait_for_review(workflow) return final_state # ------------------------------------------------------------------ # Persist metrics # ------------------------------------------------------------------ def _persist_metrics(self, workflow, report): """Save run metrics to a WorkflowCheckpoint so they survive restarts.""" if not report: return try: from app.models.workflow_checkpoint import WorkflowCheckpoint from sqlalchemy.orm import Session db = self.workflow_service.repository.db # Upsert: remove old metrics checkpoint if exists db.query(WorkflowCheckpoint).filter( WorkflowCheckpoint.workflow_run_id == workflow.id, WorkflowCheckpoint.agent_name == "METRICS", ).delete(synchronize_session=False) checkpoint = WorkflowCheckpoint( workflow_run_id=workflow.id, agent_name="METRICS", state=report, message=f"Run metrics: {report.get('total_elapsed_ms', 0):.0f}ms, {report.get('total_tokens', 0)} tokens", ) db.add(checkpoint) db.commit() except Exception: pass # Non-fatal — don't break workflow for metrics persistence # ------------------------------------------------------------------ # Cleanup # ------------------------------------------------------------------ def close(self): self.checkpoint_connection.close()