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()