307 lines
14 KiB
Python
307 lines
14 KiB
Python
"""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"]
|