deerflow-code/offline-backend-20260512/backend/app/gateway/workflow_executor.py
2026-09-07 18:24:55 +08:00

479 lines
20 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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"]