deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/workflows/runtime/engine.py
2026-09-07 18:24:55 +08:00

532 lines
21 KiB
Python
Raw Permalink 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.

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