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

187 lines
6.3 KiB
Python
Raw 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.

"""iframe embedding support: origin allowlist + short-lived signed tickets.
Scope, stated plainly: a ticket proves *"this workflow may be embedded by this
origin, for this user, right now"*. It is **not** an API credential — embedded
pages still authenticate their API calls with the normal session, exactly like
the rest of the app's iframe surfaces. What the ticket adds is an explicit
origin decision that can be audited, plus the ``frame-ancestors`` value the host
should send.
Tickets are stateless HS256 JWTs (so any worker can verify them) with a 60s
default TTL. Single-use is enforced best-effort per process via a replay cache;
the short TTL is the real protection.
"""
from __future__ import annotations
import hashlib
import logging
import os
import time
from typing import Any
from urllib.parse import urlparse
import jwt
from fastapi import APIRouter, HTTPException, Request
from pydantic import BaseModel, Field
from app.gateway.deps import get_current_user
from deerflow.config.app_config import get_app_config
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/workflows/embed", tags=["workflow-embed"])
_ALGORITHM = "HS256"
_AUDIENCE = "workflow-embed"
_replay_cache: dict[str, float] = {}
_REPLAY_CACHE_MAX = 4096
class TicketRequest(BaseModel):
workflow_id: str = Field(alias="workflowId")
origin: str
model_config = {"populate_by_name": True}
class VerifyRequest(BaseModel):
ticket: str
origin: str | None = None
def _signing_key() -> str:
for name in ("WORKFLOW_SECRET_KEY", "DEERFLOW_SECRET_KEY", "JWT_SECRET_KEY"):
value = os.environ.get(name)
if value and value.strip():
return hashlib.sha256(value.strip().encode("utf-8")).hexdigest()
raise HTTPException(
status_code=503,
detail={
"code": "WORKFLOW_EMBED_TICKET_INVALID",
"message": "服务端未配置签名密钥,无法签发嵌入票据",
},
)
def _normalize_origin(raw: str) -> str:
parsed = urlparse(raw.strip())
if not parsed.scheme or not parsed.netloc:
raise HTTPException(
status_code=400,
detail={"code": "WORKFLOW_EMBED_ORIGIN_DENIED", "message": "origin 格式无效"},
)
return f"{parsed.scheme.lower()}://{parsed.netloc.lower()}"
def _assert_origin_allowed(origin: str) -> str:
cfg = get_app_config().workflows.embed
normalized = _normalize_origin(origin)
allowed = {_normalize_origin(o) for o in cfg.allowed_origins if o.strip()}
if not allowed:
raise HTTPException(
status_code=403,
detail={
"code": "WORKFLOW_EMBED_ORIGIN_DENIED",
"message": "未配置任何允许的嵌入来源(workflows.embed.allowed_origins)",
},
)
if normalized not in allowed:
logger.warning("workflow embed origin denied: %s", normalized)
raise HTTPException(
status_code=403,
detail={"code": "WORKFLOW_EMBED_ORIGIN_DENIED", "message": "该来源不允许嵌入"},
)
return normalized
def _prune_replay_cache(now: float) -> None:
for key, expiry in list(_replay_cache.items()):
if expiry <= now:
_replay_cache.pop(key, None)
if len(_replay_cache) > _REPLAY_CACHE_MAX:
_replay_cache.clear()
@router.post("/tickets")
async def issue_ticket(request: Request, body: TicketRequest) -> dict[str, Any]:
cfg = get_app_config().workflows
if not cfg.enabled:
raise HTTPException(status_code=503, detail="Workflow Studio is disabled")
user_id = await get_current_user(request)
origin = _assert_origin_allowed(body.origin)
definitions = getattr(request.app.state, "workflow_store", None)
if definitions is None:
raise HTTPException(status_code=503, detail="Workflow store not available")
definition = await definitions.get_definition(body.workflow_id, include_draft=False)
if definition is None:
raise HTTPException(status_code=404, detail="工作流不存在")
if definition.get("owner_id") != user_id:
raise HTTPException(status_code=403, detail={"code": "WORKFLOW_FORBIDDEN", "message": "无权嵌入该工作流"})
now = int(time.time())
ttl = cfg.embed.ticket_ttl_seconds
payload = {
"aud": _AUDIENCE,
"sub": user_id,
"wf": body.workflow_id,
"org": origin,
"iat": now,
"exp": now + ttl,
"jti": hashlib.sha256(f"{user_id}{body.workflow_id}{origin}{now}".encode()).hexdigest()[:32],
}
ticket = jwt.encode(payload, _signing_key(), algorithm=_ALGORITHM)
logger.info(
"workflow embed ticket issued workflow=%s user=%s origin=%s ttl=%ss",
body.workflow_id,
user_id,
origin,
ttl,
)
return {
"ticket": ticket,
"expiresIn": ttl,
"origin": origin,
"frameAncestors": sorted({_normalize_origin(o) for o in cfg.embed.allowed_origins if o.strip()}),
}
@router.post("/verify")
async def verify_ticket(body: VerifyRequest) -> dict[str, Any]:
cfg = get_app_config().workflows
if not cfg.enabled:
raise HTTPException(status_code=503, detail="Workflow Studio is disabled")
try:
payload = jwt.decode(body.ticket, _signing_key(), algorithms=[_ALGORITHM], audience=_AUDIENCE)
except jwt.PyJWTError as exc:
raise HTTPException(
status_code=401,
detail={"code": "WORKFLOW_EMBED_TICKET_INVALID", "message": "票据无效或已过期"},
) from exc
if body.origin and _normalize_origin(body.origin) != payload.get("org"):
raise HTTPException(
status_code=403,
detail={"code": "WORKFLOW_EMBED_ORIGIN_DENIED", "message": "票据与当前来源不匹配"},
)
now = time.time()
_prune_replay_cache(now)
jti = str(payload.get("jti") or "")
if jti and jti in _replay_cache:
raise HTTPException(
status_code=401,
detail={"code": "WORKFLOW_EMBED_TICKET_INVALID", "message": "票据已被使用"},
)
if jti:
_replay_cache[jti] = float(payload.get("exp") or now + 60)
return {
"ok": True,
"workflowId": payload.get("wf"),
"userId": payload.get("sub"),
"origin": payload.get("org"),
"expiresAt": payload.get("exp"),
}