deerflow-code/offline-backend-20260512/backend/app/gateway/workflow_resource_checks.py
2026-09-07 18:24:55 +08:00

307 lines
14 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.

"""App-layer publish-time resource checks for workflow graphs.
The harness validator stays store-agnostic via its ``resource_checker`` hook;
this module is where the Gateway binds that hook to the real agent / skill /
data-source / published-version stores. It pre-fetches the catalogs the caller
is allowed to see, then exposes a **sync** checker so it can be passed straight
into ``validate_workflow_graph``.
Never imports from ``packages/harness`` internals beyond the public workflow
API, and the harness never imports this module.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
from fastapi import Request
from app.gateway.routers._workflow_planner_seed import WORKFLOW_PLANNER_AGENT_ID
from deerflow.config.workflow_config import WorkflowConfig
from deerflow.workflows.errors import WorkflowError
from deerflow.workflows.security.sql_policy import assert_read_only, assert_tables_allowed
# Header-ish keys that must never be configured as literal values on the canvas:
# credentials belong in an http data source referenced by ``credentialRef``.
_LITERAL_SECRET_KEYS = frozenset(
{
"authorization",
"proxy-authorization",
"cookie",
"x-api-key",
"apikey",
"api_key",
"secret",
"client_secret",
"password",
"access_token",
"refresh_token",
}
)
_HTTP_METHODS = frozenset({"GET", "POST", "PUT", "PATCH", "DELETE"})
@dataclass
class ResourceCatalog:
"""Everything the checker is allowed to reference, prefetched."""
agent_ids: set[str] = field(default_factory=set)
# agent_id -> cached skill whitelist from the agent store (empty = no
# restriction recorded).
agent_skills: dict[str, set[str]] = field(default_factory=dict)
enabled_skills: set[str] = field(default_factory=set)
# data source id -> public row (kind / enabled / allowed_* / max_rows).
data_sources: dict[str, dict[str, Any]] = field(default_factory=dict)
# workflow_id -> set of published version ids.
published_subworkflows: dict[str, set[str]] = field(default_factory=dict)
def _issue(message: str, field_path: str) -> dict[str, Any]:
return {"message": message, "field": field_path}
async def build_resource_catalog(request: Request, *, user_id: str) -> ResourceCatalog:
catalog = ResourceCatalog()
agent_store = getattr(request.app.state, "agent_store", None)
if agent_store is not None:
try:
if hasattr(agent_store, "list_visible"):
rows = await agent_store.list_visible(user_id)
elif hasattr(agent_store, "list_all"):
rows = await agent_store.list_all()
else:
rows = []
catalog.agent_ids = {
str(row.get("id") or row.get("agent_id") or "")
for row in rows or []
if isinstance(row, dict)
and (row.get("id") or row.get("agent_id"))
and str(row.get("id") or row.get("agent_id") or "") != WORKFLOW_PLANNER_AGENT_ID
}
if hasattr(agent_store, "get_extras_for"):
extras = await agent_store.get_extras_for([aid for aid in catalog.agent_ids if aid])
for agent_id, extra in (extras or {}).items():
skills = (extra or {}).get("skills") or []
if skills:
catalog.agent_skills[str(agent_id)] = {str(s) for s in skills}
except Exception: # noqa: BLE001 - publish checks degrade to strict-by-empty
catalog.agent_ids = set()
try:
from deerflow.skills import get_or_new_skill_storage
catalog.enabled_skills = {
str(getattr(skill, "name"))
for skill in get_or_new_skill_storage().load_skills(enabled_only=True) or []
if getattr(skill, "name", None)
}
except Exception: # noqa: BLE001
catalog.enabled_skills = set()
ds_store = getattr(request.app.state, "workflow_data_source_store", None)
if ds_store is not None:
try:
for row in await ds_store.list_sources(owner_id=user_id, include_shared=True) or []:
catalog.data_sources[str(row.get("id"))] = row
except Exception: # noqa: BLE001
catalog.data_sources = {}
wf_store = getattr(request.app.state, "workflow_store", None)
if wf_store is not None:
try:
for row in await wf_store.list_published_summaries() or []:
versions = catalog.published_subworkflows.setdefault(str(row.get("workflow_id")), set())
versions.add(str(row.get("version_id")))
except Exception: # noqa: BLE001
catalog.published_subworkflows = {}
return catalog
def _check_agent(config: dict[str, Any], catalog: ResourceCatalog) -> dict[str, Any] | None:
# Empty agentId means the runtime's built-in default agent — a real,
# executable fallback, not missing configuration.
agent_id = str(config.get("agentId") or config.get("agent_id") or "") or "default"
if agent_id != "default" and agent_id not in catalog.agent_ids:
return _issue("未选择可访问的智能体", "config.agentId")
if not str(config.get("promptTemplate") or config.get("prompt_template") or "").strip():
return _issue("智能体节点的提示词不能为空", "config.promptTemplate")
mode = str(config.get("responseMode") or config.get("response_mode") or "text")
schema = config.get("responseSchema") or config.get("response_schema")
if mode == "json":
if not isinstance(schema, dict) or not schema:
return _issue("JSON 输出模式必须配置 responseSchema 对象", "config.responseSchema")
if schema.get("type") not in (None, "object"):
return _issue("responseSchema 目前仅支持 object 类型", "config.responseSchema")
return None
def _check_skill(config: dict[str, Any], catalog: ResourceCatalog) -> dict[str, Any] | None:
mode = str(config.get("mode") or "agent_skill")
if mode != "agent_skill":
return _issue("callable_skill 模式暂未开放执行,请改用 agent_skill", "config.mode")
agent_id = str(config.get("agentId") or config.get("agent_id") or "") or "default"
if agent_id != "default" and agent_id not in catalog.agent_ids:
return _issue("技能节点引用的智能体不可访问", "config.agentId")
skill_names = [str(s) for s in (config.get("skillNames") or config.get("skill_names") or [])]
if not skill_names:
return _issue("技能节点至少需要选择一个技能", "config.skillNames")
for name in skill_names:
if name not in catalog.enabled_skills:
return _issue(f"技能不存在或未启用:{name}", "config.skillNames")
whitelist = catalog.agent_skills.get(agent_id)
if whitelist:
missing = [name for name in skill_names if name not in whitelist]
if missing:
return _issue(f"技能不在智能体白名单内:{'、'.join(missing)}", "config.skillNames")
return None
def _check_sql(config: dict[str, Any], catalog: ResourceCatalog, cfg: WorkflowConfig) -> dict[str, Any] | None:
source_id = str(config.get("dataSourceId") or config.get("data_source_id") or "")
if not source_id:
return _issue("SQL 节点未选择数据源", "config.dataSourceId")
source = catalog.data_sources.get(source_id)
if source is None:
return _issue("数据源不存在或无权访问", "config.dataSourceId")
if not source.get("enabled", True):
return _issue("数据源已被禁用", "config.dataSourceId")
if str(source.get("kind") or "sql") != "sql":
return _issue("该资源不是 SQL 数据源", "config.dataSourceId")
statement = str(config.get("statement") or config.get("queryTemplate") or config.get("query_template") or "")
if not statement.strip():
return _issue("SQL 节点缺少查询语句", "config.statement")
try:
assert_read_only(statement)
assert_tables_allowed(statement, list(source.get("allowed_tables") or []))
except WorkflowError as exc:
return _issue(str(exc.message if hasattr(exc, "message") else exc), "config.statement")
parameters = config.get("parameters")
if parameters is not None and not isinstance(parameters, dict):
return _issue("SQL 参数必须是对象", "config.parameters")
max_rows = config.get("maxRows") or config.get("max_rows")
if max_rows is not None:
try:
rows = int(max_rows)
except (TypeError, ValueError):
return _issue("maxRows 必须是整数", "config.maxRows")
ceiling = min(int(cfg.sql.max_rows), int(source.get("max_rows") or cfg.sql.max_rows))
if rows < 1 or rows > ceiling:
return _issue(f"maxRows 超过上限 {ceiling}", "config.maxRows")
return None
def _check_http(config: dict[str, Any], catalog: ResourceCatalog) -> dict[str, Any] | None:
method = str(config.get("method") or "GET").upper()
if method not in _HTTP_METHODS:
return _issue(f"不支持的 HTTP 方法:{method}", "config.method")
headers = config.get("headers")
if headers is not None:
if not isinstance(headers, dict):
return _issue("请求头必须是字符串到字符串的映射", "config.headers")
for name in headers:
if str(name).strip().lower() in _LITERAL_SECRET_KEYS:
return _issue(f"请勿在节点上明文配置 {name},改用凭据引用", "config.headers")
url = str(config.get("url") or "").strip()
credential_ref = str(config.get("credentialRef") or config.get("credential_ref") or "").strip()
if not url and not credential_ref:
return _issue("HTTP 节点缺少 URL 或凭据引用", "config.url")
if credential_ref:
source = catalog.data_sources.get(credential_ref)
if source is None:
return _issue("引用的凭据(HTTP 数据源)不存在或无权访问", "config.credentialRef")
if not source.get("enabled", True):
return _issue("引用的凭据已被禁用", "config.credentialRef")
if str(source.get("kind") or "") != "http":
return _issue("引用的凭据不是 HTTP 类型数据源", "config.credentialRef")
allowed = {str(m).upper() for m in (source.get("allowed_methods") or [])}
if allowed and method not in allowed:
return _issue(f"方法 {method} 不在凭据资源允许的方法列表内", "config.method")
return None
def _check_code(config: dict[str, Any], cfg: WorkflowConfig) -> dict[str, Any] | None:
if not cfg.code.enabled:
return _issue("代码节点未启用(需管理员在 config.yaml 中开启 workflows.code.enabled)", "")
language = config.get("language")
if language is not None and str(language) != "python":
return _issue("代码节点目前仅支持 python", "config.language")
source = str(config.get("source") or "")
if not source.strip():
return _issue("代码节点没有源代码", "config.source")
timeout = config.get("timeoutSeconds") or config.get("timeout_seconds")
if timeout is not None:
try:
seconds = int(timeout)
except (TypeError, ValueError):
return _issue("timeoutSeconds 必须是整数", "config.timeoutSeconds")
if seconds < 1 or seconds > int(cfg.code.timeout_seconds):
return _issue(f"timeoutSeconds 超过上限 {cfg.code.timeout_seconds}", "config.timeoutSeconds")
return None
def _check_output(config: dict[str, Any]) -> dict[str, Any] | None:
mapping = config.get("mapping")
if mapping is None:
return None
if not isinstance(mapping, dict):
return _issue("输出映射必须是对象", "config.mapping")
for key, template in mapping.items():
if not isinstance(template, str) or not template.strip():
return _issue(f"输出字段 {key} 的映射模板为空", f"config.mapping.{key}")
return None
def _check_subworkflow(config: dict[str, Any], catalog: ResourceCatalog, *, current_workflow_id: str) -> dict[str, Any] | None:
workflow_id = str(config.get("workflowId") or config.get("workflow_id") or "")
if not workflow_id:
return _issue("子工作流节点缺少 workflowId", "config.workflowId")
version_id = str(config.get("versionId") or config.get("version_id") or "")
if not version_id:
return _issue("子工作流节点缺少 versionId", "config.versionId")
if workflow_id == current_workflow_id:
return _issue("子工作流不能引用自身", "config.workflowId")
versions = catalog.published_subworkflows.get(workflow_id)
if not versions:
return _issue("子工作流未发布或无权访问", "config.workflowId")
if version_id not in versions:
return _issue("子工作流版本不存在或已下线", "config.versionId")
return None
def make_resource_checker(
catalog: ResourceCatalog,
*,
config: WorkflowConfig,
current_workflow_id: str = "",
):
"""Build the sync ``resource_checker`` consumed by ``validate_workflow_graph``."""
def checker(node_type: str, node_config: dict[str, Any]):
try:
if node_type == "agent":
return _check_agent(node_config, catalog)
if node_type == "skill":
return _check_skill(node_config, catalog)
if node_type == "sql_read":
return _check_sql(node_config, catalog, config)
if node_type == "http":
return _check_http(node_config, catalog)
if node_type == "code":
return _check_code(node_config, config)
if node_type == "output":
return _check_output(node_config)
if node_type == "subworkflow":
return _check_subworkflow(node_config, catalog, current_workflow_id=current_workflow_id)
except Exception: # noqa: BLE001 - a broken check must fail closed
return _issue("资源配置校验失败,请检查节点配置", "")
return None
return checker
__all__ = ["ResourceCatalog", "build_resource_catalog", "make_resource_checker"]