"""Persisted DAG scheduler for workflow runs. Scheduling model ---------------- Every edge carries a state: ``pending`` → ``active`` (its source chose this port) or ``pruned`` (the source chose a different port). A node is *ready* when none of its non-back incoming edges is ``pending`` and at least one is ``active``; a node whose incoming edges all resolved to ``pruned`` is *skipped* and prunes its own outgoing edges in turn. That single rule gives us fan-out parallelism, condition branching, and merge joins without a topological pre-pass, and it terminates because each superstep resolves at least one edge. Loops are the one cycle we allow: the back edge (``loop`` node, port ``continue``) is excluded from readiness gating, and activating it clears the body's results so the next iteration re-runs them. Every re-entry is capped by ``settings.max_loop_iterations``. Durability: node results are written to the run store as they complete, so a resumed or re-leased run replays only what is missing. ``human_input`` raises :class:`WorkflowPaused`, which the caller turns into ``awaiting_input``. """ from __future__ import annotations import asyncio import logging import time from dataclasses import dataclass, field from datetime import UTC, datetime from typing import Any, Protocol from deerflow.persistence.workflow_runs.base import WorkflowRunStore from deerflow.workflows.errors import WorkflowError from deerflow.workflows.ports import ACTIVE_PORTS, LOOP_CONTINUE_PORT from deerflow.workflows.runtime.context import RunContext, WorkflowPaused from deerflow.workflows.schemas import NodeResult, WorkflowEdge, WorkflowGraph, WorkflowNode class ExecutorRegistry(Protocol): """Structural view of ``deerflow.workflows.nodes.NodeRegistry``. Declared here rather than imported so the scheduler has no dependency on the node package (which depends on the runtime context). """ def get(self, node_type: str) -> Any: ... logger = logging.getLogger(__name__) PENDING = "pending" ACTIVE = "active" PRUNED = "pruned" @dataclass class EngineResult: output: dict[str, Any] = field(default_factory=dict) node_results: dict[str, dict[str, Any]] = field(default_factory=dict) loop_state: dict[str, int] = field(default_factory=dict) steps: int = 0 def _is_back_edge(edge: WorkflowEdge, nodes: dict[str, WorkflowNode]) -> bool: source = nodes.get(edge.source) return bool(source and source.type == "loop" and (edge.source_port or "") == LOOP_CONTINUE_PORT) def _node_started_event_data(node: WorkflowNode, node_attempt: int) -> dict[str, Any]: """Safe presentation metadata for workflow clients. This stays intentionally small: UI clients need stable agent identity and an optional explicit parallel group, but never prompt templates, tools, credentials, or the full node config. A missing group is valid; clients may still use their backwards-compatible name-based presentation. """ data: dict[str, Any] = { "nodeType": node.type, "nodeKind": node.type, "name": node.name, "attempt": node_attempt, } group_id = node.config.get("parallelGroupId") or node.config.get("parallel_group_id") if isinstance(group_id, str) and group_id.strip(): data["parallelGroupId"] = group_id.strip() round_index = node.config.get("roundIndex") if round_index is None: round_index = node.config.get("round_index") if isinstance(round_index, int) and round_index >= 0: data["roundIndex"] = round_index if node.type == "agent": agent_id = node.config.get("agentId") or node.config.get("agent_id") if isinstance(agent_id, str) and agent_id.strip(): agent: dict[str, str] = { "agentId": agent_id.strip(), "name": str( node.config.get("agentName") or node.config.get("agent_name") or node.name or agent_id ), } avatar_url = node.config.get("agentAvatarUrl") or node.config.get("agent_avatar_url") if isinstance(avatar_url, str) and avatar_url.strip(): agent["avatarUrl"] = avatar_url.strip() data["agent"] = agent return data class WorkflowEngine: def __init__(self, registry: ExecutorRegistry, *, run_store: WorkflowRunStore | None = None) -> None: self._registry = registry self._store = run_store async def run( self, ctx: RunContext, *, preloaded: dict[str, dict[str, Any]] | None = None, loop_state: dict[str, int] | None = None, attempt: int = 1, ) -> EngineResult: graph: WorkflowGraph = ctx.graph nodes = graph.node_map() outgoing: dict[str, list[WorkflowEdge]] = {node_id: [] for node_id in nodes} incoming: dict[str, list[WorkflowEdge]] = {node_id: [] for node_id in nodes} for edge in graph.edges: outgoing[edge.source].append(edge) incoming[edge.target].append(edge) edge_state: dict[str, str] = {edge.id: PENDING for edge in graph.edges} done: set[str] = set() skipped: set[str] = set() ctx.loop_state.update(loop_state or {}) if preloaded: ctx.node_results.update(preloaded) done |= set(preloaded) # Re-derive edge states from the ports each replayed node selected. for node_id in list(done): self._apply_ports( node_id, ctx.node_results.get(node_id, {}).get("metadata") or {}, outgoing, edge_state, ) semaphore = asyncio.Semaphore(max(1, graph.settings.max_parallelism)) steps = 0 max_steps = max(1, graph.settings.max_steps) while True: ctx.cancel.raise_if_cancelled() self._propagate_skips(nodes, incoming, outgoing, edge_state, done, skipped) ready = [node_id for node_id in nodes if node_id not in done and node_id not in skipped and self._readiness(node_id, incoming, edge_state, nodes) == "ready"] if not ready: break steps += len(ready) if steps > max_steps: raise WorkflowError( "WORKFLOW_MAX_STEPS", f"运行步数超过上限 {max_steps}", details={"maxSteps": max_steps}, ) results = await asyncio.gather( *(self._run_one(nodes[node_id], ctx, semaphore, attempt) for node_id in ready), return_exceptions=True, ) paused: WorkflowPaused | None = None failure: BaseException | None = None for node_id, outcome in zip(ready, results, strict=True): if isinstance(outcome, WorkflowPaused): paused = paused or outcome continue if isinstance(outcome, BaseException): failure = failure or outcome continue ctx.node_results[node_id] = outcome done.add(node_id) metadata = outcome.get("metadata") or {} self._apply_ports(node_id, metadata, outgoing, edge_state) if nodes[node_id].type == "loop": await self._maybe_reset_loop_body(nodes[node_id], ctx, nodes, outgoing, edge_state, done, skipped) if failure is not None: raise failure if paused is not None: raise paused return EngineResult( output=self._collect_output(nodes, ctx), node_results=dict(ctx.node_results), loop_state=dict(ctx.loop_state), steps=steps, ) # ── scheduling helpers ────────────────────────────────────────────── def _readiness( self, node_id: str, incoming: dict[str, list[WorkflowEdge]], edge_state: dict[str, str], nodes: dict[str, WorkflowNode], ) -> str: edges = incoming.get(node_id) or [] if not edges: return "ready" back = [e for e in edges if _is_back_edge(e, nodes)] if any(edge_state[e.id] == ACTIVE for e in back): return "ready" normal = [e for e in edges if not _is_back_edge(e, nodes)] if not normal: return "blocked" if any(edge_state[e.id] == PENDING for e in normal): return "blocked" if any(edge_state[e.id] == ACTIVE for e in normal): return "ready" return "skip" def _apply_ports( self, node_id: str, metadata: dict[str, Any], outgoing: dict[str, list[WorkflowEdge]], edge_state: dict[str, str], ) -> None: active_ports = metadata.get(ACTIVE_PORTS) allowed = {str(p) for p in active_ports} if isinstance(active_ports, list) else None for edge in outgoing.get(node_id) or []: port = edge.source_port or "" if allowed is None or not port or port in allowed: edge_state[edge.id] = ACTIVE else: edge_state[edge.id] = PRUNED def _propagate_skips( self, nodes: dict[str, WorkflowNode], incoming: dict[str, list[WorkflowEdge]], outgoing: dict[str, list[WorkflowEdge]], edge_state: dict[str, str], done: set[str], skipped: set[str], ) -> None: changed = True while changed: changed = False for node_id in nodes: if node_id in done or node_id in skipped: continue if self._readiness(node_id, incoming, edge_state, nodes) != "skip": continue skipped.add(node_id) for edge in outgoing.get(node_id) or []: edge_state[edge.id] = PRUNED changed = True async def _maybe_reset_loop_body( self, loop_node: WorkflowNode, ctx: RunContext, nodes: dict[str, WorkflowNode], outgoing: dict[str, list[WorkflowEdge]], edge_state: dict[str, str], done: set[str], skipped: set[str], ) -> None: back_edges = [e for e in outgoing.get(loop_node.id) or [] if _is_back_edge(e, nodes) and edge_state[e.id] == ACTIVE] if not back_edges: return body_entry = str(loop_node.config.get("bodyEntry") or loop_node.config.get("body_entry") or "") if not body_entry or body_entry not in nodes: raise WorkflowError( "WORKFLOW_LOOP_FAILED", "循环节点的 bodyEntry 不存在", node_id=loop_node.id, details={"bodyEntry": body_entry}, ) body = self._loop_body(body_entry, loop_node.id, outgoing, nodes) for node_id in body: done.discard(node_id) skipped.discard(node_id) ctx.node_results.pop(node_id, None) for edge in outgoing.get(node_id) or []: if not _is_back_edge(edge, nodes): edge_state[edge.id] = PENDING # The back edge stays ACTIVE so the body entry is reachable again; the # loop node prunes it itself once it selects the ``done`` port. if self._store is not None and body: try: await self._store.invalidate_node_runs(ctx.run_id, sorted(body)) except Exception: # noqa: BLE001 - bookkeeping must not kill the run logger.warning("failed to invalidate node runs for loop reset", exc_info=True) await ctx.emit_event( "node.progress", data={ "phase": "loop.iteration", "iteration": ctx.loop_state.get(loop_node.id, 0), "bodyEntry": body_entry, }, node_id=loop_node.id, ) def _loop_body( self, body_entry: str, loop_id: str, outgoing: dict[str, list[WorkflowEdge]], nodes: dict[str, WorkflowNode], ) -> set[str]: """Nodes reachable from ``body_entry`` up to (and including) the loop node.""" body: set[str] = set() stack = [body_entry] while stack: current = stack.pop() if current in body: continue body.add(current) if current == loop_id: continue for edge in outgoing.get(current) or []: if not _is_back_edge(edge, nodes): stack.append(edge.target) return body # ── node execution ────────────────────────────────────────────────── async def _run_one( self, node: WorkflowNode, ctx: RunContext, semaphore: asyncio.Semaphore, attempt: int, ) -> dict[str, Any]: async with semaphore: ctx.cancel.raise_if_cancelled(node.id) retry_cfg = node.config.get("retry") if isinstance(node.config.get("retry"), dict) else {} max_attempts = max(1, min(5, int(retry_cfg.get("maxAttempts") or 1))) backoff = max(0.0, float(retry_cfg.get("backoffSeconds") or 0.5)) on_error = str(node.config.get("onError") or "fail") timeout = self._node_timeout(node, ctx) last_error: WorkflowError | None = None for node_attempt in range(1, max_attempts + 1): node_run = await self._begin_node_run(ctx, node, attempt, node_attempt) started = time.monotonic() await ctx.emit_event( "node.started", data=_node_started_event_data(node, node_attempt), node_id=node.id, ) try: executor = self._registry.get(node.type) result: NodeResult = await asyncio.wait_for(executor.execute(node, ctx), timeout=timeout) except WorkflowPaused: await self._finish_node_run(ctx, node_run, status="queued", output=None) raise except TimeoutError as exc: last_error = WorkflowError( "WORKFLOW_TIMEOUT", f"节点执行超时({timeout}s)", retryable=True, node_id=node.id, details={"timeoutSeconds": timeout}, ) await self._fail_node_run(ctx, node_run, last_error, started) if node_attempt < max_attempts: await asyncio.sleep(backoff * node_attempt) continue if on_error == "continue": return self._error_result(last_error) raise last_error from exc except WorkflowError as exc: last_error = exc exc.node_id = exc.node_id or node.id await self._fail_node_run(ctx, node_run, exc, started) if exc.retryable and node_attempt < max_attempts: await asyncio.sleep(backoff * node_attempt) continue if on_error == "continue": return self._error_result(exc) raise except asyncio.CancelledError: raise except Exception as exc: # noqa: BLE001 - unexpected node bug logger.exception("workflow node %s crashed", node.id) last_error = WorkflowError( "WORKFLOW_INTERNAL", "节点执行发生内部错误", node_id=node.id, details={"reason": type(exc).__name__}, ) await self._fail_node_run(ctx, node_run, last_error, started) if on_error == "continue": return self._error_result(last_error) raise last_error from exc payload = result.model_dump(by_alias=False) duration = int((time.monotonic() - started) * 1000) await self._finish_node_run(ctx, node_run, status="completed", output=payload, duration_ms=duration) await ctx.emit_event( "node.completed", data={ "nodeType": node.type, "durationMs": duration, "data": payload.get("data"), "warnings": payload.get("warnings") or [], }, node_id=node.id, ) return payload raise last_error or WorkflowError("WORKFLOW_INTERNAL", "节点执行失败", node_id=node.id) def _node_timeout(self, node: WorkflowNode, ctx: RunContext) -> int: configured = node.config.get("timeoutSeconds") or node.config.get("timeout_seconds") graph_limit = ctx.graph.settings.node_timeout_seconds try: value = int(configured) if configured else graph_limit except (TypeError, ValueError): value = graph_limit return max(1, min(value, graph_limit)) def _error_result(self, error: WorkflowError) -> dict[str, Any]: return { "data": {}, "messages": [], "artifacts": [], "metadata": {"failed": True, "error": error.to_body().model_dump(by_alias=True)}, "warnings": [error.message], } async def _begin_node_run(self, ctx: RunContext, node: WorkflowNode, run_attempt: int, node_attempt: int) -> dict[str, Any] | None: if self._store is None: return None try: row = await self._store.record_node_run( { "run_id": ctx.run_id, "node_id": node.id, "node_type": node.type, "status": "running", "attempt": node_attempt, "iteration": int(ctx.loop_state.get(node.id, 0)), "started_at": datetime.now(UTC), "is_current": True, } ) except Exception: # noqa: BLE001 logger.warning("failed to record node run start for %s", node.id, exc_info=True) return None ctx.node_run_ids[node.id] = str(row.get("id") or "") return row async def _finish_node_run( self, ctx: RunContext, node_run: dict[str, Any] | None, *, status: str, output: dict[str, Any] | None, duration_ms: int | None = None, ) -> None: if self._store is None or node_run is None: return try: await self._store.record_node_run( { "run_id": ctx.run_id, "node_id": node_run["node_id"], "node_type": node_run["node_type"], "attempt": node_run["attempt"], "iteration": node_run["iteration"], "status": status, "output": output, "duration_ms": duration_ms, "finished_at": datetime.now(UTC), "is_current": True, } ) except Exception: # noqa: BLE001 logger.warning("failed to record node run finish", exc_info=True) async def _fail_node_run( self, ctx: RunContext, node_run: dict[str, Any] | None, error: WorkflowError, started: float, ) -> None: duration = int((time.monotonic() - started) * 1000) body = error.to_body().model_dump(by_alias=True) if self._store is not None and node_run is not None: try: await self._store.record_node_run( { "run_id": ctx.run_id, "node_id": node_run["node_id"], "node_type": node_run["node_type"], "attempt": node_run["attempt"], "iteration": node_run["iteration"], "status": "failed", "error": body, "duration_ms": duration, "finished_at": datetime.now(UTC), "is_current": True, } ) except Exception: # noqa: BLE001 logger.warning("failed to record node run failure", exc_info=True) await ctx.emit_event( "node.failed", data={"error": body, "durationMs": duration}, node_id=error.node_id, ) def _collect_output(self, nodes: dict[str, WorkflowNode], ctx: RunContext) -> dict[str, Any]: output: dict[str, Any] = {} for node_id, node in nodes.items(): if node.type != "output": continue result = ctx.node_results.get(node_id) if isinstance(result, dict) and isinstance(result.get("data"), dict): output.update(result["data"]) return output __all__ = ["EngineResult", "WorkflowEngine"]