532 lines
21 KiB
Python
532 lines
21 KiB
Python
"""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"]
|