479 lines
20 KiB
Python
479 lines
20 KiB
Python
"""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"]
|