DocWeave / backend /app /workflow /executor.py
shak3008's picture
fix: remove unsupported JsonPlusSerializer kwarg for deployed langgraph version
d797968
Raw
History Blame Contribute Delete
5.45 kB
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()