288 lines
11 KiB
Python
288 lines
11 KiB
Python
"""Workflow data sources and HTTP credentials.
|
||
|
||
Write-only secrets: the DSN (or the credential header blob) is accepted on
|
||
create/update, encrypted immediately, and never returned. Reads expose only a
|
||
masked target so an operator can tell two sources apart.
|
||
|
||
Shared sources (``ownerId`` null) are admin-only; a regular user may manage
|
||
their own.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
from typing import Any
|
||
from urllib.parse import urlsplit
|
||
|
||
from fastapi import APIRouter, HTTPException, Request
|
||
from pydantic import BaseModel, Field
|
||
|
||
from app.gateway.deps import get_current_user, get_optional_user_from_request
|
||
from app.gateway.workflow_audit import audit_workflow
|
||
from deerflow.config.app_config import get_app_config
|
||
from deerflow.workflows.security.secret_box import SecretKeyMissingError, secrets_available
|
||
from deerflow.workflows.security.sql_policy import assert_read_only, assert_tables_allowed
|
||
|
||
router = APIRouter(prefix="/api/workflows/data-sources", tags=["workflow-data-sources"])
|
||
|
||
_HTTP_METHODS = frozenset({"GET", "POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS"})
|
||
|
||
|
||
class DataSourceCreate(BaseModel):
|
||
name: str = Field(min_length=1, max_length=191)
|
||
description: str = ""
|
||
kind: str = Field(default="sql", pattern="^(sql|http)$")
|
||
dsn: str | None = None
|
||
headers: dict[str, str] | None = None
|
||
base_url: str = Field(default="", max_length=2048, alias="baseUrl")
|
||
allowed_methods: list[str] = Field(default_factory=list, alias="allowedMethods")
|
||
max_rows: int = Field(default=1000, ge=1, le=100_000, alias="maxRows")
|
||
allowed_tables: list[str] = Field(default_factory=list, alias="allowedTables")
|
||
enabled: bool = True
|
||
shared: bool = False
|
||
|
||
model_config = {"populate_by_name": True}
|
||
|
||
|
||
class QueryBody(BaseModel):
|
||
statement: str = Field(min_length=1)
|
||
|
||
|
||
class DataSourceUpdate(BaseModel):
|
||
name: str | None = Field(default=None, min_length=1, max_length=191)
|
||
description: str | None = None
|
||
dsn: str | None = None
|
||
headers: dict[str, str] | None = None
|
||
base_url: str | None = Field(default=None, max_length=2048, alias="baseUrl")
|
||
allowed_methods: list[str] | None = Field(default=None, alias="allowedMethods")
|
||
max_rows: int | None = Field(default=None, ge=1, le=100_000, alias="maxRows")
|
||
allowed_tables: list[str] | None = Field(default=None, alias="allowedTables")
|
||
enabled: bool | None = None
|
||
|
||
model_config = {"populate_by_name": True}
|
||
|
||
|
||
def _require_enabled() -> None:
|
||
if not get_app_config().workflows.enabled:
|
||
raise HTTPException(status_code=503, detail="Workflow Studio is disabled")
|
||
|
||
|
||
def _store(request: Request):
|
||
store = getattr(request.app.state, "workflow_data_source_store", None)
|
||
if store is None:
|
||
raise HTTPException(status_code=503, detail="Workflow data source store not available")
|
||
return store
|
||
|
||
|
||
async def _is_admin(request: Request) -> bool:
|
||
user = await get_optional_user_from_request(request)
|
||
return user is not None and getattr(user, "system_role", None) == "admin"
|
||
|
||
|
||
def _normalise_http_metadata(
|
||
base_url: str,
|
||
allowed_methods: list[str],
|
||
) -> tuple[str, list[str]]:
|
||
target = base_url.strip()
|
||
try:
|
||
parsed = urlsplit(target)
|
||
except ValueError as exc:
|
||
raise HTTPException(status_code=400, detail="HTTP 接口地址格式无效") from exc
|
||
if (
|
||
parsed.scheme not in {"http", "https"}
|
||
or not parsed.netloc
|
||
or parsed.username is not None
|
||
or parsed.password is not None
|
||
or parsed.query
|
||
or parsed.fragment
|
||
):
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail="HTTP 接口地址必须是无凭据、无查询参数的 http(s) 地址",
|
||
)
|
||
|
||
normalised_methods: list[str] = []
|
||
for raw_method in allowed_methods:
|
||
method = str(raw_method).strip().upper()
|
||
if method not in _HTTP_METHODS:
|
||
raise HTTPException(status_code=400, detail=f"不支持的 HTTP 方法:{method}")
|
||
if method not in normalised_methods:
|
||
normalised_methods.append(method)
|
||
if not normalised_methods:
|
||
raise HTTPException(status_code=400, detail="至少选择一种 HTTP 方法")
|
||
return target, normalised_methods
|
||
|
||
|
||
def _normalise_headers(headers: dict[str, str] | None) -> dict[str, str]:
|
||
normalised: dict[str, str] = {}
|
||
for raw_name, raw_value in (headers or {}).items():
|
||
name = str(raw_name).strip()
|
||
value = str(raw_value)
|
||
if not name or any(char in name + value for char in "\r\n"):
|
||
raise HTTPException(status_code=400, detail="HTTP 请求头格式无效")
|
||
normalised[name] = value
|
||
return normalised
|
||
|
||
|
||
def _secret_from(kind: str, dsn: str | None, headers: dict[str, str] | None) -> str:
|
||
if kind == "http":
|
||
return json.dumps({"headers": _normalise_headers(headers)}, ensure_ascii=False)
|
||
if not dsn:
|
||
raise HTTPException(status_code=400, detail="SQL 数据源需要提供 dsn")
|
||
if "://" not in dsn:
|
||
raise HTTPException(status_code=400, detail="dsn 格式无效")
|
||
return dsn
|
||
|
||
|
||
@router.get("")
|
||
async def list_data_sources(request: Request) -> dict[str, Any]:
|
||
_require_enabled()
|
||
user_id = await get_current_user(request)
|
||
rows = await _store(request).list_sources(owner_id=user_id)
|
||
return {"dataSources": rows, "secretsEncryptionAvailable": secrets_available()}
|
||
|
||
|
||
@router.post("")
|
||
async def create_data_source(request: Request, body: DataSourceCreate) -> dict[str, Any]:
|
||
_require_enabled()
|
||
user_id = await get_current_user(request)
|
||
is_admin = await _is_admin(request)
|
||
if body.shared and not is_admin:
|
||
raise HTTPException(status_code=403, detail="仅管理员可创建共享数据源")
|
||
secret = _secret_from(body.kind, body.dsn, body.headers)
|
||
base_url = ""
|
||
allowed_methods: list[str] = []
|
||
if body.kind == "http":
|
||
base_url, allowed_methods = _normalise_http_metadata(
|
||
body.base_url,
|
||
body.allowed_methods,
|
||
)
|
||
try:
|
||
created = await _store(request).create_source(
|
||
{
|
||
"name": body.name,
|
||
"description": body.description,
|
||
"kind": body.kind,
|
||
"dsn": secret,
|
||
"base_url": base_url,
|
||
"allowed_methods": allowed_methods,
|
||
"max_rows": body.max_rows,
|
||
"allowed_tables": body.allowed_tables,
|
||
"enabled": body.enabled,
|
||
"owner_id": None if body.shared else user_id,
|
||
"created_by": user_id,
|
||
}
|
||
)
|
||
except SecretKeyMissingError as exc:
|
||
raise HTTPException(
|
||
status_code=503,
|
||
detail={
|
||
"code": "WORKFLOW_RESOURCE_MISSING",
|
||
"message": "服务端未配置 WORKFLOW_SECRET_KEY,无法安全保存连接信息",
|
||
},
|
||
) from exc
|
||
audit_workflow("data_source.create", user_id=user_id, dataSourceId=created["id"], kind=body.kind, shared=body.shared)
|
||
return created
|
||
|
||
|
||
async def _owned(request: Request, source_id: str) -> dict[str, Any]:
|
||
row = await _store(request).get_source(source_id)
|
||
if row is None:
|
||
raise HTTPException(status_code=404, detail="数据源不存在")
|
||
user_id = await get_current_user(request)
|
||
if row.get("owner_id") == user_id:
|
||
return row
|
||
if await _is_admin(request):
|
||
return row
|
||
raise HTTPException(status_code=403, detail="无权操作该数据源")
|
||
|
||
|
||
async def _visible(request: Request, source_id: str) -> dict[str, Any]:
|
||
row = await _store(request).get_source(source_id)
|
||
if row is None:
|
||
raise HTTPException(status_code=404, detail="数据源不存在")
|
||
user_id = await get_current_user(request)
|
||
if row.get("owner_id") in (user_id, None):
|
||
return row
|
||
if await _is_admin(request):
|
||
return row
|
||
raise HTTPException(status_code=403, detail="无权查看该数据源")
|
||
|
||
|
||
@router.get("/{source_id}")
|
||
async def get_data_source(request: Request, source_id: str) -> dict[str, Any]:
|
||
_require_enabled()
|
||
return await _visible(request, source_id)
|
||
|
||
|
||
@router.post("/{source_id}/introspect")
|
||
async def introspect_data_source(request: Request, source_id: str) -> dict[str, Any]:
|
||
"""Return the configured table allowlist only — never sample rows."""
|
||
_require_enabled()
|
||
row = await _visible(request, source_id)
|
||
tables = list(row.get("allowed_tables") or [])
|
||
return {
|
||
"dataSourceId": row["id"],
|
||
"kind": row.get("kind"),
|
||
"tables": [{"name": name} for name in tables],
|
||
"schemas": [],
|
||
}
|
||
|
||
|
||
@router.post("/{source_id}/validate-query")
|
||
async def validate_data_source_query(request: Request, source_id: str, body: QueryBody) -> dict[str, Any]:
|
||
"""Validate SQL against the read-only policy without executing it."""
|
||
_require_enabled()
|
||
row = await _visible(request, source_id)
|
||
try:
|
||
statement = assert_read_only(body.statement)
|
||
assert_tables_allowed(statement, list(row.get("allowed_tables") or []))
|
||
except Exception as exc: # noqa: BLE001
|
||
from deerflow.workflows.errors import WorkflowError
|
||
|
||
if isinstance(exc, WorkflowError):
|
||
raise HTTPException(status_code=400, detail=exc.to_body().model_dump(by_alias=True)) from exc
|
||
raise
|
||
audit_workflow("data_source.validate_query", user_id=await get_current_user(request), dataSourceId=source_id)
|
||
return {"ok": True, "statement": statement}
|
||
|
||
|
||
@router.put("/{source_id}")
|
||
async def update_data_source(request: Request, source_id: str, body: DataSourceUpdate) -> dict[str, Any]:
|
||
_require_enabled()
|
||
row = await _owned(request, source_id)
|
||
payload = body.model_dump(exclude_unset=True, by_alias=False)
|
||
kind = str(row.get("kind") or "sql")
|
||
if body.dsn is not None or body.headers is not None:
|
||
payload["dsn"] = _secret_from(kind, body.dsn, body.headers)
|
||
payload.pop("headers", None)
|
||
if kind == "http":
|
||
base_url, allowed_methods = _normalise_http_metadata(
|
||
body.base_url if body.base_url is not None else str(row.get("base_url") or ""),
|
||
body.allowed_methods
|
||
if body.allowed_methods is not None
|
||
else list(row.get("allowed_methods") or []),
|
||
)
|
||
payload["base_url"] = base_url
|
||
payload["allowed_methods"] = allowed_methods
|
||
else:
|
||
payload.pop("base_url", None)
|
||
payload.pop("allowed_methods", None)
|
||
try:
|
||
updated = await _store(request).update_source(source_id, payload)
|
||
except SecretKeyMissingError as exc:
|
||
raise HTTPException(status_code=503, detail="服务端未配置 WORKFLOW_SECRET_KEY") from exc
|
||
if updated is None:
|
||
raise HTTPException(status_code=404, detail="数据源不存在")
|
||
audit_workflow("data_source.update", user_id=await get_current_user(request), dataSourceId=source_id)
|
||
return updated
|
||
|
||
|
||
@router.delete("/{source_id}")
|
||
async def delete_data_source(request: Request, source_id: str) -> dict[str, Any]:
|
||
_require_enabled()
|
||
await _owned(request, source_id)
|
||
deleted = await _store(request).delete_source(source_id)
|
||
audit_workflow("data_source.delete", user_id=await get_current_user(request), dataSourceId=source_id)
|
||
return {"deleted": deleted}
|