"""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)