277 lines
12 KiB
Python
277 lines
12 KiB
Python
"""External-effect nodes: ``http``, ``sql_read``, ``code``."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import logging
|
||
from typing import Any
|
||
|
||
from deerflow.config.workflow_config import WorkflowConfig
|
||
from deerflow.workflows.errors import WorkflowError
|
||
from deerflow.workflows.expressions import render_template, render_value
|
||
from deerflow.workflows.runtime.code_runner import run_python_snippet
|
||
from deerflow.workflows.runtime.context import RunContext
|
||
from deerflow.workflows.runtime.sql_runner import execute_read_query
|
||
from deerflow.workflows.schemas import NodeResult, WorkflowNode
|
||
from deerflow.workflows.security.http_policy import enforce_http_target, enforce_redirect
|
||
from deerflow.workflows.security.sql_policy import (
|
||
assert_parameters_bound,
|
||
assert_read_only,
|
||
assert_tables_allowed,
|
||
enforce_limit,
|
||
)
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
# Response headers safe to echo back into the event log.
|
||
_SAFE_RESPONSE_HEADERS = ("content-type", "content-length", "date", "etag", "location")
|
||
_REDACTED_REQUEST_HEADERS = ("authorization", "cookie", "x-api-key", "proxy-authorization")
|
||
|
||
|
||
class HttpNodeExecutor:
|
||
def __init__(self, config: WorkflowConfig) -> None:
|
||
self._config = config
|
||
|
||
async def execute(self, node: WorkflowNode, ctx: RunContext) -> NodeResult:
|
||
policy = self._config.http
|
||
state = ctx.state()
|
||
url = render_template(str(node.config.get("url") or ""), state).strip()
|
||
if not url:
|
||
raise WorkflowError("WORKFLOW_HTTP_FAILED", "HTTP 节点缺少 URL", node_id=node.id)
|
||
method = str(node.config.get("method") or "GET").upper()
|
||
headers = {str(k): str(v) for k, v in (render_value(node.config.get("headers") or {}, state) or {}).items()}
|
||
query = render_value(node.config.get("query") or {}, state) or {}
|
||
body = render_value(node.config.get("body"), state)
|
||
|
||
credential_ref = node.config.get("credentialRef") or node.config.get("credential_ref")
|
||
if credential_ref:
|
||
headers.update(await self._credential_headers(str(credential_ref), node, ctx))
|
||
|
||
try:
|
||
target = await enforce_http_target(url, policy)
|
||
except WorkflowError as exc:
|
||
exc.node_id = exc.node_id or node.id
|
||
raise
|
||
|
||
timeout = min(
|
||
int(node.config.get("timeoutSeconds") or node.config.get("timeout_seconds") or policy.timeout_seconds),
|
||
policy.timeout_seconds,
|
||
)
|
||
max_bytes = min(
|
||
int(node.config.get("maxResponseBytes") or node.config.get("max_response_bytes") or policy.max_response_bytes),
|
||
policy.max_response_bytes,
|
||
)
|
||
allow_redirects = bool(node.config.get("allowRedirects") or node.config.get("allow_redirects"))
|
||
max_hops = policy.max_redirects if allow_redirects else 0
|
||
|
||
await ctx.emit_event(
|
||
"node.progress",
|
||
data={
|
||
"phase": "http.request",
|
||
"method": method,
|
||
"host": target["host"],
|
||
"headers": [k for k in headers if k.lower() not in _REDACTED_REQUEST_HEADERS],
|
||
},
|
||
node_id=node.id,
|
||
)
|
||
|
||
import httpx
|
||
|
||
current_url = url
|
||
hops = 0
|
||
raw = b""
|
||
truncated = False
|
||
async with httpx.AsyncClient(timeout=timeout, follow_redirects=False) as client:
|
||
while True:
|
||
ctx.cancel.raise_if_cancelled(node.id)
|
||
try:
|
||
request = client.build_request(
|
||
method,
|
||
current_url,
|
||
headers=headers,
|
||
params=query or None,
|
||
json=body if isinstance(body, (dict, list)) else None,
|
||
content=body if isinstance(body, (str, bytes)) else None,
|
||
)
|
||
response = await client.send(request, stream=True)
|
||
except httpx.HTTPError as exc:
|
||
raise WorkflowError(
|
||
"WORKFLOW_HTTP_FAILED",
|
||
"HTTP 请求失败",
|
||
retryable=True,
|
||
node_id=node.id,
|
||
details={"reason": type(exc).__name__},
|
||
) from exc
|
||
|
||
if response.status_code in (301, 302, 303, 307, 308) and hops < max_hops:
|
||
location = response.headers.get("location") or ""
|
||
await response.aclose()
|
||
if not location:
|
||
break
|
||
current_url = await enforce_redirect(location, current_url, policy)
|
||
hops += 1
|
||
continue
|
||
|
||
chunks: list[bytes] = []
|
||
size = 0
|
||
try:
|
||
async for chunk in response.aiter_bytes():
|
||
size += len(chunk)
|
||
if size > max_bytes:
|
||
truncated = True
|
||
break
|
||
chunks.append(chunk)
|
||
finally:
|
||
await response.aclose()
|
||
raw = b"".join(chunks)
|
||
break
|
||
|
||
text_body = raw.decode("utf-8", errors="replace")
|
||
parsed: Any = None
|
||
content_type = response.headers.get("content-type", "")
|
||
if "json" in content_type.lower():
|
||
try:
|
||
parsed = json.loads(text_body)
|
||
except json.JSONDecodeError:
|
||
parsed = None
|
||
|
||
data = {
|
||
"status": response.status_code,
|
||
"ok": 200 <= response.status_code < 300,
|
||
"headers": {k: v for k, v in response.headers.items() if k.lower() in _SAFE_RESPONSE_HEADERS},
|
||
"json": parsed,
|
||
"text": None if parsed is not None else text_body,
|
||
"truncated": truncated,
|
||
"redirects": hops,
|
||
}
|
||
if not data["ok"]:
|
||
raise WorkflowError(
|
||
"WORKFLOW_HTTP_FAILED",
|
||
f"HTTP 请求返回 {response.status_code}",
|
||
retryable=response.status_code >= 500 or response.status_code == 429,
|
||
node_id=node.id,
|
||
details={"status": response.status_code, "preview": text_body[:500]},
|
||
)
|
||
return NodeResult(data=data, metadata={"host": target["host"]})
|
||
|
||
async def _credential_headers(self, ref: str, node: WorkflowNode, ctx: RunContext) -> dict[str, str]:
|
||
if ctx.deps.resolve_credential is None:
|
||
raise WorkflowError(
|
||
"WORKFLOW_RESOURCE_MISSING",
|
||
"运行环境未配置凭据解析能力",
|
||
node_id=node.id,
|
||
)
|
||
credential = await ctx.deps.resolve_credential(ref)
|
||
if not credential:
|
||
raise WorkflowError(
|
||
"WORKFLOW_RESOURCE_MISSING",
|
||
"引用的凭据不存在或无权访问",
|
||
node_id=node.id,
|
||
details={"credentialRef": ref},
|
||
)
|
||
headers = credential.get("headers")
|
||
return {str(k): str(v) for k, v in headers.items()} if isinstance(headers, dict) else {}
|
||
|
||
|
||
class SqlReadNodeExecutor:
|
||
def __init__(self, config: WorkflowConfig) -> None:
|
||
self._config = config
|
||
|
||
async def execute(self, node: WorkflowNode, ctx: RunContext) -> NodeResult:
|
||
policy = self._config.sql
|
||
if not policy.enabled:
|
||
raise WorkflowError("WORKFLOW_SQL_POLICY_DENIED", "SQL 节点已被系统禁用", node_id=node.id)
|
||
data_source_id = str(node.config.get("dataSourceId") or node.config.get("data_source_id") or "")
|
||
if not data_source_id:
|
||
raise WorkflowError("WORKFLOW_SQL_FAILED", "SQL 节点未选择数据源", node_id=node.id)
|
||
if ctx.deps.resolve_data_source is None:
|
||
raise WorkflowError("WORKFLOW_RESOURCE_MISSING", "运行环境未配置数据源解析能力", node_id=node.id)
|
||
source = await ctx.deps.resolve_data_source(data_source_id)
|
||
if not source or not source.get("dsn"):
|
||
raise WorkflowError(
|
||
"WORKFLOW_RESOURCE_MISSING",
|
||
"数据源不存在或无权访问",
|
||
node_id=node.id,
|
||
details={"dataSourceId": data_source_id},
|
||
)
|
||
|
||
statement = assert_read_only(str(node.config.get("statement") or ""))
|
||
parameters = render_value(node.config.get("parameters") or {}, ctx.state()) or {}
|
||
if not isinstance(parameters, dict):
|
||
raise WorkflowError("WORKFLOW_SQL_FAILED", "SQL 参数必须是对象", node_id=node.id)
|
||
assert_parameters_bound(statement, parameters)
|
||
assert_tables_allowed(statement, list(source.get("allowed_tables") or []))
|
||
max_rows = min(
|
||
int(node.config.get("maxRows") or node.config.get("max_rows") or policy.max_rows),
|
||
policy.max_rows,
|
||
int(source.get("max_rows") or policy.max_rows),
|
||
)
|
||
statement = enforce_limit(statement, max_rows)
|
||
|
||
await ctx.emit_event(
|
||
"node.progress",
|
||
data={"phase": "sql.query", "dataSourceId": data_source_id, "maxRows": max_rows},
|
||
node_id=node.id,
|
||
)
|
||
try:
|
||
result = await execute_read_query(
|
||
dsn=str(source["dsn"]),
|
||
statement=statement,
|
||
parameters={str(k): v for k, v in parameters.items()},
|
||
max_rows=max_rows,
|
||
timeout_seconds=policy.statement_timeout_seconds,
|
||
max_cell_chars=policy.max_cell_chars,
|
||
)
|
||
except WorkflowError as exc:
|
||
exc.node_id = exc.node_id or node.id
|
||
raise
|
||
return NodeResult(data=result)
|
||
|
||
|
||
class CodeNodeExecutor:
|
||
def __init__(self, config: WorkflowConfig) -> None:
|
||
self._config = config
|
||
|
||
async def execute(self, node: WorkflowNode, ctx: RunContext) -> NodeResult:
|
||
policy = self._config.code
|
||
if not policy.enabled:
|
||
raise WorkflowError(
|
||
"WORKFLOW_CODE_SANDBOX_FAILED",
|
||
"代码节点未启用(需管理员在 config.yaml 中开启 workflows.code.enabled)",
|
||
node_id=node.id,
|
||
)
|
||
source = str(node.config.get("source") or "")
|
||
if not source.strip():
|
||
raise WorkflowError("WORKFLOW_CODE_SANDBOX_FAILED", "代码节点没有源代码", node_id=node.id)
|
||
if len(source) > policy.max_source_chars:
|
||
raise WorkflowError(
|
||
"WORKFLOW_LIMIT_EXCEEDED",
|
||
f"代码长度超过上限 {policy.max_source_chars}",
|
||
node_id=node.id,
|
||
)
|
||
inputs = render_value(node.config.get("inputs") or {}, ctx.state()) or {}
|
||
timeout = min(
|
||
int(node.config.get("timeoutSeconds") or node.config.get("timeout_seconds") or policy.timeout_seconds),
|
||
policy.timeout_seconds,
|
||
)
|
||
allow_network = bool((node.config.get("allowNetwork") or node.config.get("allow_network")) and policy.allow_network)
|
||
try:
|
||
outcome = await run_python_snippet(
|
||
source=source,
|
||
inputs=inputs if isinstance(inputs, dict) else {"value": inputs},
|
||
timeout_seconds=timeout,
|
||
max_output_chars=policy.max_output_chars,
|
||
memory_limit_mb=policy.memory_limit_mb,
|
||
allow_network=allow_network,
|
||
)
|
||
except WorkflowError as exc:
|
||
exc.node_id = exc.node_id or node.id
|
||
raise
|
||
return NodeResult(
|
||
data=outcome["result"],
|
||
metadata={"stdout": outcome["stdout"][-4000:], "stderr": outcome["stderr"][-2000:]},
|
||
)
|
||
|
||
|
||
__all__ = ["CodeNodeExecutor", "HttpNodeExecutor", "SqlReadNodeExecutor"]
|