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