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

319 lines
11 KiB
Python

"""Public (unauthenticated) API for embedded workspace chat integrations.
Parent pages embed DeerFlow in a plain iframe and cannot read ``thread_id``
from the child frame. They pass a ``session_id`` (their own correlation id)
in the iframe URL; the embed page registers ``session_id → thread_id`` here,
then the parent polls ``GET /sessions/{session_id}/status``.
"""
from __future__ import annotations
import logging
import re
from fastapi import APIRouter, HTTPException, Request
from pydantic import BaseModel, Field, field_validator
from app.gateway.deps import get_checkpointer, get_embed_session_store, get_run_manager, get_thread_store
from app.gateway.routers.thread_shares import _load_thread_snapshot
from app.gateway.routers.threads import _derive_thread_status
from deerflow.persistence.embed_sessions.base import EmbedSessionStore
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/public/embed", tags=["public-embed"])
_SESSION_ID_RE = re.compile(r"^[A-Za-z0-9_-]{8,128}$")
_UNFINISHED_THREAD_STATUSES = frozenset({"running", "busy"})
def compute_embed_finished(
*,
has_active_run: bool,
meta_status: str,
derived_status: str,
) -> bool:
"""Return whether an embed parent may treat the Q&A session as complete."""
if has_active_run or meta_status in _UNFINISHED_THREAD_STATUSES:
return False
if derived_status == "interrupted":
return False
return True
def _validate_session_id(session_id: str) -> str:
value = session_id.strip()
if not _SESSION_ID_RE.fullmatch(value):
raise HTTPException(
status_code=422,
detail="session_id must be 8-128 characters (letters, digits, underscore, hyphen)",
)
return value
class RegisterEmbedSessionRequest(BaseModel):
session_id: str = Field(..., description="Correlation id from the parent page (iframe query param)")
thread_id: str = Field(..., description="DeerFlow conversation thread id")
@field_validator("session_id")
@classmethod
def _check_session_id(cls, value: str) -> str:
if not _SESSION_ID_RE.fullmatch(value.strip()):
raise ValueError("session_id must be 8-128 characters (letters, digits, underscore, hyphen)")
return value.strip()
@field_validator("thread_id")
@classmethod
def _check_thread_id(cls, value: str) -> str:
trimmed = value.strip()
if not trimmed:
raise ValueError("thread_id is required")
return trimmed
class RegisterEmbedSessionResponse(BaseModel):
session_id: str
thread_id: str
registered: bool = True
class EmbedThreadStatusFields(BaseModel):
finished: bool
thread_status: str
has_active_run: bool
latest_run_status: str | None = None
class EmbedThreadStatusResponse(EmbedThreadStatusFields):
thread_id: str
class EmbedSessionStatusResponse(EmbedThreadStatusFields):
session_id: str
registered: bool = Field(description="False until the iframe has registered this session_id")
thread_id: str | None = Field(default=None, description="Resolved thread id when registered")
class EmbedSessionResultResponse(EmbedThreadStatusFields):
session_id: str
registered: bool = Field(description="False until the iframe has registered this session_id")
thread_id: str | None = Field(default=None, description="Resolved thread id when registered")
messages: list[dict] = Field(default_factory=list, description="Full serialized conversation messages")
generated_text: str = Field(default="", description="Concatenated assistant/AI message text")
generated_files: list[str] = Field(
default_factory=list,
description="Virtual artifact paths produced during the conversation (e.g. /mnt/user-data/outputs/...)",
)
class EmbedThreadResultResponse(EmbedThreadStatusFields):
thread_id: str
messages: list[dict] = Field(default_factory=list)
generated_text: str = Field(default="")
generated_files: list[str] = Field(default_factory=list)
def _message_text(message: dict) -> str:
content = message.get("content")
if isinstance(content, str):
return content.strip()
if isinstance(content, list):
parts: list[str] = []
for part in content:
if isinstance(part, dict) and part.get("type") == "text" and isinstance(part.get("text"), str):
text = part["text"].strip()
if text:
parts.append(text)
return "\n".join(parts)
return ""
def extract_generated_text(messages: list[dict]) -> str:
"""Join all assistant/AI message bodies in conversation order."""
parts: list[str] = []
for message in messages:
if not isinstance(message, dict):
continue
message_type = message.get("type") or message.get("role")
if message_type not in {"ai", "assistant"}:
continue
text = _message_text(message)
if text:
parts.append(text)
return "\n\n".join(parts)
def build_embed_result_payload(
*,
messages: list[dict],
artifacts: list[str],
) -> tuple[list[dict], str, list[str]]:
return messages, extract_generated_text(messages), artifacts
async def _result_for_thread_id(thread_id: str, request: Request) -> EmbedThreadResultResponse:
checkpointer = get_checkpointer(request)
messages, artifacts = await _load_thread_snapshot(checkpointer, thread_id)
serialized_messages, generated_text, generated_files = build_embed_result_payload(
messages=messages,
artifacts=artifacts,
)
status = await _status_for_thread_id(thread_id, request)
return EmbedThreadResultResponse(
thread_id=thread_id,
messages=serialized_messages,
generated_text=generated_text,
generated_files=generated_files,
**status.model_dump(),
)
async def _status_for_thread_id(thread_id: str, request: Request) -> EmbedThreadStatusFields:
thread_store = get_thread_store(request)
run_mgr = get_run_manager(request)
checkpointer = get_checkpointer(request)
record = await thread_store.get(thread_id, user_id=None)
config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
try:
checkpoint_tuple = await checkpointer.aget_tuple(config)
except Exception:
logger.exception("Failed to read checkpoint for embed status on %s", thread_id)
raise HTTPException(status_code=500, detail="Failed to read thread state") from None
if record is None and checkpoint_tuple is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
derived_status = (
_derive_thread_status(checkpoint_tuple)
if checkpoint_tuple is not None
else (record or {}).get("status", "idle")
)
meta_status = str((record or {}).get("status") or derived_status)
has_active_run = await run_mgr.has_inflight(thread_id)
runs = await run_mgr.list_by_thread(thread_id)
latest_run_status = runs[-1].status.value if runs else None
finished = compute_embed_finished(
has_active_run=has_active_run,
meta_status=meta_status,
derived_status=derived_status,
)
return EmbedThreadStatusFields(
finished=finished,
thread_status=derived_status,
has_active_run=has_active_run,
latest_run_status=latest_run_status,
)
@router.post(
"/sessions",
response_model=RegisterEmbedSessionResponse,
summary="Register embed session_id → thread_id (called from iframe)",
)
async def register_embed_session(
body: RegisterEmbedSessionRequest,
request: Request,
) -> RegisterEmbedSessionResponse:
"""Upsert a mapping so the parent page can poll status by session_id only."""
store: EmbedSessionStore = get_embed_session_store(request)
record = await store.upsert(body.session_id, body.thread_id)
return RegisterEmbedSessionResponse(
session_id=record["session_id"],
thread_id=record["thread_id"],
)
@router.get(
"/sessions/{session_id}/status",
response_model=EmbedSessionStatusResponse,
summary="Poll embed Q&A completion by parent session_id",
)
async def get_embed_session_status(session_id: str, request: Request) -> EmbedSessionStatusResponse:
"""Primary status endpoint for iframe parents — no thread_id required."""
session_id = _validate_session_id(session_id)
store: EmbedSessionStore = get_embed_session_store(request)
record = await store.get(session_id)
if record is None:
return EmbedSessionStatusResponse(
session_id=session_id,
registered=False,
thread_id=None,
finished=False,
thread_status="pending",
has_active_run=False,
latest_run_status=None,
)
thread_id = record["thread_id"]
status = await _status_for_thread_id(thread_id, request)
return EmbedSessionStatusResponse(
session_id=session_id,
registered=True,
thread_id=thread_id,
**status.model_dump(),
)
@router.get(
"/threads/{thread_id}/status",
response_model=EmbedThreadStatusResponse,
summary="Poll embed Q&A completion by thread_id (optional / debugging)",
)
async def get_embed_thread_status(thread_id: str, request: Request) -> EmbedThreadStatusResponse:
"""Direct thread lookup when the parent already knows thread_id."""
status = await _status_for_thread_id(thread_id, request)
return EmbedThreadStatusResponse(thread_id=thread_id, **status.model_dump())
@router.get(
"/sessions/{session_id}/result",
response_model=EmbedSessionResultResponse,
summary="Fetch embed Q&A messages, generated text, and files by session_id",
)
async def get_embed_session_result(session_id: str, request: Request) -> EmbedSessionResultResponse:
"""Return model output for iframe parents once they know ``session_id`` only."""
session_id = _validate_session_id(session_id)
store: EmbedSessionStore = get_embed_session_store(request)
record = await store.get(session_id)
if record is None:
return EmbedSessionResultResponse(
session_id=session_id,
registered=False,
thread_id=None,
finished=False,
thread_status="pending",
has_active_run=False,
latest_run_status=None,
messages=[],
generated_text="",
generated_files=[],
)
thread_result = await _result_for_thread_id(record["thread_id"], request)
return EmbedSessionResultResponse(
session_id=session_id,
registered=True,
thread_id=thread_result.thread_id,
finished=thread_result.finished,
thread_status=thread_result.thread_status,
has_active_run=thread_result.has_active_run,
latest_run_status=thread_result.latest_run_status,
messages=thread_result.messages,
generated_text=thread_result.generated_text,
generated_files=thread_result.generated_files,
)
@router.get(
"/threads/{thread_id}/result",
response_model=EmbedThreadResultResponse,
summary="Fetch embed Q&A messages, generated text, and files by thread_id",
)
async def get_embed_thread_result(thread_id: str, request: Request) -> EmbedThreadResultResponse:
"""Direct thread lookup when the parent already knows thread_id."""
return await _result_for_thread_id(thread_id, request)