319 lines
11 KiB
Python
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)
|