| 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, |
| ) |
|
|
| |
| |
| |
|
|
| @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)) |
|
|
| |
| |
| |
|
|
| def execute( |
| self, |
| state: WorkflowState | None, |
| workflow, |
| ) -> WorkflowState: |
|
|
| config = { |
| "configurable": { |
| "thread_id": str(workflow.id), |
| } |
| } |
|
|
| |
| 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) |
| |
| tracker.end_stage("total") |
| |
| self._persist_metrics(workflow, tracker.get_report()) |
|
|
| return final_state |
|
|
| |
| |
| |
|
|
| def resume( |
| self, |
| workflow, |
| ) -> WorkflowState: |
|
|
| config = { |
| "configurable": { |
| "thread_id": str(workflow.id), |
| } |
| } |
|
|
| |
| 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 |
|
|
| |
| |
| |
|
|
| 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 |
| |
| 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 |
|
|
| |
| |
| |
|
|
| def close(self): |
| self.checkpoint_connection.close() |