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

1639 lines
67 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.

"""Thread CRUD, state, and history endpoints.
Combines the existing thread-local filesystem cleanup with LangGraph
Platform-compatible thread management backed by the checkpointer.
Channel values returned in state responses are serialized through
:func:`deerflow.runtime.serialization.serialize_channel_values` to
ensure LangChain message objects are converted to JSON-safe dicts
matching the LangGraph Platform wire format expected by the
``useStream`` React hook.
"""
from __future__ import annotations
import asyncio
import json
import logging
import os
import uuid
from pathlib import Path
from typing import Any
from fastapi import APIRouter, HTTPException, Query, Request
from langgraph.checkpoint.base import empty_checkpoint, uuid6
from pydantic import BaseModel, Field, field_validator
from sqlalchemy import select
from app.gateway.authz import require_permission
from app.gateway.deps import get_checkpointer
from app.gateway.llmwiki_deposit import delete_conversation_deposit_documents
from app.gateway.utils import sanitize_log_param
from deerflow.config.extensions_config import SkillDisplayConfig
from deerflow.config.paths import Paths, get_paths
from deerflow.persistence.engine import get_engine, get_session_factory
from deerflow.persistence.skills.model import SkillDisplayAdapterRow
from deerflow.agents.roundtable_orchestrator.delivery import (
classify_seat_delivery,
visible_seat_text,
)
from deerflow.runtime import serialize_channel_values
from deerflow.runtime.references import (
adapt_reference_batches_with_displays,
build_reference_debug_payload,
build_reference_debug_summary,
)
from deerflow.runtime.references import (
build_reference_batches as build_runtime_reference_batches,
)
from deerflow.runtime.user_context import get_effective_user_id
from deerflow.utils.time import coerce_iso, now_iso
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/threads", tags=["threads"])
_RAG_DEBUG_LOGGER_NAME = "rag_citation_debug"
# Metadata keys that the server controls; clients are not allowed to set
# them. Pydantic ``@field_validator("metadata")`` strips them on every
# inbound model below so a malicious client cannot reflect a forged
# owner identity through the API surface. Defense-in-depth — the
# row-level invariant is still ``threads_meta.user_id`` populated from
# the auth contextvar; this list closes the metadata-blob echo gap.
_SERVER_RESERVED_METADATA_KEYS: frozenset[str] = frozenset({"owner_id", "user_id"})
def _strip_reserved_metadata(metadata: dict[str, Any] | None) -> dict[str, Any]:
"""Return ``metadata`` with server-controlled keys removed."""
if not metadata:
return metadata or {}
return {k: v for k, v in metadata.items() if k not in _SERVER_RESERVED_METADATA_KEYS}
def _get_rag_debug_logger() -> logging.Logger:
debug_logger = logging.getLogger(_RAG_DEBUG_LOGGER_NAME)
if getattr(debug_logger, "_rag_file_configured", False):
return debug_logger
log_path = Path(os.environ.get("RAG_CITATION_DEBUG_LOG", "logs/rag_citation_debug.log"))
if not log_path.is_absolute():
log_path = Path.cwd() / log_path
log_path.parent.mkdir(parents=True, exist_ok=True)
handler = logging.FileHandler(log_path, encoding="utf-8")
handler.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s"))
debug_logger.addHandler(handler)
debug_logger.setLevel(logging.INFO)
debug_logger.propagate = False
setattr(debug_logger, "_rag_file_configured", True)
setattr(debug_logger, "_rag_file_path", str(log_path))
return debug_logger
def _rag_debug_file_path() -> str:
return str(getattr(_get_rag_debug_logger(), "_rag_file_path", "logs/rag_citation_debug.log"))
def _log_rag_debug(message: str, *args: Any) -> None:
logger.info(message, *args)
_get_rag_debug_logger().info(message, *args)
def _json_text_loads(value: str | None) -> dict[str, Any] | None:
if not value:
return None
try:
parsed = json.loads(value)
return parsed if isinstance(parsed, dict) else None
except Exception:
logger.debug("Failed to parse persisted skill display JSON", exc_info=True)
return None
async def _load_persisted_skill_displays() -> dict[str, SkillDisplayConfig]:
engine = get_engine()
sf = get_session_factory()
if engine is None or sf is None:
return {}
try:
async with engine.begin() as conn:
await conn.run_sync(SkillDisplayAdapterRow.__table__.create, checkfirst=True)
async with sf() as session:
result = await session.execute(select(SkillDisplayAdapterRow))
displays: dict[str, SkillDisplayConfig] = {}
for row in result.scalars():
payload = _json_text_loads(row.display_config)
if payload is None:
continue
try:
display = SkillDisplayConfig.model_validate(payload)
except Exception:
logger.debug(
"Failed to validate persisted skill display for %s",
row.skill_name,
exc_info=True,
)
continue
if display.enabled and display.mode != "none":
displays[row.skill_name] = display
return displays
except Exception:
logger.debug("Failed to load persisted skill display adapters", exc_info=True)
return {}
# ---------------------------------------------------------------------------
# Response / request models
# ---------------------------------------------------------------------------
class ThreadDeleteResponse(BaseModel):
"""Response model for thread cleanup."""
success: bool
message: str
class ThreadResponse(BaseModel):
"""Response model for a single thread."""
thread_id: str = Field(description="Unique thread identifier")
status: str = Field(default="idle", description="Thread status: idle, busy, interrupted, error")
created_at: str = Field(default="", description="ISO timestamp")
updated_at: str = Field(default="", description="ISO timestamp")
metadata: dict[str, Any] = Field(default_factory=dict, description="Thread metadata")
values: dict[str, Any] = Field(default_factory=dict, description="Current state channel values")
interrupts: dict[str, Any] = Field(default_factory=dict, description="Pending interrupts")
class ThreadCreateRequest(BaseModel):
"""Request body for creating a thread."""
thread_id: str | None = Field(default=None, description="Optional thread ID (auto-generated if omitted)")
assistant_id: str | None = Field(default=None, description="Associate thread with an assistant")
metadata: dict[str, Any] = Field(default_factory=dict, description="Initial metadata")
_strip_reserved = field_validator("metadata")(classmethod(lambda cls, v: _strip_reserved_metadata(v)))
class ThreadSearchRequest(BaseModel):
"""Request body for searching threads."""
metadata: dict[str, Any] = Field(default_factory=dict, description="Metadata filter (exact match)")
limit: int = Field(default=100, ge=1, le=1000, description="Maximum results")
offset: int = Field(default=0, ge=0, description="Pagination offset")
status: str | None = Field(default=None, description="Filter by thread status")
query: str | None = Field(default=None, description="Case-insensitive substring match on the thread title")
class ThreadStateResponse(BaseModel):
"""Response model for thread state."""
values: dict[str, Any] = Field(default_factory=dict, description="Current channel values")
next: list[str] = Field(default_factory=list, description="Next tasks to execute")
metadata: dict[str, Any] = Field(default_factory=dict, description="Checkpoint metadata")
checkpoint: dict[str, Any] = Field(default_factory=dict, description="Checkpoint info")
checkpoint_id: str | None = Field(default=None, description="Current checkpoint ID")
parent_checkpoint_id: str | None = Field(default=None, description="Parent checkpoint ID")
created_at: str | None = Field(default=None, description="Checkpoint timestamp")
tasks: list[dict[str, Any]] = Field(default_factory=list, description="Interrupted task details")
class ThreadMessagesResponse(BaseModel):
"""A page of the latest thread message state.
``messages`` deliberately has no narrower response model: LangChain message
fields such as ``additional_kwargs`` and provider-specific tool-call payloads
must reach integration clients unchanged.
"""
total: int = Field(description="Total number of messages in the thread")
messages: list[Any] = Field(default_factory=list, description="Messages in chronological order")
class ThreadPatchRequest(BaseModel):
"""Request body for patching thread metadata."""
metadata: dict[str, Any] = Field(default_factory=dict, description="Metadata to merge")
_strip_reserved = field_validator("metadata")(classmethod(lambda cls, v: _strip_reserved_metadata(v)))
class ThreadStateUpdateRequest(BaseModel):
"""Request body for updating thread state (human-in-the-loop resume)."""
values: dict[str, Any] | None = Field(default=None, description="Channel values to merge")
checkpoint_id: str | None = Field(default=None, description="Checkpoint to branch from")
checkpoint: dict[str, Any] | None = Field(default=None, description="Full checkpoint object")
as_node: str | None = Field(default=None, description="Node identity for the update")
class HistoryEntry(BaseModel):
"""Single checkpoint history entry."""
checkpoint_id: str
parent_checkpoint_id: str | None = None
metadata: dict[str, Any] = Field(default_factory=dict)
values: dict[str, Any] = Field(default_factory=dict)
created_at: str | None = None
next: list[str] = Field(default_factory=list)
class ThreadHistoryRequest(BaseModel):
"""Request body for checkpoint history."""
limit: int = Field(default=10, ge=1, le=100, description="Maximum entries")
before: str | None = Field(default=None, description="Cursor for pagination")
class ReferenceSourceResponse(BaseModel):
index: int
id: str | None = None
skillName: str = "backend-reference"
mode: str = "citation"
resultSetKey: str | None = None
resultSetLabel: str | None = None
title: str
type: str = "result"
snippet: str = ""
source: str | None = None
author: str | None = None
time: str | None = None
url: str | None = None
content: str | None = None
score: float | None = None
tableColumns: list[dict[str, Any]] = Field(default_factory=list)
raw: dict[str, Any] = Field(default_factory=dict)
class ReferenceBatchResponse(BaseModel):
id: str
assistant_message_id: str
sources: list[ReferenceSourceResponse] = Field(default_factory=list)
class ReferenceBatchesResponse(BaseModel):
batches: list[ReferenceBatchResponse] = Field(default_factory=list)
class ReferenceDebugLogRequest(BaseModel):
event: str
thread_id: str | None = None
payload: dict[str, Any] = Field(default_factory=dict)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _delete_thread_data(thread_id: str, paths: Paths | None = None, *, user_id: str | None = None) -> ThreadDeleteResponse:
"""Delete local persisted filesystem data for a thread."""
path_manager = paths or get_paths()
try:
path_manager.delete_thread_dir(thread_id, user_id=user_id)
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
except FileNotFoundError:
# Not critical — thread data may not exist on disk
logger.debug("No local thread data to delete for %s", sanitize_log_param(thread_id))
return ThreadDeleteResponse(success=True, message=f"No local data for {thread_id}")
except Exception as exc:
logger.exception("Failed to delete thread data for %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to delete local thread data.") from exc
logger.info("Deleted local thread data for %s", sanitize_log_param(thread_id))
return ThreadDeleteResponse(success=True, message=f"Deleted local thread data for {thread_id}")
def _derive_thread_status(checkpoint_tuple) -> str:
"""Derive thread status from checkpoint metadata."""
if checkpoint_tuple is None:
return "idle"
pending_writes = getattr(checkpoint_tuple, "pending_writes", None) or []
# Check for error in pending writes
for pw in pending_writes:
if len(pw) >= 2 and pw[1] == "__error__":
return "error"
# Check for pending next tasks (indicates interrupt)
tasks = getattr(checkpoint_tuple, "tasks", None)
if tasks:
return "interrupted"
return "idle"
def _message_type(message: Any) -> str:
if isinstance(message, dict):
return str(message.get("type") or message.get("role") or "")
return str(getattr(message, "type", "") or "")
def _message_id(message: Any) -> str | None:
if isinstance(message, dict):
value = message.get("id")
else:
value = getattr(message, "id", None)
return str(value) if value else None
def _serialize_thread_messages(messages: Any) -> list[Any]:
"""Return checkpoint messages in the same wire form as stream/state APIs.
The latest checkpoint is the canonical conversation transcript. Reusing the
runtime serializer keeps every standard and provider-specific message field
(including ``additional_kwargs``, ``tool_calls`` and hidden context flags)
while converting LangChain message objects into JSON-safe values.
"""
if not isinstance(messages, list):
return []
serialized = serialize_channel_values({"messages": messages}).get("messages", [])
return serialized if isinstance(serialized, list) else []
def _message_content(message: Any) -> Any:
if isinstance(message, dict):
return message.get("content")
return getattr(message, "content", None)
def _message_name(message: Any) -> str:
if isinstance(message, dict):
return str(message.get("name") or (message.get("additional_kwargs") or {}).get("name") or "tool")
return str(getattr(message, "name", None) or "tool")
def _has_tool_calls(message: Any) -> bool:
if isinstance(message, dict):
direct = message.get("tool_calls")
chunks = message.get("tool_call_chunks")
extra = (message.get("additional_kwargs") or {}).get("tool_calls")
else:
direct = getattr(message, "tool_calls", None)
chunks = getattr(message, "tool_call_chunks", None)
extra = (getattr(message, "additional_kwargs", None) or {}).get("tool_calls")
return bool(direct or chunks or extra)
def _has_answer_content(message: Any) -> bool:
content = _message_content(message)
if isinstance(content, str):
return bool(content.strip())
if isinstance(content, list):
for part in content:
if isinstance(part, dict) and part.get("type") == "text" and str(part.get("text") or "").strip():
return True
return False
def _extract_tool_payload(message: Any) -> Any:
content = _message_content(message)
if isinstance(content, dict):
return _unwrap_tool_payload(content)
if isinstance(content, list):
texts: list[str] = []
for part in content:
if not isinstance(part, dict):
continue
if isinstance(part.get("json"), dict):
return _unwrap_tool_payload(part["json"])
if part.get("type") == "json" and isinstance(part.get("json"), dict):
return _unwrap_tool_payload(part["json"])
if part.get("type") == "text" and isinstance(part.get("text"), str):
texts.append(part["text"])
unwrapped = _unwrap_tool_payload(part)
if unwrapped is not None:
return unwrapped
content = "\n".join(texts)
if isinstance(content, str):
return _parse_json_like(content)
return None
def _unwrap_tool_payload(value: dict[str, Any]) -> Any:
if any(isinstance(value.get(key), list) for key in ("results", "data", "items", "hits", "documents", "records", "sources")):
return value
for key in ("result", "output", "stdout", "content", "text", "payload", "response"):
nested = value.get(key)
if isinstance(nested, str):
parsed = _parse_json_like(nested)
if parsed is not None:
return parsed
if isinstance(nested, dict):
unwrapped = _unwrap_tool_payload(nested)
if unwrapped is not None:
return unwrapped
return value
def _parse_json_like(text: str) -> Any:
trimmed = text.strip()
if not trimmed:
return None
try:
return json.loads(trimmed)
except Exception:
pass
first_obj = trimmed.find("{")
last_obj = trimmed.rfind("}")
if 0 <= first_obj < last_obj:
try:
return json.loads(trimmed[first_obj : last_obj + 1])
except Exception:
return None
return None
def _first_list(payload: Any) -> list[Any]:
if isinstance(payload, list):
return payload
if not isinstance(payload, dict):
return []
for key in ("results", "data", "items", "hits", "documents", "records", "entries", "matches", "chunks", "sources"):
value = payload.get(key)
if isinstance(value, list):
return value
for key in ("result", "payload", "response", "output", "data"):
nested = payload.get(key)
found = _first_list(nested)
if found:
return found
return []
def _nested_records(raw: Any) -> list[dict[str, Any]]:
if not isinstance(raw, dict):
return []
layers = [raw]
for key in ("metadata", "meta", "attrs", "attributes", "properties", "extra"):
nested = raw.get(key)
if isinstance(nested, dict):
layers.append(nested)
return layers
def _pick(layers: list[dict[str, Any]], keys: tuple[str, ...]) -> str | None:
for layer in layers:
for key in keys:
value = layer.get(key)
if isinstance(value, (str, int, float)) and str(value).strip():
return str(value).strip()
return None
def _normalize_reference_item(raw: Any, fallback_index: int) -> dict[str, Any] | None:
if not isinstance(raw, dict):
return None
layers = _nested_records(raw)
title = _pick(layers, ("title", "m_title", "name", "question", "keyword", "query", "text"))
content = _pick(layers, ("content", "page_content", "summary", "snippet", "description", "content_preview", "abstract", "body"))
url = _pick(layers, ("url", "link", "href", "source_url", "web_url"))
source = _pick(layers, ("source", "source1", "site", "provider", "from"))
time = _pick(layers, ("time", "date", "published_date", "publishedDate", "m_publish", "publish_time", "created_at"))
item_id = _pick(layers, ("recUuid", "rec_uuid", "uuid", "docId", "doc_id", "id", "recordId", "taskId"))
score_raw = _pick(layers, ("score", "relevance", "similarity"))
score = None
if score_raw is not None:
try:
score = float(score_raw)
except Exception:
score = None
if not title:
title = source or f"结果 {fallback_index}"
if not content:
content = ""
if not title and not content and not url:
return None
return {
"id": item_id or url,
"title": title,
"type": "result",
"snippet": content,
"content": content,
"source": source,
"time": time,
"url": url,
"score": score,
"raw": raw,
}
def _reference_key(source: dict[str, Any]) -> str:
raw = source.get("raw") if isinstance(source.get("raw"), dict) else {}
item_id = (
source.get("id")
or source.get("url")
or raw.get("recUuid")
or raw.get("uuid")
or raw.get("docId")
or raw.get("id")
)
if item_id:
return str(item_id).casefold().strip()
return "|".join(
str(source.get(key) or "").casefold().strip()
for key in ("title", "source", "time", "content", "snippet")
)
def _build_reference_batches(messages: list[Any]) -> list[ReferenceBatchResponse]:
batches: list[ReferenceBatchResponse] = []
for index, message in enumerate(messages):
if _message_type(message) != "human":
continue
end = len(messages)
for cursor in range(index + 1, len(messages)):
if _message_type(messages[cursor]) == "human":
end = cursor
break
assistant_message: Any | None = None
for cursor in range(end - 1, index, -1):
candidate = messages[cursor]
if (
_message_type(candidate) == "ai"
and _message_id(candidate)
and _has_answer_content(candidate)
and not _has_tool_calls(candidate)
):
assistant_message = candidate
break
if assistant_message is None:
continue
seen: set[str] = set()
sources: list[ReferenceSourceResponse] = []
for tool_message in messages[index + 1 : end]:
if _message_type(tool_message) != "tool":
continue
payload = _extract_tool_payload(tool_message)
for raw_item in _first_list(payload):
normalized = _normalize_reference_item(raw_item, len(sources) + 1)
if not normalized:
continue
key = _reference_key(normalized)
if key and key in seen:
continue
if key:
seen.add(key)
sources.append(ReferenceSourceResponse(index=len(sources) + 1, **normalized))
if sources:
mid = _message_id(assistant_message) or str(index)
batches.append(ReferenceBatchResponse(id=mid, assistant_message_id=mid, sources=sources))
return batches
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
@router.delete("/{thread_id}", response_model=ThreadDeleteResponse)
@require_permission("threads", "delete", owner_check=True, require_existing=True)
async def delete_thread_data(thread_id: str, request: Request) -> ThreadDeleteResponse:
"""Delete local persisted filesystem data for a thread.
Cleans DeerFlow-managed thread directories, removes checkpoint data,
and removes the thread_meta row from the configured ThreadMetaStore
(sqlite or memory).
"""
from app.gateway.deps import get_thread_store
# Clean local filesystem
response = _delete_thread_data(thread_id, user_id=get_effective_user_id())
# Remove checkpoints (best-effort)
checkpointer = getattr(request.app.state, "checkpointer", None)
if checkpointer is not None:
try:
if hasattr(checkpointer, "adelete_thread"):
await checkpointer.adelete_thread(thread_id)
except Exception:
logger.debug("Could not delete checkpoints for thread %s (not critical)", sanitize_log_param(thread_id))
# Remove thread_meta row (best-effort) — required for sqlite backend
# so the deleted thread no longer appears in /threads/search.
try:
thread_store = get_thread_store(request)
await thread_store.delete(thread_id)
except Exception:
logger.debug("Could not delete thread_meta for %s (not critical)", sanitize_log_param(thread_id))
# Conversation deposits are an auxiliary artifact index. Remove their
# records/documents as well so a whole-thread delete cannot leave a stale
# searchable copy behind. The helper is intentionally best-effort (the
# primary thread state is already deleted above).
asyncio.create_task(
delete_conversation_deposit_documents(
request.app,
thread_id=thread_id,
message_ids=None,
)
)
return response
@router.post("", response_model=ThreadResponse)
async def create_thread(body: ThreadCreateRequest, request: Request) -> ThreadResponse:
"""Create a new thread.
Writes a thread_meta record (so the thread appears in /threads/search)
and an empty checkpoint (so state endpoints work immediately).
Idempotent: returns the existing record when ``thread_id`` already exists.
"""
from app.gateway.deps import get_thread_store
checkpointer = get_checkpointer(request)
thread_store = get_thread_store(request)
thread_id = body.thread_id or str(uuid.uuid4())
now = now_iso()
# ``body.metadata`` is already stripped of server-reserved keys by
# ``ThreadCreateRequest._strip_reserved`` — see the model definition.
# Idempotency: return existing record when already present
existing_record = await thread_store.get(thread_id)
if existing_record is not None:
return ThreadResponse(
thread_id=thread_id,
status=existing_record.get("status", "idle"),
created_at=coerce_iso(existing_record.get("created_at", "")),
updated_at=coerce_iso(existing_record.get("updated_at", "")),
metadata=existing_record.get("metadata", {}),
)
# Write thread_meta so the thread appears in /threads/search immediately
try:
await thread_store.create(
thread_id,
assistant_id=getattr(body, "assistant_id", None),
metadata=body.metadata,
)
except Exception:
logger.exception("Failed to write thread_meta for %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to create thread")
# Write an empty checkpoint so state endpoints work immediately
config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
try:
ckpt_metadata = {
"step": -1,
"source": "input",
"writes": None,
"parents": {},
**body.metadata,
"created_at": now,
}
await checkpointer.aput(config, empty_checkpoint(), ckpt_metadata, {})
except Exception:
logger.exception("Failed to create checkpoint for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to create thread")
logger.info("Thread created: %s", sanitize_log_param(thread_id))
return ThreadResponse(
thread_id=thread_id,
status="idle",
created_at=now,
updated_at=now,
metadata=body.metadata,
)
@router.post("/search", response_model=list[ThreadResponse])
async def search_threads(body: ThreadSearchRequest, request: Request) -> list[ThreadResponse]:
"""Search and list threads.
Delegates to the configured ThreadMetaStore implementation
(SQL-backed for sqlite/postgres, Store-backed for memory mode).
"""
from app.gateway.deps import get_thread_store
repo = get_thread_store(request)
# System threads (scheduler runs, roundtable workers, report-structure /
# writing setup wizards, …) carry ``metadata.system=true`` and/or a known
# ``thread_type`` — hide them from the default conversation list so they
# don't crowd out real chats. A caller that explicitly filters on
# ``thread_type``/``system`` (e.g. the studio notebook list) opts out of
# the exclusion.
explicit_system_filter = bool(body.metadata) and ("thread_type" in body.metadata or "system" in body.metadata)
rows = await repo.search(
metadata=body.metadata or None,
status=body.status,
limit=body.limit,
offset=body.offset,
exclude_system=not explicit_system_filter,
query=body.query,
)
return [
ThreadResponse(
thread_id=r["thread_id"],
status=r.get("status", "idle"),
# ``coerce_iso`` heals legacy unix-second values that
# ``MemoryThreadMetaStore`` historically wrote with ``time.time()``;
# SQL-backed rows already arrive as ISO strings and pass through.
created_at=coerce_iso(r.get("created_at", "")),
updated_at=coerce_iso(r.get("updated_at", "")),
metadata=r.get("metadata", {}),
values={"title": r["display_name"]} if r.get("display_name") else {},
interrupts={},
)
for r in rows
]
class ThreadCountResponse(BaseModel):
"""Response model for the thread count endpoint."""
total: int = Field(description="Number of threads matching the filters")
@router.post("/count", response_model=ThreadCountResponse)
async def count_threads(body: ThreadSearchRequest, request: Request) -> ThreadCountResponse:
"""Count threads matching the same filters as ``/threads/search``.
``limit``/``offset`` in the body are ignored. Backs the chats page
pagination total ("共 N 条 / X 页") without fetching any rows.
"""
from app.gateway.deps import get_thread_store
repo = get_thread_store(request)
explicit_system_filter = bool(body.metadata) and ("thread_type" in body.metadata or "system" in body.metadata)
total = await repo.count(
metadata=body.metadata or None,
status=body.status,
exclude_system=not explicit_system_filter,
query=body.query,
)
return ThreadCountResponse(total=total)
@router.patch("/{thread_id}", response_model=ThreadResponse)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def patch_thread(thread_id: str, body: ThreadPatchRequest, request: Request) -> ThreadResponse:
"""Merge metadata into a thread record."""
from app.gateway.deps import get_thread_store
thread_store = get_thread_store(request)
record = await thread_store.get(thread_id)
if record is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
# ``body.metadata`` already stripped by ``ThreadPatchRequest._strip_reserved``.
try:
await thread_store.update_metadata(thread_id, body.metadata)
except Exception:
logger.exception("Failed to patch thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to update thread")
# Re-read to get the merged metadata + refreshed updated_at
record = await thread_store.get(thread_id) or record
return ThreadResponse(
thread_id=thread_id,
status=record.get("status", "idle"),
created_at=coerce_iso(record.get("created_at", "")),
updated_at=coerce_iso(record.get("updated_at", "")),
metadata=record.get("metadata", {}),
)
@router.get("/{thread_id}", response_model=ThreadResponse)
@require_permission("threads", "read", owner_check=True)
async def get_thread(thread_id: str, request: Request) -> ThreadResponse:
"""Get thread info.
Reads metadata from the ThreadMetaStore and derives the accurate
execution status from the checkpointer. Falls back to the checkpointer
alone for threads that pre-date ThreadMetaStore adoption (backward compat).
"""
from app.gateway.deps import get_thread_store
thread_store = get_thread_store(request)
checkpointer = get_checkpointer(request)
record: dict | None = await thread_store.get(thread_id)
# Derive accurate status from the checkpointer
config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
try:
checkpoint_tuple = await checkpointer.aget_tuple(config)
except Exception:
logger.exception("Failed to get checkpoint for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to get thread")
if record is None and checkpoint_tuple is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
# If the thread exists in the checkpointer but not in thread_meta (e.g.
# legacy data created before thread_meta adoption), synthesize a minimal
# record from the checkpoint metadata.
if record is None and checkpoint_tuple is not None:
ckpt_meta = getattr(checkpoint_tuple, "metadata", {}) or {}
record = {
"thread_id": thread_id,
"status": "idle",
"created_at": coerce_iso(ckpt_meta.get("created_at", "")),
"updated_at": coerce_iso(ckpt_meta.get("updated_at", ckpt_meta.get("created_at", ""))),
"metadata": {k: v for k, v in ckpt_meta.items() if k not in ("created_at", "updated_at", "step", "source", "writes", "parents")},
}
if record is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
status = _derive_thread_status(checkpoint_tuple) if checkpoint_tuple is not None else record.get("status", "idle")
checkpoint = getattr(checkpoint_tuple, "checkpoint", {}) or {} if checkpoint_tuple is not None else {}
channel_values = checkpoint.get("channel_values", {})
return ThreadResponse(
thread_id=thread_id,
status=status,
created_at=coerce_iso(record.get("created_at", "")),
updated_at=coerce_iso(record.get("updated_at", "")),
metadata=record.get("metadata", {}),
values=serialize_channel_values(channel_values),
)
# ---------------------------------------------------------------------------
@router.get("/{thread_id}/state", response_model=ThreadStateResponse)
@require_permission("threads", "read", owner_check=True)
async def get_thread_state(thread_id: str, request: Request) -> ThreadStateResponse:
"""Get the latest state snapshot for a thread.
Channel values are serialized to ensure LangChain message objects
are converted to JSON-safe dicts.
"""
checkpointer = get_checkpointer(request)
config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
try:
checkpoint_tuple = await checkpointer.aget_tuple(config)
except Exception:
logger.exception("Failed to get state for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to get thread state")
if checkpoint_tuple is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
checkpoint = getattr(checkpoint_tuple, "checkpoint", {}) or {}
metadata = getattr(checkpoint_tuple, "metadata", {}) or {}
checkpoint_id = None
ckpt_config = getattr(checkpoint_tuple, "config", {})
if ckpt_config:
checkpoint_id = ckpt_config.get("configurable", {}).get("checkpoint_id")
channel_values = checkpoint.get("channel_values", {})
parent_config = getattr(checkpoint_tuple, "parent_config", None)
parent_checkpoint_id = None
if parent_config:
parent_checkpoint_id = parent_config.get("configurable", {}).get("checkpoint_id")
tasks_raw = getattr(checkpoint_tuple, "tasks", []) or []
next_tasks = [t.name for t in tasks_raw if hasattr(t, "name")]
tasks = [{"id": getattr(t, "id", ""), "name": getattr(t, "name", "")} for t in tasks_raw]
values = serialize_channel_values(channel_values)
return ThreadStateResponse(
values=values,
next=next_tasks,
metadata=metadata,
checkpoint={"id": checkpoint_id, "ts": coerce_iso(metadata.get("created_at", ""))},
checkpoint_id=checkpoint_id,
parent_checkpoint_id=parent_checkpoint_id,
created_at=coerce_iso(metadata.get("created_at", "")),
tasks=tasks,
)
@router.get("/{thread_id}/reference-batches", response_model=ReferenceBatchesResponse)
@require_permission("threads", "read", owner_check=True)
async def get_thread_reference_batches(thread_id: str, request: Request) -> ReferenceBatchesResponse:
"""Return backend-built citation batches for completed assistant turns."""
checkpointer = get_checkpointer(request)
config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
try:
checkpoint_tuple = await checkpointer.aget_tuple(config)
except Exception:
logger.exception("Failed to get reference batches for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to get reference batches")
if checkpoint_tuple is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
checkpoint = getattr(checkpoint_tuple, "checkpoint", {}) or {}
values = serialize_channel_values(checkpoint.get("channel_values", {}) or {})
messages = values.get("messages")
if not isinstance(messages, list):
_log_rag_debug(
"backend reference-batches thread=%s messages=missing batches=0 log_file=%s",
sanitize_log_param(thread_id),
_rag_debug_file_path(),
)
return ReferenceBatchesResponse(batches=[])
batches_before_display_adapt = build_runtime_reference_batches(messages)
raw_batches = batches_before_display_adapt
persisted_displays = await _load_persisted_skill_displays()
if persisted_displays:
raw_batches = adapt_reference_batches_with_displays(raw_batches, persisted_displays)
debug_summary = build_reference_debug_summary(messages, raw_batches)
_log_rag_debug(
"backend reference-batches thread=%s messages=%s turns=%s tool_messages=%s raw_items=%s "
"normalized_items=%s kept_items=%s batches=%s batch_sources=%s batch_ids=%s "
"display_adapters=%s turns_detail=%s log_file=%s",
sanitize_log_param(thread_id),
debug_summary["messages"],
debug_summary["turns"],
debug_summary["tool_messages"],
debug_summary["raw_items"],
debug_summary["normalized_items"],
debug_summary["kept_items"],
len(raw_batches),
",".join(str(batch.get("source_count", 0)) for batch in debug_summary["batches"]),
",".join(str(batch.get("id") or "") for batch in debug_summary["batches"]),
len(persisted_displays),
json.dumps(debug_summary["turns_detail"], ensure_ascii=False),
_rag_debug_file_path(),
)
_log_rag_debug(
"backend reference-batches-payload thread=%s %s",
sanitize_log_param(thread_id),
json.dumps(
{
"display_adapter_names": sorted(persisted_displays.keys()),
"before_display_adapt": build_reference_debug_payload(
messages=messages,
batches=batches_before_display_adapt,
),
"after_display_adapt": build_reference_debug_payload(
messages=messages,
batches=raw_batches,
),
},
ensure_ascii=False,
default=str,
),
)
return ReferenceBatchesResponse(
batches=[
ReferenceBatchResponse(
id=str(batch.get("id") or ""),
assistant_message_id=str(batch.get("assistant_message_id") or batch.get("id") or ""),
sources=[
ReferenceSourceResponse(**source)
for source in batch.get("sources", [])
if isinstance(source, dict)
],
)
for batch in raw_batches
]
)
@router.post("/reference-debug-log")
async def write_reference_debug_log(body: ReferenceDebugLogRequest) -> dict[str, str | bool]:
"""Append frontend citation diagnostics to the backend RAG debug log."""
_get_rag_debug_logger().info(
"frontend %s thread=%s payload=%s log_file=%s",
body.event,
sanitize_log_param(body.thread_id or ""),
json.dumps(body.payload, ensure_ascii=False, default=str),
_rag_debug_file_path(),
)
return {"ok": True, "log_file": _rag_debug_file_path()}
@router.post("/{thread_id}/state", response_model=ThreadStateResponse)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def update_thread_state(thread_id: str, body: ThreadStateUpdateRequest, request: Request) -> ThreadStateResponse:
"""Update thread state (e.g. for human-in-the-loop resume or title rename).
Writes a new checkpoint that merges *body.values* into the latest
channel values, then syncs any updated ``title`` field through the
ThreadMetaStore abstraction so that ``/threads/search`` reflects the
change immediately in both sqlite and memory backends.
"""
from app.gateway.deps import get_thread_store
checkpointer = get_checkpointer(request)
thread_store = get_thread_store(request)
# checkpoint_ns must be present in the config for aput — default to ""
# (the root graph namespace). checkpoint_id is optional; omitting it
# fetches the latest checkpoint for the thread.
read_config: dict[str, Any] = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": "",
}
}
if body.checkpoint_id:
read_config["configurable"]["checkpoint_id"] = body.checkpoint_id
try:
checkpoint_tuple = await checkpointer.aget_tuple(read_config)
except Exception:
logger.exception("Failed to get state for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to get thread state")
if checkpoint_tuple is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
# Work on mutable copies so we don't accidentally mutate cached objects.
checkpoint: dict[str, Any] = dict(getattr(checkpoint_tuple, "checkpoint", {}) or {})
metadata: dict[str, Any] = dict(getattr(checkpoint_tuple, "metadata", {}) or {})
channel_values: dict[str, Any] = dict(checkpoint.get("channel_values", {}))
if body.values:
channel_values.update(body.values)
checkpoint["channel_values"] = channel_values
metadata["updated_at"] = now_iso()
if body.as_node:
metadata["source"] = "update"
metadata["step"] = metadata.get("step", 0) + 1
metadata["writes"] = {body.as_node: body.values}
# aput requires checkpoint_ns in the config — use the same config used for the
# read (which always includes checkpoint_ns=""). Do NOT include checkpoint_id
# so that aput generates a fresh checkpoint ID for the new snapshot.
write_config: dict[str, Any] = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": "",
}
}
try:
new_config = await checkpointer.aput(write_config, checkpoint, metadata, {})
except Exception:
logger.exception("Failed to update state for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to update thread state")
new_checkpoint_id: str | None = None
if isinstance(new_config, dict):
new_checkpoint_id = new_config.get("configurable", {}).get("checkpoint_id")
# Sync title changes through the ThreadMetaStore abstraction so /threads/search
# reflects them immediately in both sqlite and memory backends.
if body.values and "title" in body.values:
new_title = body.values["title"]
if new_title: # Skip empty strings and None
try:
await thread_store.update_display_name(thread_id, new_title)
except Exception:
logger.debug("Failed to sync title to thread_meta for %s (non-fatal)", sanitize_log_param(thread_id))
return ThreadStateResponse(
values=serialize_channel_values(channel_values),
next=[],
metadata=metadata,
checkpoint_id=new_checkpoint_id,
created_at=coerce_iso(metadata.get("created_at", "")),
)
# ---------------------------------------------------------------------------
# Lightweight message append / last-AI read
#
# 圆桌广播与「席位交付状态」核对原本都走 GET /state + POST /state,要把动辄几万字的整份
# thread 状态经 loopback 传一两遍 —— 实测让总控每轮 payload_built 飙到 ~11s、广播也很慢。
# 下面两个轻量接口把重活留在服务端,loopback 只传一条小消息 / 截断后的摘要:
# • POST /{id}/messages/append —— 只传 1 条消息,服务端读 checkpoint+追加+写回;
# • GET /{id}/last-ai-message —— 只回最后一条 AI 交付(可 max_chars 截断),核对交付状态用。
# 注意:Postgres checkpoint 的 aput/aget 本身仍是 O(state) 的(快照模型),这里省掉的是
# 「整份 state 的 JSON 序列化 + loopback 往返」,这是可避免的那部分开销。
# ---------------------------------------------------------------------------
class AppendMessageRequest(BaseModel):
"""Request body for the lightweight single-message append."""
content: str = Field(..., description="Message text to append")
type: str = Field(default="human", description="Message role: human/ai/system")
class EditMessageRequest(BaseModel):
"""Request body for editing one persisted chat message."""
content: str = Field(..., description="Replacement message text")
@field_validator("content")
@classmethod
def _content_not_blank(cls, value: str) -> str:
if not value.strip():
raise ValueError("Message content cannot be empty")
return value
def _flatten_msg_content(content: Any) -> str:
"""把 message.content(str / 结构化 block 列表)拍平成纯文本。"""
if isinstance(content, str):
return content
if isinstance(content, list):
parts: list[str] = []
for item in content:
if isinstance(item, str):
parts.append(item)
elif isinstance(item, dict):
parts.append(str(item.get("text", "")))
return "".join(parts)
return str(content) if content else ""
def _set_message_content(message: Any, content: str) -> None:
"""Replace a message's visible content while preserving id/metadata/files."""
if isinstance(message, dict):
message["content"] = content
return
try:
setattr(message, "content", content)
except Exception as exc:
raise TypeError("Message content cannot be edited") from exc
def _msg_type_and_content(msg: Any) -> tuple[str, Any]:
if isinstance(msg, dict):
return str(msg.get("type") or "").lower(), msg.get("content", "")
return str(getattr(msg, "type", "") or "").lower(), getattr(msg, "content", "")
def _is_broadcast_human(text: str) -> bool:
"""圆桌后广播进席位 thread 的 human 消息,不算新一轮任务边界。"""
t = (text or "").strip()
return t.startswith("总控智能体收到") or t.startswith("子智能体")
def _extract_seat_delivery(messages: list[Any]) -> tuple[str, str]:
"""核对席位本轮是否产出**可见正文**。
返回 ``(visible_text, invalid_reason)``。``invalid_reason`` 为空表示已交付;
``thinking_only`` = 最新 AI 只有思考;``empty`` = 本轮尚无可见正文。
从后往前扫:跳过后广播 human、跳过空壳 AI(纯 tool_call);一旦碰到思考-only 的
最新 AI 就停(不继承更早一轮的旧交付);碰到真实派活 human 也停。
"""
for msg in reversed(messages or []):
mtype, content = _msg_type_and_content(msg)
if mtype in ("human", "user"):
if _is_broadcast_human(_flatten_msg_content(content)):
continue
return "", "empty"
if mtype not in ("ai", "aimessage", "aimessagechunk"):
continue
raw = _flatten_msg_content(content)
kind = classify_seat_delivery(raw)
if kind == "valid":
return visible_seat_text(raw), ""
if kind == "thinking_only":
return "", "thinking_only"
return "", "empty"
def _extract_last_ai_text(messages: list[Any]) -> str:
"""返回席位最后一条**可见**交付正文;思考-only / 空串都不算。"""
visible, _reason = _extract_seat_delivery(messages)
return visible
async def _append_message_to_checkpoint(checkpointer: Any, thread_id: str, content: str, type_: str) -> int | None:
"""把一条消息追加进 thread 最新 checkpoint 的 ``messages`` 通道,返回新消息总数。
``None`` 表示线程不存在(无 checkpoint)。纯粹围绕 checkpointer 的 aget/aput,便于单测。
语义与 ``update_thread_state`` 一致:仅替换 ``messages`` 通道为追加后的列表,其它通道原样保留。
"""
config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
ckpt_tuple = await checkpointer.aget_tuple(config)
if ckpt_tuple is None:
return None
checkpoint: dict[str, Any] = dict(getattr(ckpt_tuple, "checkpoint", {}) or {})
metadata: dict[str, Any] = dict(getattr(ckpt_tuple, "metadata", {}) or {})
channel_values: dict[str, Any] = dict(checkpoint.get("channel_values", {}))
messages = list(channel_values.get("messages") or [])
messages.append({"type": type_, "content": content})
channel_values["messages"] = messages
checkpoint["channel_values"] = channel_values
# 生成**新的** checkpoint 快照 id + ts —— 每次追加都落成一个新快照(uuid6 时间可排序,新 id
# 天然成为最新),而不是复用旧 id 原地覆盖(后端语义不一致,会让连续追加互相丢失)。
checkpoint["id"] = str(uuid6(clock_seq=-2))
checkpoint["ts"] = now_iso()
# **关键**:checkpointer 用「版本化 blob」存 channel —— 必须 bump ``messages`` 的版本并把它
# 放进 aput 的 new_versions,aput 才会真正写入该 channel 的新内容、aget 才能读回。否则(像原
# GET/POST /state 那样只改 channel_values、new_versions 传 {})对没有现成 messages 版本的线程
# 不生效。
versions = dict(checkpoint.get("channel_versions", {}))
new_ver = checkpointer.get_next_version(versions.get("messages"), None)
versions["messages"] = new_ver
checkpoint["channel_versions"] = versions
metadata["updated_at"] = now_iso()
await checkpointer.aput(config, checkpoint, metadata, {"messages": new_ver})
return len(messages)
@router.post("/{thread_id}/messages/append", response_model=None)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def append_thread_message(thread_id: str, body: AppendMessageRequest, request: Request) -> dict[str, Any]:
"""轻量追加一条消息到 thread checkpoint —— 调用方只传一条消息,不必 GET+POST 整份 state。
语义与 ``update_thread_state`` 的 messages 写入一致:直接追加到 ``messages`` 通道,下次图
运行时 ``add_messages`` 做规范化/按 id 去重。失败抛 5xx 由调用方按 best-effort 处理。
"""
checkpointer = get_checkpointer(request)
try:
count = await _append_message_to_checkpoint(checkpointer, thread_id, body.content, body.type)
except Exception:
logger.exception("Failed to append message on thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to append thread message")
if count is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
return {"ok": True, "message_count": count}
@router.get("/{thread_id}/last-ai-message", response_model=None)
@require_permission("threads", "read", owner_check=True)
async def get_thread_last_ai_message(thread_id: str, request: Request, max_chars: int = 0) -> dict[str, Any]:
"""轻量读取 thread 最后一条 AI 消息(可截断),供圆桌核对席位交付状态。
返回 ``{delivered, content, length, invalid_reason}``:``delivered`` = 本轮是否存在
**可见正文**(剥掉 ``<think>`` 之后非空;不要求 md 文件);``max_chars>0`` 时 ``content``
截断到该长度(派活轮取摘要),``<=0`` 取全文(综合轮)。``length`` 是可见正文全文长度。
``invalid_reason`` 在未交付时为 ``thinking_only`` / ``empty``,已交付时为空串。
线程不存在时返回 ``delivered=False``(而非 404),方便调用方把"还没建好/没交付"统一当未交付。
"""
checkpointer = get_checkpointer(request)
config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
try:
ckpt_tuple = await checkpointer.aget_tuple(config)
except Exception:
logger.exception("Failed to read last-ai-message on thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to get thread state")
if ckpt_tuple is None:
return {"delivered": False, "content": "", "length": 0, "invalid_reason": "empty"}
channel_values = (getattr(ckpt_tuple, "checkpoint", {}) or {}).get("channel_values", {}) or {}
messages = channel_values.get("messages") or []
text, invalid_reason = _extract_seat_delivery(messages if isinstance(messages, list) else [])
full_len = len(text)
if max_chars and max_chars > 0:
text = text[:max_chars]
return {
"delivered": full_len > 0,
"content": text,
"length": full_len,
"invalid_reason": invalid_reason,
}
@router.get("/{thread_id}/messages", response_model=ThreadMessagesResponse)
@require_permission("threads", "read", owner_check=True)
async def get_thread_messages(
thread_id: str,
request: Request,
limit: int = Query(default=200, ge=1, le=1000, description="Maximum messages to return"),
offset: int = Query(default=0, ge=0, description="Number of leading messages to skip"),
) -> ThreadMessagesResponse:
"""Return a chronological, paginated view of a thread's message transcript.
This reads the latest checkpoint rather than a derived display transcript, so
hidden context messages and all original message fields are retained. The
response serialization is shared with ``GET /state`` and streaming events.
"""
checkpointer = get_checkpointer(request)
config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
try:
checkpoint_tuple = await checkpointer.aget_tuple(config)
except Exception:
logger.exception("Failed to get messages for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to get thread messages")
if checkpoint_tuple is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
checkpoint = getattr(checkpoint_tuple, "checkpoint", {}) or {}
channel_values = checkpoint.get("channel_values", {}) or {}
messages = _serialize_thread_messages(channel_values.get("messages"))
total = len(messages)
return ThreadMessagesResponse(total=total, messages=messages[offset : offset + limit])
@router.delete("/{thread_id}/messages")
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def clear_thread_messages(thread_id: str, request: Request) -> dict[str, Any]:
"""Clear all messages from a thread's latest checkpoint."""
checkpointer = get_checkpointer(request)
config: dict[str, Any] = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": "",
}
}
try:
checkpoint_tuple = await checkpointer.aget_tuple(config)
except Exception:
logger.exception("Failed to get state for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to get thread state")
if checkpoint_tuple is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
checkpoint: dict[str, Any] = dict(getattr(checkpoint_tuple, "checkpoint", {}) or {})
metadata: dict[str, Any] = dict(getattr(checkpoint_tuple, "metadata", {}) or {})
channel_values: dict[str, Any] = dict(checkpoint.get("channel_values", {}))
messages = channel_values.get("messages")
if not isinstance(messages, list):
messages = []
deleted_ids = [mid for m in messages if (mid := _message_id(m))]
channel_values["messages"] = []
checkpoint["channel_values"] = channel_values
checkpoint["id"] = str(uuid6(clock_seq=-2))
checkpoint["ts"] = now_iso()
versions = dict(checkpoint.get("channel_versions", {}))
new_ver = checkpointer.get_next_version(versions.get("messages"), None)
versions["messages"] = new_ver
checkpoint["channel_versions"] = versions
metadata["updated_at"] = now_iso()
try:
await checkpointer.aput(config, checkpoint, metadata, {"messages": new_ver})
except Exception:
logger.exception("Failed to clear messages from thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to clear thread messages")
logger.info(
"Cleared %d message(s) from thread %s",
len(messages),
sanitize_log_param(thread_id),
)
asyncio.create_task(
delete_conversation_deposit_documents(
request.app,
thread_id=thread_id,
message_ids=None,
)
)
return {
"ok": True,
"deleted_ids": deleted_ids,
"deleted_count": len(messages),
"remaining": 0,
}
@router.delete("/{thread_id}/messages/{message_id}")
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def delete_thread_message(thread_id: str, message_id: str, request: Request) -> dict[str, Any]:
"""Delete a conversation turn from a thread's latest checkpoint.
``message_id`` identifies a user question. The whole turn is removed: the
question plus every following message (its answer and any tool/intermediate
messages) up to — but not including — the next user question. Removing a
complete turn keeps the remaining checkpoint valid (no orphaned tool
results / dangling answers) and lets the graph still resume the thread.
If ``message_id`` points at a non-human message, only that single message is
removed.
"""
checkpointer = get_checkpointer(request)
read_config: dict[str, Any] = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": "",
}
}
try:
checkpoint_tuple = await checkpointer.aget_tuple(read_config)
except Exception:
logger.exception("Failed to get state for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to get thread state")
if checkpoint_tuple is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
checkpoint: dict[str, Any] = dict(getattr(checkpoint_tuple, "checkpoint", {}) or {})
metadata: dict[str, Any] = dict(getattr(checkpoint_tuple, "metadata", {}) or {})
channel_values: dict[str, Any] = dict(checkpoint.get("channel_values", {}))
messages = channel_values.get("messages")
if not isinstance(messages, list):
raise HTTPException(status_code=404, detail="Thread has no messages")
target_index = next(
(i for i, m in enumerate(messages) if _message_id(m) == message_id),
None,
)
if target_index is None:
raise HTTPException(status_code=404, detail=f"Message {message_id} not found")
# Delete the whole turn: from the question to just before the next question.
end = len(messages)
if _message_type(messages[target_index]) == "human":
for j in range(target_index + 1, len(messages)):
if _message_type(messages[j]) == "human":
end = j
break
else:
end = target_index + 1
deleted_ids = [
mid for m in messages[target_index:end] if (mid := _message_id(m))
]
filtered = messages[:target_index] + messages[end:]
channel_values["messages"] = filtered
checkpoint["channel_values"] = channel_values
checkpoint["id"] = str(uuid6(clock_seq=-2))
checkpoint["ts"] = now_iso()
versions = dict(checkpoint.get("channel_versions", {}))
new_ver = checkpointer.get_next_version(versions.get("messages"), None)
versions["messages"] = new_ver
checkpoint["channel_versions"] = versions
metadata["updated_at"] = now_iso()
write_config: dict[str, Any] = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": "",
}
}
try:
await checkpointer.aput(write_config, checkpoint, metadata, {"messages": new_ver})
except Exception:
logger.exception(
"Failed to delete message %s from thread %s",
sanitize_log_param(message_id),
sanitize_log_param(thread_id),
)
raise HTTPException(status_code=500, detail="Failed to delete message")
logger.info(
"Deleted turn starting at message %s (%d message(s)) from thread %s",
sanitize_log_param(message_id),
len(messages) - len(filtered),
sanitize_log_param(thread_id),
)
asyncio.create_task(
delete_conversation_deposit_documents(
request.app,
thread_id=thread_id,
message_ids=deleted_ids,
)
)
return {"ok": True, "deleted_ids": deleted_ids, "remaining": len(filtered)}
@router.patch("/{thread_id}/messages/{message_id}")
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def edit_thread_message(
thread_id: str,
message_id: str,
body: EditMessageRequest,
request: Request,
) -> dict[str, Any]:
"""Edit one message in a thread's latest checkpoint.
This is intentionally narrower than ``POST /state``: callers send only the
target id and replacement text, while the server preserves all other
message fields (attachments, timestamps, tool metadata, ids) and bumps the
``messages`` channel version so the edited checkpoint is durable.
"""
checkpointer = get_checkpointer(request)
config: dict[str, Any] = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": "",
}
}
try:
checkpoint_tuple = await checkpointer.aget_tuple(config)
except Exception:
logger.exception("Failed to get state for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to get thread state")
if checkpoint_tuple is None:
raise HTTPException(status_code=404, detail=f"Thread {thread_id} not found")
checkpoint: dict[str, Any] = dict(getattr(checkpoint_tuple, "checkpoint", {}) or {})
metadata: dict[str, Any] = dict(getattr(checkpoint_tuple, "metadata", {}) or {})
channel_values: dict[str, Any] = dict(checkpoint.get("channel_values", {}))
messages = list(channel_values.get("messages") or [])
if not messages:
raise HTTPException(status_code=404, detail="Thread has no messages")
target_index = next(
(i for i, m in enumerate(messages) if _message_id(m) == message_id),
None,
)
if target_index is None:
raise HTTPException(status_code=404, detail=f"Message {message_id} not found")
try:
_set_message_content(messages[target_index], body.content)
except TypeError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
channel_values["messages"] = messages
checkpoint["channel_values"] = channel_values
checkpoint["id"] = str(uuid6(clock_seq=-2))
checkpoint["ts"] = now_iso()
versions = dict(checkpoint.get("channel_versions", {}))
new_ver = checkpointer.get_next_version(versions.get("messages"), None)
versions["messages"] = new_ver
checkpoint["channel_versions"] = versions
metadata["updated_at"] = now_iso()
try:
await checkpointer.aput(config, checkpoint, metadata, {"messages": new_ver})
except Exception:
logger.exception(
"Failed to edit message %s from thread %s",
sanitize_log_param(message_id),
sanitize_log_param(thread_id),
)
raise HTTPException(status_code=500, detail="Failed to edit message")
serialized_messages = serialize_channel_values({"messages": [messages[target_index]]}).get("messages", [])
logger.info(
"Edited message %s in thread %s",
sanitize_log_param(message_id),
sanitize_log_param(thread_id),
)
return {
"ok": True,
"message": serialized_messages[0] if serialized_messages else {"id": message_id, "content": body.content},
}
@router.post("/{thread_id}/history", response_model=list[HistoryEntry])
@require_permission("threads", "read", owner_check=True)
async def get_thread_history(thread_id: str, body: ThreadHistoryRequest, request: Request) -> list[HistoryEntry]:
"""Get checkpoint history for a thread.
Messages are read from the checkpointer's channel values (the
authoritative source) and serialized via
:func:`~deerflow.runtime.serialization.serialize_channel_values`.
Only the latest (first) checkpoint carries the ``messages`` key to
avoid duplicating them across every entry.
"""
checkpointer = get_checkpointer(request)
config: dict[str, Any] = {"configurable": {"thread_id": thread_id}}
if body.before:
config["configurable"]["checkpoint_id"] = body.before
entries: list[HistoryEntry] = []
is_latest_checkpoint = True
try:
async for checkpoint_tuple in checkpointer.alist(config, limit=body.limit):
ckpt_config = getattr(checkpoint_tuple, "config", {})
parent_config = getattr(checkpoint_tuple, "parent_config", None)
metadata = getattr(checkpoint_tuple, "metadata", {}) or {}
checkpoint = getattr(checkpoint_tuple, "checkpoint", {}) or {}
checkpoint_id = ckpt_config.get("configurable", {}).get("checkpoint_id", "")
parent_id = None
if parent_config:
parent_id = parent_config.get("configurable", {}).get("checkpoint_id")
channel_values = checkpoint.get("channel_values", {})
# Build values from checkpoint channel_values
values: dict[str, Any] = {}
if title := channel_values.get("title"):
values["title"] = title
if thread_data := channel_values.get("thread_data"):
values["thread_data"] = thread_data
# Attach messages only to the latest checkpoint entry.
if is_latest_checkpoint:
messages = channel_values.get("messages")
if messages:
values["messages"] = serialize_channel_values({"messages": messages}).get("messages", [])
is_latest_checkpoint = False
# Derive next tasks
tasks_raw = getattr(checkpoint_tuple, "tasks", []) or []
next_tasks = [t.name for t in tasks_raw if hasattr(t, "name")]
# Strip LangGraph internal keys from metadata
user_meta = {k: v for k, v in metadata.items() if k not in ("created_at", "updated_at", "step", "source", "writes", "parents")}
# Keep step for ordering context
if "step" in metadata:
user_meta["step"] = metadata["step"]
entries.append(
HistoryEntry(
checkpoint_id=checkpoint_id,
parent_checkpoint_id=parent_id,
metadata=user_meta,
values=values,
created_at=coerce_iso(metadata.get("created_at", "")),
next=next_tasks,
)
)
except Exception:
logger.exception("Failed to get history for thread %s", sanitize_log_param(thread_id))
raise HTTPException(status_code=500, detail="Failed to get thread history")
return entries