435 lines
17 KiB
Python
435 lines
17 KiB
Python
"""Publish-time structural validator for workflow graphs (phase 1).
|
||
|
||
Full resource-permission and SQL EXPLAIN checks land in later phases.
|
||
This module never imports ``app.*``.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from collections import defaultdict, deque
|
||
from collections.abc import Callable
|
||
from typing import Any
|
||
|
||
from pydantic import ValidationError
|
||
|
||
from deerflow.workflows.errors import WorkflowError, WorkflowErrorBody
|
||
from deerflow.workflows.expressions import collect_referenced_nodes
|
||
from deerflow.workflows.schemas import ALL_NODE_TYPES, WorkflowGraph
|
||
|
||
ResourceChecker = Callable[[str, dict[str, Any]], "str | dict[str, Any] | None"]
|
||
"""Return an error message when a node resource is missing/forbidden.
|
||
|
||
May instead return ``{"message": ..., "field": "config.agentId"}`` so API
|
||
responses can point the canvas form at the exact offending field. ``None``
|
||
means the node's resources check out."""
|
||
|
||
|
||
class ValidationIssue(WorkflowErrorBody):
|
||
"""Alias kept for readability in API responses."""
|
||
|
||
|
||
def _collect_reachable(start_ids: set[str], adjacency: dict[str, list[str]]) -> set[str]:
|
||
seen: set[str] = set()
|
||
queue: deque[str] = deque(start_ids)
|
||
while queue:
|
||
node_id = queue.popleft()
|
||
if node_id in seen:
|
||
continue
|
||
seen.add(node_id)
|
||
for nxt in adjacency.get(node_id, []):
|
||
if nxt not in seen:
|
||
queue.append(nxt)
|
||
return seen
|
||
|
||
|
||
def _has_uncontrolled_cycle(graph: WorkflowGraph) -> bool:
|
||
"""Detect cycles that are not entered via an explicit ``loop`` node.
|
||
|
||
Phase-1 heuristic: if the condensation of the graph (ignoring edges whose
|
||
source is a ``loop`` node) still has a back-edge, reject.
|
||
"""
|
||
loop_ids = {n.id for n in graph.nodes if n.type == "loop"}
|
||
adjacency: dict[str, list[str]] = defaultdict(list)
|
||
for edge in graph.edges:
|
||
if edge.source in loop_ids:
|
||
continue
|
||
adjacency[edge.source].append(edge.target)
|
||
|
||
visiting: set[str] = set()
|
||
visited: set[str] = set()
|
||
|
||
def dfs(node_id: str) -> bool:
|
||
visiting.add(node_id)
|
||
for nxt in adjacency.get(node_id, []):
|
||
if nxt in visiting:
|
||
return True
|
||
if nxt not in visited and dfs(nxt):
|
||
return True
|
||
visiting.discard(node_id)
|
||
visited.add(node_id)
|
||
return False
|
||
|
||
for node in graph.nodes:
|
||
if node.id not in visited and dfs(node.id):
|
||
return True
|
||
return False
|
||
|
||
|
||
def validate_workflow_graph(
|
||
raw: dict[str, Any] | WorkflowGraph,
|
||
*,
|
||
system_max_steps: int = 500,
|
||
system_max_loop_iterations: int = 20,
|
||
system_max_parallelism: int = 16,
|
||
resource_checker: ResourceChecker | None = None,
|
||
) -> list[ValidationIssue]:
|
||
"""Validate a graph. Empty list means publishable."""
|
||
|
||
issues: list[ValidationIssue] = []
|
||
|
||
try:
|
||
graph = raw if isinstance(raw, WorkflowGraph) else WorkflowGraph.model_validate(raw)
|
||
except ValidationError as exc:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="图结构不符合 schema v1.0",
|
||
details={"errors": exc.errors()},
|
||
)
|
||
)
|
||
return issues
|
||
|
||
starts = [n for n in graph.nodes if n.type == "start"]
|
||
outputs = [n for n in graph.nodes if n.type == "output"]
|
||
if len(starts) != 1:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message=f"工作流必须恰好有 1 个开始节点,当前为 {len(starts)}",
|
||
details={"count": len(starts)},
|
||
)
|
||
)
|
||
if len(outputs) < 1:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="工作流至少需要 1 个输出节点",
|
||
)
|
||
)
|
||
|
||
unknown = [n for n in graph.nodes if n.type not in ALL_NODE_TYPES]
|
||
for node in unknown:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message=f"未知节点类型: {node.type}",
|
||
node_id=node.id,
|
||
)
|
||
)
|
||
|
||
if _has_uncontrolled_cycle(graph):
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_GRAPH_CYCLE",
|
||
message="检测到未受控环;循环必须通过显式 loop 节点",
|
||
)
|
||
)
|
||
|
||
if starts:
|
||
adjacency: dict[str, list[str]] = defaultdict(list)
|
||
for edge in graph.edges:
|
||
adjacency[edge.source].append(edge.target)
|
||
reachable = _collect_reachable({starts[0].id}, adjacency)
|
||
for node in graph.nodes:
|
||
if node.id not in reachable:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_NODE_NOT_FOUND",
|
||
message=f"节点不可从开始节点到达: {node.id}",
|
||
node_id=node.id,
|
||
)
|
||
)
|
||
for out in outputs:
|
||
if out.id not in reachable:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message=f"输出节点不可达: {out.id}",
|
||
node_id=out.id,
|
||
)
|
||
)
|
||
|
||
settings = graph.settings
|
||
if settings.max_steps > system_max_steps:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_LIMIT_EXCEEDED",
|
||
message=f"maxSteps 超过系统上限 {system_max_steps}",
|
||
details={"maxSteps": settings.max_steps, "limit": system_max_steps},
|
||
)
|
||
)
|
||
if settings.max_loop_iterations > system_max_loop_iterations:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_LIMIT_EXCEEDED",
|
||
message=f"maxLoopIterations 超过系统上限 {system_max_loop_iterations}",
|
||
details={
|
||
"maxLoopIterations": settings.max_loop_iterations,
|
||
"limit": system_max_loop_iterations,
|
||
},
|
||
)
|
||
)
|
||
if settings.max_parallelism > system_max_parallelism:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_LIMIT_EXCEEDED",
|
||
message=f"maxParallelism 超过系统上限 {system_max_parallelism}",
|
||
details={"maxParallelism": settings.max_parallelism, "limit": system_max_parallelism},
|
||
)
|
||
)
|
||
|
||
predecessors = _predecessors(graph)
|
||
for node in graph.nodes:
|
||
if node.type == "condition":
|
||
branches = node.config.get("branches") or []
|
||
ports = {str(b.get("name")) for b in branches if isinstance(b, dict) and b.get("name")}
|
||
default_branch = node.config.get("defaultBranch") or node.config.get("default_branch")
|
||
if default_branch:
|
||
ports.add(str(default_branch))
|
||
edge_ports = {e.source_port for e in graph.edges if e.source == node.id and e.source_port}
|
||
missing = edge_ports - ports
|
||
for port in missing:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_EDGE_INVALID",
|
||
message=f"条件节点出口未声明: {port}",
|
||
node_id=node.id,
|
||
details={"sourcePort": port},
|
||
)
|
||
)
|
||
if node.type == "loop":
|
||
max_iter = node.config.get("maxIterations") or node.config.get("max_iterations")
|
||
if max_iter is None:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="loop 节点必须配置 maxIterations",
|
||
node_id=node.id,
|
||
)
|
||
)
|
||
body_entry = node.config.get("bodyEntry") or node.config.get("body_entry")
|
||
if not body_entry:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="loop 节点必须配置 bodyEntry",
|
||
node_id=node.id,
|
||
)
|
||
)
|
||
elif body_entry not in graph.node_map():
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_NODE_NOT_FOUND",
|
||
message=f"loop bodyEntry 不存在: {body_entry}",
|
||
node_id=node.id,
|
||
)
|
||
)
|
||
|
||
if node.type == "evidence_normalizer":
|
||
raw_sources = node.config.get("sources")
|
||
if raw_sources is not None and not isinstance(raw_sources, list):
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="证据归一化节点的 sources 必须是数组",
|
||
node_id=node.id,
|
||
details={"field": "config.sources"},
|
||
)
|
||
)
|
||
source_ids: list[str] = []
|
||
if isinstance(raw_sources, list):
|
||
for item in raw_sources:
|
||
source_id = str(item.get("nodeId") or item.get("node_id") or "").strip() if isinstance(item, dict) else ""
|
||
if not source_id:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="证据归一化 sources 的每项都需要 nodeId",
|
||
node_id=node.id,
|
||
details={"field": "config.sources"},
|
||
)
|
||
)
|
||
continue
|
||
source_ids.append(source_id)
|
||
else:
|
||
source_ids = [edge.source for edge in graph.edges if edge.target == node.id]
|
||
if not source_ids:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="证据归一化节点至少需要一个上游研究节点",
|
||
node_id=node.id,
|
||
details={"field": "config.sources"},
|
||
)
|
||
)
|
||
ancestors = _ancestors(node.id, predecessors)
|
||
for source_id in source_ids:
|
||
if source_id not in graph.node_map():
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_NODE_NOT_FOUND",
|
||
message=f"证据归一化引用了不存在的节点: {source_id}",
|
||
node_id=node.id,
|
||
details={"referencedNodeId": source_id},
|
||
)
|
||
)
|
||
elif source_id not in ancestors:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_EXPRESSION_INVALID",
|
||
message=f"证据归一化只能引用上游节点,{source_id} 不在 {node.id} 的上游",
|
||
node_id=node.id,
|
||
details={"referencedNodeId": source_id},
|
||
)
|
||
)
|
||
|
||
if node.type == "deep_research_write":
|
||
raw_evidence = node.config.get("evidenceBinding") or node.config.get("evidence_binding")
|
||
if raw_evidence is None:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="深度研究写作节点必须绑定 Evidence Pack",
|
||
node_id=node.id,
|
||
details={"field": "config.evidenceBinding"},
|
||
)
|
||
)
|
||
elif not isinstance(raw_evidence, (str, dict)):
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="深度研究写作节点的 Evidence Pack 绑定格式无效",
|
||
node_id=node.id,
|
||
details={"field": "config.evidenceBinding"},
|
||
)
|
||
)
|
||
raw_config = node.config.get("researchConfig") or node.config.get("research_config") or {}
|
||
if not isinstance(raw_config, dict):
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="深度研究写作节点的 researchConfig 必须是对象",
|
||
node_id=node.id,
|
||
details={"field": "config.researchConfig"},
|
||
)
|
||
)
|
||
elif raw_config.get("mode") == "multi_agent":
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_SCHEMA_INVALID",
|
||
message="深度研究写作节点暂不支持 multi_agent;请在工作流中显式使用 human_input 节点",
|
||
node_id=node.id,
|
||
details={"field": "config.researchConfig.mode"},
|
||
)
|
||
)
|
||
|
||
if resource_checker is not None:
|
||
verdict = resource_checker(node.type, node.config or {})
|
||
if verdict:
|
||
message = verdict if isinstance(verdict, str) else str(verdict.get("message") or "节点资源配置无效")
|
||
field = verdict.get("field") if isinstance(verdict, dict) else None
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_RESOURCE_MISSING",
|
||
message=message,
|
||
node_id=node.id,
|
||
details={"field": field} if field else None,
|
||
)
|
||
)
|
||
|
||
issues.extend(_check_upstream_references(graph))
|
||
return issues
|
||
|
||
|
||
def _predecessors(graph: WorkflowGraph) -> dict[str, set[str]]:
|
||
preds: dict[str, set[str]] = defaultdict(set)
|
||
for edge in graph.edges:
|
||
preds[edge.target].add(edge.source)
|
||
return preds
|
||
|
||
|
||
def _ancestors(node_id: str, preds: dict[str, set[str]]) -> set[str]:
|
||
seen: set[str] = set()
|
||
stack = list(preds.get(node_id) or ())
|
||
while stack:
|
||
current = stack.pop()
|
||
if current in seen:
|
||
continue
|
||
seen.add(current)
|
||
stack.extend(preds.get(current) or ())
|
||
return seen
|
||
|
||
|
||
def _check_upstream_references(graph: WorkflowGraph) -> list[ValidationIssue]:
|
||
"""§8.1: node input references must come from upstream nodes."""
|
||
issues: list[ValidationIssue] = []
|
||
known = graph.node_map()
|
||
preds = _predecessors(graph)
|
||
for node in graph.nodes:
|
||
refs = collect_referenced_nodes(node.config)
|
||
ancestors = _ancestors(node.id, preds)
|
||
for ref in sorted(refs):
|
||
if ref not in known:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_NODE_NOT_FOUND",
|
||
message=f"表达式引用了不存在的节点: {ref}",
|
||
node_id=node.id,
|
||
details={"referencedNodeId": ref},
|
||
)
|
||
)
|
||
continue
|
||
if ref == node.id and node.type != "loop":
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_EXPRESSION_INVALID",
|
||
message="节点不能引用自身输出",
|
||
node_id=node.id,
|
||
details={"referencedNodeId": ref},
|
||
)
|
||
)
|
||
continue
|
||
if ref != node.id and ref not in ancestors:
|
||
issues.append(
|
||
ValidationIssue(
|
||
code="WORKFLOW_EXPRESSION_INVALID",
|
||
message=f"节点引用必须来自上游,{ref} 不在 {node.id} 的上游",
|
||
node_id=node.id,
|
||
details={"referencedNodeId": ref},
|
||
)
|
||
)
|
||
return issues
|
||
|
||
|
||
def assert_publishable(raw: dict[str, Any] | WorkflowGraph, **kwargs: Any) -> WorkflowGraph:
|
||
"""Validate and return the typed graph, or raise ``WorkflowError``."""
|
||
|
||
issues = validate_workflow_graph(raw, **kwargs)
|
||
if issues:
|
||
first = issues[0]
|
||
raise WorkflowError(
|
||
first.code,
|
||
first.message,
|
||
node_id=first.node_id,
|
||
details={"issues": [i.model_dump(by_alias=True) for i in issues]},
|
||
)
|
||
return raw if isinstance(raw, WorkflowGraph) else WorkflowGraph.model_validate(raw)
|
||
|
||
|
||
__all__ = [
|
||
"ResourceChecker",
|
||
"ValidationIssue",
|
||
"assert_publishable",
|
||
"validate_workflow_graph",
|
||
]
|