mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-08-31 17:18:44 +00:00
* fix: generate title for interrupted first turn * test(title): cover partial-exchange + dict-form messages Harden the interrupted-run fallback path added in 19fc34fd: - TitleMiddleware._should_generate_title now accepts a lone first-turn user message when allow_partial_exchange=True, so the worker can still derive a title if cancellation lands before any AI chunk is checkpointed. - runtime/runs/worker._ensure_interrupted_title computes the next checkpoint step defensively (treat missing/non-int step as 0) and renames a shadowed ckpt_config local for readability. - Add four unit tests in tests/test_title_middleware_core_logic.py: partial-exchange allows user-only, partial-exchange still respects an existing title, dict-form messages are recognized, and the sync fallback path derives a title from dict-form messages — matching what channel_values stores in the checkpoint. Refs #3859. * fix: persist interrupted-title via channel_versions bump Address PR #3874 review feedback: ``_ensure_interrupted_title`` previously called ``aput(..., new_versions={})``. LangGraph's DB-backed savers (``PostgresSaver`` and the v4 ``SqliteSaver`` blob layout) strip inline ``channel_values`` from ``put`` and only persist blobs for channels named in ``new_versions`` — so the fallback ``title`` channel was dropped on read-back and ``threads_meta.display_name`` stayed ``"Untitled"`` after refresh on those backends. The original in-memory e2e passed because ``InMemorySaver`` keeps the inline snapshot verbatim. Fix mirrors ``_rollback_to_pre_run_checkpoint`` in the same file: bump ``channel_versions["title"]`` (via ``checkpointer.get_next_version`` when available, else int/string fallbacks), persist the new version on the checkpoint, and declare it in ``new_versions`` so the DB savers actually write the blob. Regression coverage in ``tests/test_run_worker_rollback.py``: - ``test_ensure_interrupted_title_bumps_channel_version_and_declares_it_in_new_versions`` — exact ``aput`` invariants: ``new_versions == {"title": 1}``, the written checkpoint's ``channel_versions["title"]`` is bumped, and the pre-existing ``messages`` version is preserved. - ``test_ensure_interrupted_title_bumps_existing_string_version`` — string-shaped prior version (some savers use UUID-style versions); bumped value must differ from the prior, no overwrite-in-place. - ``test_ensure_interrupted_title_skips_when_title_already_set`` — title short-circuit; no extra ``aput``. - ``test_ensure_interrupted_title_returns_none_when_no_checkpoint`` — no checkpoint yet; returns ``None`` without writing. - ``test_ensure_interrupted_title_round_trip_with_real_sqlite_checkpointer`` — full round-trip against a real ``AsyncSqliteSaver`` on a disk-backed DB, then closes and re-opens the saver to simulate a fresh connection. The fallback title must still be present on the second ``aget_tuple``. This is the exact scenario the review flagged. Validated locally with the full backend suite: 5195 passed, 18 skipped. Refs #3859. Addresses review on #3874. * test(worker, title): harden interrupted-title fallback for every saver Defensive coverage on top of the channel_versions fix (commit 05253957), addressing edge cases surfaced during a second-pass review of #3874. Worker: - Extract version bump into ``_bump_channel_version(checkpointer, current)`` with explicit fallbacks for int / float / numeric-string / UUID-shaped string / None / bool, AND a wrap-around defense when the saver's ``get_next_version`` raises or returns an unchanged value. The invariant is: returned version MUST differ from the prior. Without this, a saver bug (or a custom backend) could leave ``new_versions={"title": v}`` no-op on DB savers — the very class of bug the original review pointed out. Title middleware: - Coerce ``state.get("messages")`` from ``None`` to ``[]`` on both ``_should_generate_title`` and ``_build_title_prompt``. A partially-initialized checkpoint can carry ``messages=None`` on the channel_values channel (the worker reads raw channel_values, not BaseMessages), and the default kwarg only protects against a missing key. Repro: ``TypeError: 'NoneType' object is not iterable`` from the next() generator — confirmed by reverting the fix and watching ``test_*_handles_none_messages_channel`` go red. Tests (TDD-verified red→green for the new asserts): - ``test_run_worker_rollback.py``: * ``_bump_channel_version`` — 8 tests covering every version type (int, float, numeric string, UUID-style string, None, bool) and every saver-side fault mode (no ``get_next_version`` / raising / stuck on identity). * ``test_ensure_interrupted_title_*`` — 5 additional helper boundary tests: title.enabled=false short-circuit; empty messages list; messages=None; aput-error propagation (helper contract: caller swallows, not the helper); idempotency on a real InMemorySaver across two invocations. * ``test_ensure_interrupted_title_preserves_non_title_channel_versions`` — pins that ``new_versions`` only contains ``"title"`` and that other channels' versions are untouched (regression anchor for a sloppier draft that bumped every channel). * ``test_worker_finally_block_swallows_helper_exceptions`` — pins the integration contract: even if the helper raises, the worker's threads_meta status sync still runs and ``publish_end`` is still awaited so the SSE stream closes cleanly. - ``test_title_middleware_core_logic.py``: * 4 additional tests: ``messages=None`` on both ``_should_generate_title`` and ``_build_title_prompt``; the ``role: user`` / ``role: assistant`` (OpenAI-style) dict normalization; partial-exchange path with a dict-form message. Verification: - ``PYTHONPATH=. uv run pytest tests/ -x --ignore=tests/blocking_io -q`` → 5215 passed, 18 skipped. - ``ruff check`` + ``ruff format --check`` clean on every touched file. - Red/green TDD verification: temporarily reverted the ``new_versions={}`` fix → 4 new tests went red as expected; restored and the suite is green again. Same red/green dance for the ``messages=None`` coercion. Refs #3859. Addresses second-pass review on #3874. * fix(title): ignore dict context reminders in fallback * fix(worker): link interrupted-title checkpoint to its parent The title-bump checkpoint written by ``_ensure_interrupted_title`` was landing without a ``parent_checkpoint_id`` — a real orphan in the LangGraph history graph. Reproduction (disk-backed AsyncSqliteSaver): [seed] checkpoint_id = 1f173dbc... [helper] wrote title = "Why is the sky blue?" [issue 1] new checkpoint = 1f173dbc..., parent = None [issue 1] is new checkpoint orphaned? True Root cause: ``_ensure_interrupted_title`` built ``write_config`` as ``{"thread_id": ..., "checkpoint_ns": ...}`` only. ``BaseCheckpointSaver`` implementations read ``configurable.checkpoint_id`` from that config as the *parent* id when inserting (see ``langgraph/checkpoint/sqlite/aio.py`` ``aput``: ``config["configurable"].get("checkpoint_id")`` becomes the ``parent_checkpoint_id`` column). With no value, the saver writes NULL — the new checkpoint is a tree root. Consequences: - Any future LangGraph ``runs.resume_from`` / time-travel feature has no backward edge to walk past the title-bump. - History-visualization UIs built on ``alist()`` render the title-bump as a sibling of the prior checkpoint, not its descendant. Fix: read ``checkpoint_id`` off the tuple's own config and thread it into ``write_config["configurable"]["checkpoint_id"]`` before calling ``aput``, the same pattern every middleware-driven write uses. Three new regression tests against real ``AsyncSqliteSaver`` (disk-backed, fresh connections so we exercise the on-disk read path): - ``test_ensure_interrupted_title_links_new_checkpoint_to_its_parent`` — asserts ``latest.parent_config["configurable"]["checkpoint_id"]`` equals the seeded checkpoint id. TDD red-green verified: reverting the fix flips this test red with ``AssertionError: title-bump checkpoint must have a parent_config``. - ``test_ensure_interrupted_title_appears_in_history_with_audit_marker`` — pins the audit contract: the title-bump entry in ``alist()`` carries ``metadata.source == "update"`` and ``metadata.writes`` contains ``runtime_interrupt_title``. This is a deliberate design choice — we do NOT hide the entry from history (audit trail belongs in the saver), but its source and writes marker MUST be unambiguous so UIs/tools can identify it. - ``test_ensure_interrupted_title_survives_immediate_next_turn`` — cancel → immediate user follow-up scenario. Simulates the agent's next turn appending a (user, ai) pair without touching the title channel, then opens a fresh saver and verifies the title is still present after the next-turn checkpoint write. Pins the channel-version-blob invariant established by commit 05253957 — without the ``new_versions={"title": v}`` declaration there, the title blob would vanish from the DB and this test would read back ``None``. Verification: - ``PYTHONPATH=. uv run pytest tests/ -x --ignore=tests/blocking_io -q`` → 5222 passed, 15 skipped. - ``ruff check`` + ``ruff format --check`` clean on every touched file. - Reproduction script confirms ``parent_checkpoint_id`` is now non-null and the next-turn read-back preserves the fallback title. Refs #3859. * Revert "fix(worker): link interrupted-title checkpoint to its parent" This reverts commit c763ed9334781db1acdce0f5f33d663d8d5f80ff. * test: trim over-engineered test coverage Reduce review surface area on PR #3874 by dropping defensive tests that don't pin a real invariant. After self-review: - ``_bump_channel_version``: 8 tests → 2 (happy path + saver-error fallback). Dropped float / bool / numeric-string / UUID-string / missing-get-next-version / stuck-get-next-version branches — those are speculative scaffolding for savers we don't ship. - ``_ensure_interrupted_title``: dropped ``returns_none_when_title_disabled``, ``returns_none_with_no_user_message``, ``returns_none_when_no_checkpoint`` — boundary guards already exercised by the e2e test and the ``handles_none_messages_channel`` regression anchor. Net: -107 lines of test code. Remaining coverage still pins every red-green-verified invariant (channel_versions bump, string-version bump, idempotency, sqlite round-trip, non-title channel preservation, aput-error contract, worker finally swallowing, partial-exchange). Verification: 5209 passed, 15 skipped. * fix: harden interrupted title finalization * fix: serialize interrupted title finalization * fix: preserve interrupt semantics during title finalization * fix: preserve delayed interrupted title recovery --------- Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
232 lines
9.6 KiB
Python
232 lines
9.6 KiB
Python
"""Middleware for automatic thread title generation."""
|
|
|
|
import logging
|
|
import re
|
|
from typing import TYPE_CHECKING, Any, NotRequired, override
|
|
|
|
from langchain.agents import AgentState
|
|
from langchain.agents.middleware import AgentMiddleware
|
|
from langgraph.config import get_config
|
|
from langgraph.constants import TAG_NOSTREAM
|
|
from langgraph.runtime import Runtime
|
|
|
|
from deerflow.agents.middlewares.dynamic_context_middleware import is_dynamic_context_reminder
|
|
from deerflow.config.title_config import get_title_config
|
|
from deerflow.models import create_chat_model
|
|
|
|
if TYPE_CHECKING:
|
|
from deerflow.config.app_config import AppConfig
|
|
from deerflow.config.title_config import TitleConfig
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class TitleMiddlewareState(AgentState):
|
|
"""Compatible with the `ThreadState` schema."""
|
|
|
|
title: NotRequired[str | None]
|
|
|
|
|
|
class TitleMiddleware(AgentMiddleware[TitleMiddlewareState]):
|
|
"""Automatically generate a title for the thread after the first user message."""
|
|
|
|
state_schema = TitleMiddlewareState
|
|
|
|
def __init__(self, *, app_config: "AppConfig | None" = None, title_config: "TitleConfig | None" = None):
|
|
super().__init__()
|
|
self._app_config = app_config
|
|
self._title_config = title_config
|
|
|
|
def _get_title_config(self):
|
|
if self._title_config is not None:
|
|
return self._title_config
|
|
if self._app_config is not None:
|
|
return self._app_config.title
|
|
return get_title_config()
|
|
|
|
def _normalize_content(self, content: object) -> str:
|
|
if isinstance(content, str):
|
|
return content
|
|
|
|
if isinstance(content, list):
|
|
parts = [self._normalize_content(item) for item in content]
|
|
return "\n".join(part for part in parts if part)
|
|
|
|
if isinstance(content, dict):
|
|
text_value = content.get("text")
|
|
if isinstance(text_value, str):
|
|
return text_value
|
|
|
|
nested_content = content.get("content")
|
|
if nested_content is not None:
|
|
return self._normalize_content(nested_content)
|
|
|
|
return ""
|
|
|
|
@staticmethod
|
|
def _message_type(message: object) -> str | None:
|
|
message_type = getattr(message, "type", None)
|
|
if message_type is None and isinstance(message, dict):
|
|
message_type = message.get("type") or message.get("role")
|
|
if message_type == "user":
|
|
return "human"
|
|
if message_type == "assistant":
|
|
return "ai"
|
|
return message_type if isinstance(message_type, str) else None
|
|
|
|
@staticmethod
|
|
def _message_content(message: object) -> object:
|
|
if isinstance(message, dict):
|
|
return message.get("content", "")
|
|
return getattr(message, "content", "")
|
|
|
|
@staticmethod
|
|
def _is_dynamic_context_reminder_message(message: object) -> bool:
|
|
if is_dynamic_context_reminder(message):
|
|
return True
|
|
if isinstance(message, dict):
|
|
additional_kwargs = message.get("additional_kwargs")
|
|
return isinstance(additional_kwargs, dict) and bool(additional_kwargs.get("dynamic_context_reminder"))
|
|
return False
|
|
|
|
@staticmethod
|
|
def _is_user_message_for_title(message: object) -> bool:
|
|
return TitleMiddleware._message_type(message) == "human" and not TitleMiddleware._is_dynamic_context_reminder_message(message)
|
|
|
|
def _get_title_user_message(self, state: TitleMiddlewareState) -> str:
|
|
messages = state.get("messages") or []
|
|
user_msg_content = next((self._message_content(m) for m in messages if self._is_user_message_for_title(m)), "")
|
|
return self._normalize_content(user_msg_content)
|
|
|
|
def _should_generate_title(self, state: TitleMiddlewareState, *, allow_partial_exchange: bool = False) -> bool:
|
|
"""Check if we should generate a title for this thread."""
|
|
config = self._get_title_config()
|
|
if not config.enabled:
|
|
return False
|
|
|
|
# Check if thread already has a title in state
|
|
if state.get("title"):
|
|
return False
|
|
|
|
# Check if this is the first turn (has at least one user message and one assistant response).
|
|
# Defensively coerce a None ``messages`` channel (possible when reading a
|
|
# partially-initialized checkpoint) into an empty list so ``len()`` is safe.
|
|
messages = state.get("messages") or []
|
|
min_messages = 1 if allow_partial_exchange else 2
|
|
if len(messages) < min_messages:
|
|
return False
|
|
|
|
# Count user and assistant messages
|
|
user_messages = [m for m in messages if self._is_user_message_for_title(m)]
|
|
assistant_messages = [m for m in messages if self._message_type(m) == "ai"]
|
|
|
|
# Normal path: title only after first complete exchange. Interrupted path
|
|
# (``allow_partial_exchange=True``) accepts a lone first-turn user message
|
|
# so a fallback title can still be persisted when the run is cancelled
|
|
# before any AI chunk reaches the checkpoint.
|
|
return len(user_messages) == 1 and (len(assistant_messages) >= 1 or allow_partial_exchange)
|
|
|
|
def _build_title_prompt(self, state: TitleMiddlewareState) -> tuple[str, str]:
|
|
"""Extract user/assistant messages and build the title prompt.
|
|
|
|
Returns (prompt_string, user_msg) so callers can use user_msg as fallback.
|
|
"""
|
|
config = self._get_title_config()
|
|
messages = state.get("messages") or []
|
|
|
|
assistant_msg_content = next((self._message_content(m) for m in messages if self._message_type(m) == "ai"), "")
|
|
|
|
user_msg = self._get_title_user_message(state)
|
|
assistant_msg = self._strip_think_tags(self._normalize_content(assistant_msg_content))
|
|
|
|
prompt = config.prompt_template.format(
|
|
max_words=config.max_words,
|
|
user_msg=user_msg[:500],
|
|
assistant_msg=assistant_msg[:500],
|
|
)
|
|
return prompt, user_msg
|
|
|
|
def _strip_think_tags(self, text: str) -> str:
|
|
"""Remove <think>...</think> blocks emitted by reasoning models (e.g. minimax, DeepSeek-R1)."""
|
|
return re.sub(r"<think>[\s\S]*?</think>", "", text, flags=re.IGNORECASE).strip()
|
|
|
|
def _parse_title(self, content: object) -> str:
|
|
"""Normalize model output into a clean title string."""
|
|
config = self._get_title_config()
|
|
title_content = self._normalize_content(content)
|
|
title_content = self._strip_think_tags(title_content)
|
|
title = title_content.strip().strip('"').strip("'")
|
|
return title[: config.max_chars] if len(title) > config.max_chars else title
|
|
|
|
def _fallback_title(self, user_msg: str) -> str:
|
|
config = self._get_title_config()
|
|
fallback_chars = min(config.max_chars, 50)
|
|
if len(user_msg) > fallback_chars:
|
|
return user_msg[:fallback_chars].rstrip() + "..."
|
|
return user_msg if user_msg else "New Conversation"
|
|
|
|
def _get_runnable_config(self) -> dict[str, Any]:
|
|
"""Inherit the parent RunnableConfig and add middleware tag.
|
|
|
|
This ensures RunJournal identifies LLM calls from this middleware
|
|
as ``middleware:title`` instead of ``lead_agent``.
|
|
"""
|
|
try:
|
|
parent = get_config()
|
|
except Exception:
|
|
parent = {}
|
|
config = {**parent}
|
|
config["run_name"] = "title_agent"
|
|
config["tags"] = [
|
|
*(config.get("tags") or []),
|
|
"middleware:title",
|
|
TAG_NOSTREAM,
|
|
]
|
|
return config
|
|
|
|
def _generate_title_result(self, state: TitleMiddlewareState, *, allow_partial_exchange: bool = False) -> dict | None:
|
|
"""Generate a local fallback title without blocking on an LLM call."""
|
|
if not self._should_generate_title(state, allow_partial_exchange=allow_partial_exchange):
|
|
return None
|
|
|
|
user_msg = self._get_title_user_message(state)
|
|
return {"title": self._fallback_title(user_msg)}
|
|
|
|
async def _agenerate_title_result(self, state: TitleMiddlewareState) -> dict | None:
|
|
"""Generate a configured LLM title asynchronously and fall back locally."""
|
|
if not self._should_generate_title(state):
|
|
return None
|
|
|
|
config = self._get_title_config()
|
|
if not config.model_name:
|
|
user_msg = self._get_title_user_message(state)
|
|
return {"title": self._fallback_title(user_msg)}
|
|
|
|
user_msg = self._get_title_user_message(state)
|
|
|
|
try:
|
|
prompt, user_msg = self._build_title_prompt(state)
|
|
# attach_tracing=False because ``_get_runnable_config()`` inherits
|
|
# the graph-level RunnableConfig (set in ``_make_lead_agent``) whose
|
|
# callbacks already carry tracing handlers; binding them again at
|
|
# the model level would emit duplicate spans.
|
|
model_kwargs = {"thinking_enabled": False, "attach_tracing": False}
|
|
if self._app_config is not None:
|
|
model_kwargs["app_config"] = self._app_config
|
|
model = create_chat_model(name=config.model_name, **model_kwargs)
|
|
response = await model.ainvoke(prompt, config=self._get_runnable_config())
|
|
title = self._parse_title(response.content)
|
|
if title:
|
|
return {"title": title}
|
|
except Exception:
|
|
logger.debug("Failed to generate async title; falling back to local title", exc_info=True)
|
|
return {"title": self._fallback_title(user_msg)}
|
|
|
|
@override
|
|
def after_model(self, state: TitleMiddlewareState, runtime: Runtime) -> dict | None:
|
|
return self._generate_title_result(state)
|
|
|
|
@override
|
|
async def aafter_model(self, state: TitleMiddlewareState, runtime: Runtime) -> dict | None:
|
|
return await self._agenerate_title_result(state)
|