"""Executes leased workflow runs and owns every terminal state transition. Lifecycle for one lease: 1. resolve the published version → immutable graph snapshot, 2. replay completed node results from ``workflow_node_runs`` (so a re-lease or a resume never re-runs finished work), 3. drive :class:`WorkflowEngine` under a run-level timeout with a heartbeat renewing the lease and a watcher polling for ``cancel_requested``, 4. write exactly one terminal outcome: ``completed`` / ``failed`` / ``cancelled`` / ``awaiting_input``. Every exit path emits a matching terminal event, because the SSE endpoint uses those events to decide when to close a stream. """ from __future__ import annotations import asyncio import json import logging import secrets from dataclasses import dataclass from datetime import UTC, datetime, timedelta from typing import Any from deerflow.config.workflow_config import WorkflowConfig from deerflow.persistence.workflow_events.base import WorkflowEventStore from deerflow.persistence.workflow_runs.base import WorkflowRunStore from deerflow.workflows.errors import WorkflowError from deerflow.workflows.nodes import build_default_registry from deerflow.workflows.runtime import ( CancelToken, PersistingWorkflowEventSink, RunContext, WorkflowEngine, WorkflowPaused, WorkflowRuntimeDeps, ) from deerflow.workflows.schemas import WorkflowGraph logger = logging.getLogger(__name__) _HEARTBEAT_INTERVAL_SECONDS = 15 _CANCEL_POLL_SECONDS = 2.0 _MAX_SUBWORKFLOW_DEPTH = 3 @dataclass class LeaseParams: run_id: str lease_owner: str attempt: int = 1 class WorkflowRunExecutor: def __init__( self, app: Any, *, run_store: WorkflowRunStore, event_store: WorkflowEventStore, workflow_store: Any, data_source_store: Any = None, live_publisher: Any = None, config: WorkflowConfig | None = None, ) -> None: self._app = app self._runs = run_store self._events = event_store self._workflows = workflow_store self._data_sources = data_source_store self._publish = live_publisher self._config = config or WorkflowConfig() self._registry = build_default_registry(self._config) self._tasks: dict[str, asyncio.Task] = {} # ── public API ────────────────────────────────────────────────────── def start(self, params: LeaseParams) -> bool: """Launch a background task for a leased run. Idempotent per run id.""" existing = self._tasks.get(params.run_id) if existing is not None and not existing.done(): return False task = asyncio.create_task(self._guarded(params), name=f"workflow-run-{params.run_id}") self._tasks[params.run_id] = task task.add_done_callback(lambda _t, rid=params.run_id: self._tasks.pop(rid, None)) return True async def drain(self, timeout: float = 5.0) -> None: tasks = [t for t in self._tasks.values() if not t.done()] if not tasks: return for task in tasks: task.cancel() await asyncio.wait(tasks, timeout=timeout) # ── run driver ────────────────────────────────────────────────────── async def _guarded(self, params: LeaseParams) -> None: try: await self._execute_lease(params) except asyncio.CancelledError: raise except Exception: # noqa: BLE001 - never let a run kill the dispatcher logger.exception("workflow run %s crashed outside the engine", params.run_id) await self._fail( params, WorkflowError("WORKFLOW_INTERNAL", "运行发生内部错误,请查看服务端日志"), ) async def _execute_lease(self, params: LeaseParams) -> None: run = await self._runs.get_run(params.run_id) if run is None: return sink = self._sink(run) graph = await self._load_graph(run) if graph is None: await self._fail( params, WorkflowError("WORKFLOW_VERSION_NOT_FOUND", "运行引用的工作流版本不存在"), sink=sink, ) return cancel = CancelToken() run_context = run.get("context") or {} resume_payload = run_context.get("resume_payload") raw_env = run_context.get("env") env = dict(raw_env) if isinstance(raw_env, dict) else {} preloaded, loop_state = await self._replay_state(params.run_id) ctx = RunContext( run_id=run["id"], workflow_id=run["workflow_id"], version_id=run["workflow_version_id"], owner_id=run["owner_id"], graph=graph, inputs=run.get("input") or {}, deps=self._deps(depth=int((run.get("context") or {}).get("depth") or 0)), cancel=cancel, emit=sink.emit, resume_payload=resume_payload if isinstance(resume_payload, dict) else None, env=env, ) await sink.emit( "run.started" if not preloaded else "run.resumed", data={"attempt": params.attempt, "replayedNodes": sorted(preloaded)}, ) engine = WorkflowEngine(self._registry, run_store=self._runs) main = asyncio.create_task( asyncio.wait_for( engine.run(ctx, preloaded=preloaded, loop_state=loop_state, attempt=params.attempt), timeout=min(graph.settings.run_timeout_seconds, self._config.run_timeout_seconds), ) ) heartbeat = asyncio.create_task(self._heartbeat(params, main)) watcher = asyncio.create_task(self._watch_cancel(params, cancel, main)) try: result = await main except WorkflowPaused as paused: await self._pause(params, paused, ctx, sink) return except TimeoutError: await self._fail( params, WorkflowError( "WORKFLOW_TIMEOUT", f"运行超时({graph.settings.run_timeout_seconds}s)", details={"timeoutSeconds": graph.settings.run_timeout_seconds}, ), sink=sink, ) return except WorkflowError as exc: if exc.code == "WORKFLOW_CANCELLED" or cancel.cancelled: await self._cancel(params, sink) else: await self._fail(params, exc, sink=sink) return except asyncio.CancelledError: # Lease lost or process shutting down: leave the run claimable again. await self._runs.update_run(params.run_id, lease_owner=params.lease_owner, status="queued") raise finally: heartbeat.cancel() watcher.cancel() for task in (heartbeat, watcher): try: await task except (asyncio.CancelledError, Exception): # noqa: BLE001 pass await self._complete(params, result, sink) # ── terminal transitions ──────────────────────────────────────────── async def _complete(self, params: LeaseParams, result: Any, sink: Any) -> None: context = await self._context_with( params.run_id, loop_state=result.loop_state, steps=result.steps, ) await self._runs.update_run( params.run_id, lease_owner=params.lease_owner, status="completed", output_json=result.output, context_json=context, finished_at=datetime.now(UTC), ) await sink.emit("run.completed", data={"output": result.output, "steps": result.steps}) async def _fail(self, params: LeaseParams, error: WorkflowError, *, sink: Any = None) -> None: run = await self._runs.get_run(params.run_id) if run is None: return body = error.to_body().model_dump(by_alias=True) retry = error.retryable and int(run.get("attempt") or 1) < int(run.get("max_attempts") or 1) sink = sink or self._sink(run) if retry: await self._runs.update_run( params.run_id, lease_owner=params.lease_owner, status="queued", error_json=body, ) await sink.emit("run.queued", data={"retryAfterError": body, "attempt": run.get("attempt")}) return await self._runs.update_run( params.run_id, lease_owner=params.lease_owner, status="failed", error_json=body, finished_at=datetime.now(UTC), ) await sink.emit("run.failed", data={"error": body}) async def _cancel(self, params: LeaseParams, sink: Any) -> None: await self._runs.request_cancel(params.run_id) finalized = await self._runs.finalize_cancel(params.run_id) if finalized is None: await self._runs.update_run(params.run_id, lease_owner=params.lease_owner, status="cancelled") await sink.emit("run.cancelled", data={}) async def _pause(self, params: LeaseParams, paused: WorkflowPaused, ctx: RunContext, sink: Any) -> None: token = secrets.token_urlsafe(24) # One conditional update writes status + token + pending descriptor + # run context together; a separate follow-up context write could clobber # a resume that lands in between. context = await self._context_with( params.run_id, loop_state=dict(ctx.loop_state), resume_payload=None, ) updated = await self._runs.set_awaiting_input( params.run_id, lease_owner=params.lease_owner, pending_input=paused.pending_input, resume_token=token, context=context, ) if updated is None: # Lost the lease (cancel or steal) — do not fabricate a pause. logger.info("workflow run %s could not enter awaiting_input", params.run_id) return await sink.emit( "run.awaiting_input", data={ # Flattened wire contract: the full descriptor keeps living in # the run row (pending_input), the event exposes exactly what a # client needs to render the intervention card. "resumeToken": token, "prompt": paused.pending_input.get("prompt") or "", "formSchema": paused.pending_input.get("formSchema") or {}, "actions": paused.pending_input.get("actions") or ["submit"], "toolCallId": paused.pending_input.get("toolCallId") or "", }, node_id=paused.node_id, node_run_id=ctx.node_run_ids.get(paused.node_id), ) # ── helpers ───────────────────────────────────────────────────────── async def _context_with(self, run_id: str, **updates: Any) -> dict[str, Any]: """Merge checkpoints without dropping a draft execution graph snapshot.""" run = await self._runs.get_run(run_id) current = dict(run.get("context") or {}) if run else {} return {**current, **updates} def _sink(self, run: dict[str, Any]) -> PersistingWorkflowEventSink: return PersistingWorkflowEventSink( self._events, run_id=run["id"], workflow_id=run["workflow_id"], version_id=run["workflow_version_id"], publish=self._publish, ) async def _load_graph(self, run: dict[str, Any]) -> WorkflowGraph | None: # Draft test-runs carry their execution graph in run.context (version id # is the literal "draft", no published row exists for it). draft_graph = (run.get("context") or {}).get("draftGraph") if draft_graph: try: return WorkflowGraph.model_validate(draft_graph) except Exception: # noqa: BLE001 logger.warning("draft run %s has an invalid graph", run["id"], exc_info=True) return None version = await self._workflows.get_version(run["workflow_version_id"]) if version is None or version.get("workflow_id") != run["workflow_id"]: return None try: return WorkflowGraph.model_validate(version.get("graph") or {}) except Exception: # noqa: BLE001 logger.warning("workflow version %s has an invalid graph", run["workflow_version_id"], exc_info=True) return None async def _replay_state(self, run_id: str) -> tuple[dict[str, Any], dict[str, int]]: run = await self._runs.get_run(run_id) loop_state = {} raw_loop = (run or {}).get("context", {}).get("loop_state") if isinstance(raw_loop, dict): loop_state = {str(k): int(v) for k, v in raw_loop.items() if isinstance(v, int)} completed: dict[str, Any] = {} for row in await self._runs.list_node_runs(run_id, current_only=True): if row.get("status") != "completed": continue output = row.get("output") if isinstance(output, dict): completed[str(row["node_id"])] = output return completed, loop_state def _deps(self, *, depth: int = 0) -> WorkflowRuntimeDeps: from app.gateway.workflow_agent_runner import WorkflowAgentRunner from app.gateway.workflow_deep_research_adapter import WorkflowDeepResearchAdapter agent_runner = WorkflowAgentRunner(self._app) deep_research_adapter = WorkflowDeepResearchAdapter(self._app) async def save_artifact(**kwargs: Any) -> dict[str, Any]: return await self._runs.add_artifact(kwargs) async def resolve_data_source(source_id: str) -> dict[str, Any] | None: if self._data_sources is None: return None return await self._data_sources.resolve_dsn(source_id) async def resolve_credential(ref: str) -> dict[str, Any] | None: """HTTP credentials are data-source rows of kind ``http`` whose secret is a JSON blob of headers to inject.""" if self._data_sources is None: return None resolved = await self._data_sources.resolve_dsn(ref) if not resolved: return None try: payload = json.loads(resolved.get("dsn") or "{}") except json.JSONDecodeError: return None headers = payload.get("headers") if isinstance(payload, dict) else None return {"headers": headers} if isinstance(headers, dict) else None async def run_subworkflow(**kwargs: Any) -> dict[str, Any]: if depth + 1 > _MAX_SUBWORKFLOW_DEPTH: raise WorkflowError( "WORKFLOW_LIMIT_EXCEEDED", f"子工作流嵌套超过 {_MAX_SUBWORKFLOW_DEPTH} 层", details={"depth": depth + 1}, ) return await self._run_child(depth=depth + 1, **kwargs) return WorkflowRuntimeDeps( run_agent=agent_runner.run_agent, save_artifact=save_artifact, resolve_data_source=resolve_data_source, resolve_credential=resolve_credential, run_subworkflow=run_subworkflow, run_deep_research=deep_research_adapter.run, ) async def _run_child( self, *, depth: int, parent_run_id: str, node_id: str, owner_id: str, workflow_id: str, version_id: str, inputs: dict[str, Any], cancel: CancelToken | None = None, ) -> dict[str, Any]: """Run a child workflow inline: same process, its own run row and events.""" child, _ = await self._runs.create_run( { "workflow_id": workflow_id, "workflow_version_id": version_id, "owner_id": owner_id, "input": inputs, "status": "running", "context": { "depth": depth, "parent_run_id": parent_run_id, "parent_node_id": node_id, }, } ) graph = await self._load_graph(child) if graph is None: raise WorkflowError("WORKFLOW_VERSION_NOT_FOUND", "子工作流版本不存在", node_id=node_id) sink = self._sink(child) ctx = RunContext( run_id=child["id"], workflow_id=workflow_id, version_id=version_id, owner_id=owner_id, graph=graph, inputs=inputs, deps=self._deps(depth=depth), cancel=cancel or CancelToken(), emit=sink.emit, ) await sink.emit("run.started", data={"parentRunId": parent_run_id, "depth": depth}) engine = WorkflowEngine(self._registry, run_store=self._runs) try: result = await engine.run(ctx) except WorkflowPaused as exc: await self._runs.update_run(child["id"], status="failed") await sink.emit( "run.failed", data={ "error": { "code": "WORKFLOW_SUBWORKFLOW_FAILED", "message": "子工作流不支持人工节点", } }, ) raise WorkflowError( "WORKFLOW_SUBWORKFLOW_FAILED", "子工作流不支持人工输入节点", node_id=node_id, details={"childRunId": child["id"], "pausedAt": exc.node_id}, ) from exc except WorkflowError as exc: await self._runs.update_run(child["id"], status="failed", error_json=exc.to_body().model_dump(by_alias=True)) await sink.emit("run.failed", data={"error": exc.to_body().model_dump(by_alias=True)}) raise await self._runs.update_run(child["id"], status="completed", output_json=result.output) await sink.emit("run.completed", data={"output": result.output}) return {"run_id": child["id"], "status": "completed", "output": result.output} async def _heartbeat(self, params: LeaseParams, main: asyncio.Task) -> None: ttl = max(30, self._config.lease_ttl_seconds) while not main.done(): await asyncio.sleep(_HEARTBEAT_INTERVAL_SECONDS) if main.done(): return ok = await self._runs.renew_lease( params.run_id, lease_owner=params.lease_owner, lease_until=datetime.now(UTC) + timedelta(seconds=ttl), ) if not ok: logger.warning("workflow run %s lost its lease; aborting", params.run_id) main.cancel() return async def _watch_cancel(self, params: LeaseParams, cancel: CancelToken, main: asyncio.Task) -> None: while not main.done(): await asyncio.sleep(_CANCEL_POLL_SECONDS) run = await self._runs.get_run(params.run_id) if run is None: return if run.get("status") in ("cancel_requested", "cancelled"): cancel.cancel() return __all__ = ["LeaseParams", "WorkflowRunExecutor"]