File size: 5,452 Bytes
34c89e2 9308228 0868c75 e2532ee 9308228 fcacf10 e2532ee af21975 34c89e2 9308228 34c89e2 9308228 e2532ee 9308228 d797968 9308228 e2532ee 0868c75 34c89e2 9308228 e2532ee 34c89e2 9308228 fcacf10 9308228 34c89e2 0868c75 fcacf10 a6f082f 0868c75 fcacf10 a6f082f 0868c75 fcacf10 0868c75 fcacf10 0868c75 a6f082f 0868c75 34c89e2 9308228 34c89e2 a6f082f 0868c75 9308228 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 | 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() |