187 lines
6.3 KiB
Python
187 lines
6.3 KiB
Python
"""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"),
|
||
}
|