Vanzeren c48de5e70b
feat(checkpoint): make delta snapshot_frequency configurable (#4516)
* feat(checkpoint): make delta snapshot_frequency configurable

* fix(config): carry legacy checkpoint_delta_snapshot_frequency with warning

Addresses review on #4516: the rename from the flat
database.checkpoint_delta_snapshot_frequency key to nested
database.checkpoint_delta.snapshot_frequency silently dropped the old
value (pydantic extra="ignore"). Add a before-validator that maps the
legacy key onto the nested one with a deprecation warning (nested key
wins when both are set), plus a CHANGELOG breaking-change note covering
the rename and the 1000 -> 10 default change.

* fix(checkpoint): validate frozen snapshot frequency
2026-07-28 23:21:23 +08:00

2293 lines
102 KiB
Python

"""Background agent execution.
Runs an agent graph inside an ``asyncio.Task``, publishing events to
a :class:`StreamBridge` as they are produced.
Uses ``graph.astream(stream_mode=[...])`` which gives correct full-state
snapshots for ``values`` mode, proper ``{node: writes}`` for ``updates``,
and ``(chunk, metadata)`` tuples for ``messages`` mode.
Note: ``events`` mode is rejected by the gateway — it requires
``graph.astream_events()`` which cannot simultaneously produce ``values``
snapshots. The JS open-source LangGraph API server works around this via
internal checkpoint callbacks that are not exposed in the Python public API.
"""
from __future__ import annotations
import asyncio
import copy
import inspect
import logging
import os
import sys
import threading
import weakref
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from dataclasses import dataclass, field
from datetime import datetime
from functools import lru_cache
from typing import Any, Literal, cast
from langgraph.checkpoint.base import empty_checkpoint
from langgraph.types import Overwrite
from deerflow.agents.goal_state import GoalEvaluation, GoalState
from deerflow.config.app_config import AppConfig
from deerflow.config.database_config import CheckpointChannelMode
from deerflow.runtime.checkpoint_mode import (
aensure_checkpoint_mode_compatible,
inject_checkpoint_mode,
)
from deerflow.runtime.checkpoint_state import (
CheckpointStateAccessor,
build_state_mutation_graph,
graph_reducer_channels,
graph_state_schema,
graph_writable_channels,
)
from deerflow.runtime.context_keys import CURRENT_RUN_PRE_EXISTING_MESSAGE_IDS_KEY
from deerflow.runtime.goal import (
DEFAULT_MAX_GOAL_CONTINUATIONS,
DEFAULT_MAX_NO_PROGRESS_CONTINUATIONS,
GoalWriteConflict,
_call_checkpointer_method,
_is_visible_message,
_message_type,
attach_goal_evaluation,
compute_no_progress_count,
create_goal_evaluator_model,
evaluate_goal_completion,
goal_thread_lock,
latest_visible_assistant_signature,
make_goal_continuation_message,
read_thread_goal,
should_continue_goal,
visible_conversation_signature,
write_thread_goal,
)
from deerflow.runtime.serialization import serialize
from deerflow.runtime.stream_bridge import StreamBridge
from deerflow.runtime.stream_modes import normalize_stream_modes, to_langgraph_stream_modes
from deerflow.runtime.user_context import get_effective_user_id, resolve_runtime_user_id
from deerflow.trace_context import (
DEERFLOW_TRACE_METADATA_KEY,
is_trace_id_from_request_header,
resolve_deerflow_trace_id,
)
from deerflow.tracing import inject_langfuse_metadata
from deerflow.utils.messages import message_to_text
from deerflow.workspace_changes import capture_workspace_snapshot, get_changed_output_paths, record_workspace_changes
from deerflow.workspace_changes.types import WorkspaceSnapshot
from .manager import RunManager, RunRecord, RunStartOutcome
from .naming import resolve_root_run_name
from .schemas import RunStatus
logger = logging.getLogger(__name__)
_checkpoint_locks_guard = threading.Lock()
_checkpoint_locks_by_loop: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, dict[str, asyncio.Lock]] = weakref.WeakKeyDictionary()
@asynccontextmanager
async def _checkpoint_thread_lock(thread_id: str) -> AsyncIterator[None]:
"""Serialize checkpoint mutations for one thread without blocking goal commands."""
loop = asyncio.get_running_loop()
with _checkpoint_locks_guard:
locks = _checkpoint_locks_by_loop.get(loop)
if locks is None:
locks = {}
_checkpoint_locks_by_loop[loop] = locks
lock = locks.get(thread_id)
if lock is None:
lock = asyncio.Lock()
locks[thread_id] = lock
async with lock:
yield
_DELIVERY_RECEIPT_RETRY_DELAYS_SECONDS = (0.1, 0.5)
async def _persist_delivery_receipt(
event_store: Any,
*,
thread_id: str,
run_id: str,
content: dict[str, Any],
) -> bool:
"""Persist a terminal receipt with short bounded retries.
The owning worker still knows the real terminal outcome and renews its
lease while this coroutine runs. Retrying here handles transient event
store failures without handing a successful run to orphan recovery, which
cannot reconstruct either the terminal status or the detailed receipt.
"""
attempts = len(_DELIVERY_RECEIPT_RETRY_DELAYS_SECONDS) + 1
for attempt in range(attempts):
try:
await event_store.put_if_absent(
thread_id=thread_id,
run_id=run_id,
event_type="run.delivery",
category="outputs",
content=content,
)
return True
except Exception:
if attempt == attempts - 1:
logger.warning(
"Failed to persist delivery receipt for run %s after %d attempts; applying terminal delivery semantics without a receipt",
run_id,
attempts,
exc_info=True,
)
return False
delay = _DELIVERY_RECEIPT_RETRY_DELAYS_SECONDS[attempt]
logger.warning(
"Failed to persist delivery receipt for run %s (attempt %d/%d); retrying in %.1fs",
run_id,
attempt + 1,
attempts,
delay,
exc_info=True,
)
await asyncio.sleep(delay)
return False # pragma: no cover - loop always returns
_DELIVERY_INCOMPLETE_ERROR = "Artifact delivery incomplete: no produced output artifact was presented"
_DELIVERY_RECEIPT_FAILED_ERROR = "Artifact delivery verification failed: terminal delivery receipt could not be persisted"
def _empty_delivery_content() -> dict[str, Any]:
return {"presented": 0, "paths": [], "by_tool": {}}
def _presented_path_covers_output(presented_path: str, produced_path: str) -> bool:
presented_path = presented_path.rstrip("/")
return bool(presented_path) and (produced_path == presented_path or produced_path.startswith(f"{presented_path}/"))
def _delivery_content_with_outputs(
content: dict[str, Any],
produced_paths: list[str],
) -> dict[str, Any]:
"""Attach a delivery verdict when this run created or modified outputs."""
if not produced_paths:
return content
presented_paths = content.get("by_tool", {}).get("present_files", [])
matched_paths = [produced_path for produced_path in produced_paths if any(_presented_path_covers_output(presented_path, produced_path) for presented_path in presented_paths)]
satisfied = bool(matched_paths)
return {
**content,
"verification": {
"source": "outputs_changed",
"requirement": "present_files_matches_produced_output",
},
"produced_paths": produced_paths,
"presented_paths": presented_paths,
"matched_paths": matched_paths,
"stage": "presented" if satisfied else ("mismatched" if presented_paths else "not_started"),
"satisfied": satisfied,
}
def _delivery_error(content: dict[str, Any]) -> str | None:
"""Return the terminal error when no changed output was presented."""
if not content.get("produced_paths") or content.get("satisfied") is True:
return None
return _DELIVERY_INCOMPLETE_ERROR
async def _produced_output_paths(
before: WorkspaceSnapshot | None,
*,
thread_id: str,
user_id: str | None,
) -> list[str]:
"""Detect regular output files created or modified by this run."""
if before is None:
return []
try:
after = await capture_workspace_snapshot(thread_id, user_id=user_id, include_text=False)
return get_changed_output_paths(before, after)
except Exception:
logger.warning("Could not detect produced output artifacts for run thread %s", thread_id, exc_info=True)
return []
# Keep this streaming policy separate from middleware write-authorization sets.
_LARGE_FILE_TOOL_NAMES = frozenset({"str_replace", "write_file"})
_LARGE_FILE_TOOL_BATCH_SIZE = 32
@dataclass
class _LargeFileToolChunkBatcher:
"""Batch file-body argument deltas to avoid quadratic browser parsing.
Normal assistant text and non-file tool calls remain token-streamed. Large
file arguments still update progressively, but in bounded batches instead
of forcing the browser to reparse the growing JSON on every model token.
"""
batch_size: int = _LARGE_FILE_TOOL_BATCH_SIZE
tool_names: dict[tuple[str, str, str], str] = field(default_factory=dict)
pending_identity: tuple[str, str, str] | None = None
pending_message: Any | None = None
pending_metadata: dict[str, Any] = field(default_factory=dict)
pending_count: int = 0
def push(self, chunk: Any) -> list[Any]:
if not isinstance(chunk, tuple) or len(chunk) != 2:
return [*self.flush(), chunk]
message, metadata = chunk
message_id = getattr(message, "id", None)
tool_call_chunks = getattr(message, "tool_call_chunks", None)
if not isinstance(message_id, str) or not message_id or not isinstance(tool_call_chunks, list) or len(tool_call_chunks) != 1:
return [*self.flush(), chunk]
tool_chunk = tool_call_chunks[0]
if not isinstance(tool_chunk, dict):
return [*self.flush(), chunk]
index = tool_chunk.get("index")
tool_call_id = tool_chunk.get("id")
if isinstance(index, int):
discriminator = f"index:{index}"
elif isinstance(tool_call_id, str) and tool_call_id:
discriminator = f"id:{tool_call_id}"
else:
discriminator = "single"
raw_namespace = None
if isinstance(metadata, dict):
raw_namespace = metadata.get("langgraph_checkpoint_ns") or metadata.get("checkpoint_ns")
namespace = raw_namespace if isinstance(raw_namespace, str) else ""
identity = (namespace, message_id, discriminator)
name_fragment = tool_chunk.get("name")
tool_name = self.tool_names.get(identity, "")
if tool_name not in _LARGE_FILE_TOOL_NAMES and isinstance(name_fragment, str) and name_fragment:
tool_name += name_fragment
if any(candidate.startswith(tool_name) for candidate in _LARGE_FILE_TOOL_NAMES):
self.tool_names[identity] = tool_name
else:
self.tool_names.pop(identity, None)
# Batching starts only after the accumulated name matches; split or
# incomplete name fragments stream per-chunk until then.
if tool_name not in _LARGE_FILE_TOOL_NAMES:
return [*self.flush(), chunk]
model_copy = getattr(message, "model_copy", None)
if not callable(model_copy):
return [*self.flush(), chunk]
additional_kwargs = getattr(message, "additional_kwargs", None)
sanitized_additional_kwargs = additional_kwargs
if isinstance(additional_kwargs, dict) and ("function_call" in additional_kwargs or "tool_calls" in additional_kwargs):
sanitized_additional_kwargs = {key: value for key, value in additional_kwargs.items() if key not in {"function_call", "tool_calls"}}
has_non_tool_payload = bool(getattr(message, "content", None) or sanitized_additional_kwargs or getattr(message, "usage_metadata", None) or getattr(message, "response_metadata", None))
outputs: list[Any] = []
if self.pending_identity is not None and self.pending_identity != identity:
outputs.extend(self.flush())
if has_non_tool_payload:
visible_message = model_copy(
update={
"additional_kwargs": sanitized_additional_kwargs,
"invalid_tool_calls": [],
"tool_call_chunks": [],
"tool_calls": [],
}
)
outputs.append((visible_message, metadata))
tool_only_message = model_copy(
update={
"additional_kwargs": {},
"content": "",
"invalid_tool_calls": [],
"response_metadata": {},
"tool_calls": [],
"usage_metadata": None,
}
)
self.pending_identity = identity
self.pending_message = tool_only_message if self.pending_message is None else self.pending_message + tool_only_message
if isinstance(metadata, dict):
self.pending_metadata.update(metadata)
self.pending_count += 1
if self.pending_count >= self.batch_size:
outputs.extend(self.flush())
return outputs
def flush(self) -> list[Any]:
if self.pending_message is None:
return []
chunk = (self.pending_message, self.pending_metadata)
self.pending_identity = None
self.pending_message = None
self.pending_metadata = {}
self.pending_count = 0
return [chunk]
def finish(self) -> list[Any]:
"""Flush and release identities at a values or end-of-stream boundary.
A regular batch-size or interleaved-mode flush must retain identities
because continuation chunks commonly omit the tool name.
"""
chunks = self.flush()
self.tool_names.clear()
return chunks
def _build_runtime_context(
thread_id: str,
run_id: str,
caller_context: Any | None,
app_config: AppConfig | None = None,
) -> dict[str, Any]:
"""Build the dict that becomes ``ToolRuntime.context`` for the run.
Always includes ``thread_id`` and ``run_id``. Additional keys from the caller's
``config['context']`` (e.g. ``agent_name`` for the bootstrap flow — issue #2677)
are merged in but never override ``thread_id``/``run_id``. The resolved
``AppConfig`` is added by the worker so tools can consume it without ambient
global lookups.
langgraph 1.1+ surfaces this as ``runtime.context`` via the parent runtime stored
under ``config['configurable']['__pregel_runtime']`` — see
``langgraph.pregel.main`` where ``parent_runtime.merge(...)`` is invoked.
"""
runtime_ctx: dict[str, Any] = {"thread_id": thread_id, "run_id": run_id}
if isinstance(caller_context, dict):
for key, value in caller_context.items():
if key == CURRENT_RUN_PRE_EXISTING_MESSAGE_IDS_KEY:
continue
runtime_ctx.setdefault(key, value)
if app_config is not None:
runtime_ctx["app_config"] = app_config
return runtime_ctx
@dataclass(frozen=True)
class RunContext:
"""Infrastructure dependencies for a single agent run.
Groups checkpointer, store, and persistence-related singletons so that
``run_agent`` (and any future callers) receive one object instead of a
growing list of keyword arguments.
"""
checkpointer: Any
store: Any | None = field(default=None)
event_store: Any | None = field(default=None)
run_events_config: Any | None = field(default=None)
thread_store: Any | None = field(default=None)
app_config: AppConfig | None = field(default=None)
checkpoint_channel_mode: CheckpointChannelMode = "full"
# Delta snapshot cadence frozen at startup; ``None`` means "not frozen in
# this process" (embedded/tests) and resolves to the config default.
checkpoint_snapshot_frequency: int | None = None
on_run_completed: Any | None = field(default=None)
def _install_runtime_context(config: dict, runtime_context: dict[str, Any]) -> None:
existing_context = config.get("context")
if isinstance(existing_context, dict):
existing_context.setdefault("thread_id", runtime_context["thread_id"])
existing_context.setdefault("run_id", runtime_context["run_id"])
if DEERFLOW_TRACE_METADATA_KEY in runtime_context:
existing_context.setdefault(DEERFLOW_TRACE_METADATA_KEY, runtime_context[DEERFLOW_TRACE_METADATA_KEY])
if "app_config" in runtime_context:
existing_context["app_config"] = runtime_context["app_config"]
if CURRENT_RUN_PRE_EXISTING_MESSAGE_IDS_KEY in runtime_context:
existing_context[CURRENT_RUN_PRE_EXISTING_MESSAGE_IDS_KEY] = runtime_context[CURRENT_RUN_PRE_EXISTING_MESSAGE_IDS_KEY]
return
config["context"] = dict(runtime_context)
def _compute_agent_factory_supports_app_config(agent_factory: Any) -> bool:
try:
return "app_config" in inspect.signature(agent_factory).parameters
except (TypeError, ValueError):
return False
@lru_cache(maxsize=128)
def _cached_agent_factory_supports_app_config(agent_factory: Any) -> bool:
return _compute_agent_factory_supports_app_config(agent_factory)
def _agent_factory_supports_app_config(agent_factory: Any) -> bool:
try:
return _cached_agent_factory_supports_app_config(agent_factory)
except TypeError:
# Some callable instances are unhashable; fall back to a direct check.
return _compute_agent_factory_supports_app_config(agent_factory)
class _SubagentEventBuffer:
"""Buffer subagent ``task_*`` step events and flush them in one locked batch (#3779).
The live SSE bridge already forwards these events for real-time display; this
additionally writes them so the subtask card's step history survives a reload.
``RunEventStore.put`` is documented as a low-frequency path — on Postgres each
call opens its own transaction and takes a per-thread advisory lock. A deep
subagent (``general-purpose`` runs up to ``max_turns=150``) emits hundreds of
``task_running`` steps on the hot stream loop, so persisting each with
``put()`` would serialize against the run's own message-batch writer. This
accumulates recognized subagent events and writes them with ``put_batch``,
which acquires the lock once per batch, honoring the store's contract.
Best-effort: a missing store (run_events not configured) or an unrecognized
chunk is a no-op, flush failures are logged but never propagate into the
stream loop, and terminal ``subagent.end`` events flush eagerly so a completed
subagent's step history is durable promptly rather than only at run end.
"""
#: Flush once this many events are buffered, bounding memory and reload lag on
#: a single deep subagent without paying a per-step lock.
FLUSH_THRESHOLD = 25
def __init__(self, event_store: Any | None, thread_id: str, run_id: str) -> None:
self._event_store = event_store
self._thread_id = thread_id
self._run_id = run_id
self._pending: list[dict[str, Any]] = []
async def add(self, chunk: Any) -> None:
"""Buffer one custom stream chunk; flush on a terminal event or threshold."""
if self._event_store is None:
return
# Lazy import: importing deerflow.subagents at module load triggers its
# package __init__ (executor → agents → tools → task_tool), which imports
# back from deerflow.subagents and deadlocks at gateway startup. Deferring
# it to call time (after all modules are loaded) breaks that cycle.
from deerflow.subagents.step_events import subagent_run_event
record = subagent_run_event(chunk)
if record is None:
return
self._pending.append({"thread_id": self._thread_id, "run_id": self._run_id, **record})
if record["event_type"] == "subagent.end" or len(self._pending) >= self.FLUSH_THRESHOLD:
await self.flush()
async def flush(self) -> None:
"""Persist buffered events in one ``put_batch`` call; swallow store errors."""
if self._event_store is None or not self._pending:
return
batch = self._pending
self._pending = []
try:
await self._event_store.put_batch(batch)
except Exception:
# Re-buffer the failed batch (ahead of any events queued since) so a
# transient store error does not silently drop subagent step events.
self._pending = batch + self._pending
logger.warning("Run %s: failed to persist %d subagent step event(s)", self._run_id, len(batch), exc_info=True)
async def run_agent(
bridge: StreamBridge,
run_manager: RunManager,
record: RunRecord,
*,
ctx: RunContext,
agent_factory: Any,
graph_input: dict,
config: dict,
stream_modes: list[str] | None = None,
stream_subgraphs: bool = False,
interrupt_before: list[str] | Literal["*"] | None = None,
interrupt_after: list[str] | Literal["*"] | None = None,
) -> None:
"""Execute an agent in the background, publishing events to *bridge*."""
# Unpack infrastructure dependencies from RunContext.
checkpointer = ctx.checkpointer
store = ctx.store
event_store = ctx.event_store
run_events_config = ctx.run_events_config
thread_store = ctx.thread_store
terminal_status_kwargs = {"persist": False} if event_store is not None else {}
run_id = record.run_id
thread_id = record.thread_id
pre_run_checkpoint_id: str | None = None
pre_run_workspace_snapshot: WorkspaceSnapshot | None = None
workspace_changes_user_id: str | None = None
snapshot_capture_failed = False
llm_error_fallback_message: str | None = None
checkpoint_rollback_completed = False
# Message ids checkpointed *before* this run started. The stream loop uses
# this set to mask out ``deerflow_error_fallback`` markers that belong to
# earlier runs on the same thread — without it, one stale fallback in
# history would mark every subsequent run on this thread as ``error``.
pre_existing_message_ids: set[str] = set()
# Bound agent graph accessor + captured pre-run rollback point; assigned
# inside the try block so the finally rollback path can fork the pre-run
# checkpoint lineage (see below).
accessor: CheckpointStateAccessor | None = None
rollback_point: RollbackPoint | None = None
journal = None
delivery_content: dict[str, Any] | None = None
produced_output_paths: list[str] | None = None
# Journal construction moved ahead of preflight so every terminal run can
# emit a receipt. Completion persistence keeps its prior boundary: before
# #4272 the journal did not exist until preflight had succeeded, so early
# checkpoint failures / cancellation while waiting did not write an empty
# completion snapshot into RunStore.
persist_completion = False
# Buffers subagent step events for batched persistence (#3779); assigned once
# streaming starts and flushed in the finally block. Pre-bound to None so the
# finally is safe even if an exception fires before streaming begins.
subagent_events: _SubagentEventBuffer | None = None
started = False
async def _finish_cancellation(
action: str,
*,
restore_checkpoint: bool = True,
) -> None:
nonlocal checkpoint_rollback_completed
await run_manager.set_finalizing(run_id, True)
if action == "rollback":
await run_manager.set_status(
run_id,
RunStatus.error,
error="Rolled back by user",
**terminal_status_kwargs,
)
if not restore_checkpoint:
return
try:
checkpoint_rollback_completed = await _rollback_to_pre_run_checkpoint(
accessor=accessor,
checkpointer=checkpointer,
thread_id=thread_id,
run_id=run_id,
rollback_point=rollback_point,
snapshot_capture_failed=snapshot_capture_failed,
)
logger.info(
"Run %s rolled back to pre-run checkpoint %s",
run_id,
pre_run_checkpoint_id,
)
except Exception:
logger.warning(
"Run %s cancellation rollback failed",
run_id,
exc_info=True,
)
else:
await run_manager.set_status(
run_id,
RunStatus.interrupted,
**terminal_status_kwargs,
)
logger.info("Run %s was cancelled", run_id)
try:
normalized_stream_modes = normalize_stream_modes(stream_modes)
requested_modes: set[str] = set(normalized_stream_modes)
lg_modes = to_langgraph_stream_modes(normalized_stream_modes)
# Initialize the run-scoped journal before any fallible or cancellable
# preflight work. Every terminal run with an event store must reach the
# shared finally block with a journal available for its run.delivery
# receipt, including checkpoint validation failures and cancellation
# while waiting for an earlier run to finish finalizing.
if event_store is not None:
from deerflow.runtime.journal import RunJournal
journal = RunJournal(
run_id=run_id,
thread_id=thread_id,
event_store=event_store,
track_token_usage=getattr(run_events_config, "track_token_usage", True),
progress_reporter=lambda snapshot: run_manager.update_run_progress(run_id, **snapshot),
)
await run_manager.wait_for_prior_finalizing(
thread_id,
run_id,
abort_event=record.abort_event,
)
start_outcome = await run_manager.try_start(run_id)
if start_outcome is not RunStartOutcome.started:
if record.abort_event.is_set():
await _finish_cancellation(
record.abort_action,
restore_checkpoint=False,
)
return
started = True
if not record.ownership_lost and thread_store is not None:
try:
await thread_store.update_status(thread_id, "running")
except Exception:
logger.debug("Failed to update thread_meta status for %s (non-fatal)", thread_id)
mode = ctx.checkpoint_channel_mode
inject_checkpoint_mode(config, mode)
checkpoint_config = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": "",
}
}
if checkpointer is not None:
await aensure_checkpoint_mode_compatible(
checkpointer,
checkpoint_config,
mode,
)
configurable = config["configurable"]
selected_configurable = {
"thread_id": thread_id,
"checkpoint_ns": configurable.get("checkpoint_ns", ""),
}
for selector_key in ("checkpoint_id", "checkpoint_map"):
if selector_key in configurable:
selected_configurable[selector_key] = configurable[selector_key]
selected_checkpoint_config = {
"configurable": selected_configurable,
}
if selected_checkpoint_config != checkpoint_config:
await aensure_checkpoint_mode_compatible(
checkpointer,
selected_checkpoint_config,
mode,
)
persist_completion = True
if event_store is not None:
workspace_changes_user_id = get_effective_user_id()
try:
pre_run_workspace_snapshot = await capture_workspace_snapshot(
thread_id,
user_id=workspace_changes_user_id,
)
except Exception:
logger.warning("Could not capture pre-run workspace snapshot for run %s", run_id, exc_info=True)
# 2. Publish metadata — useStream needs both run_id AND thread_id
await bridge.publish(
run_id,
"metadata",
{
"run_id": run_id,
"thread_id": thread_id,
},
)
# 3. Build the agent
from langchain_core.runnables import RunnableConfig
from langgraph.runtime import Runtime
# Inject runtime context so middlewares and tools (via ToolRuntime.context) can
# access thread-level data. langgraph-cli does this automatically; we must do it
# manually here because we drive the graph through ``agent.astream(config=...)``
# without passing the official ``context=`` parameter.
runtime_ctx = _build_runtime_context(thread_id, run_id, config.get("context"), ctx.app_config)
incoming_metadata = config.get("metadata") if isinstance(config.get("metadata"), dict) else {}
deerflow_trace_id = resolve_deerflow_trace_id(incoming_metadata.get(DEERFLOW_TRACE_METADATA_KEY))
if deerflow_trace_id:
runtime_ctx[DEERFLOW_TRACE_METADATA_KEY] = deerflow_trace_id
if is_trace_id_from_request_header():
merged_metadata = dict(incoming_metadata)
merged_metadata[DEERFLOW_TRACE_METADATA_KEY] = deerflow_trace_id
config["metadata"] = merged_metadata
# Expose the run-scoped journal under a sentinel key so middleware can
# write audit events (e.g. SafetyFinishReasonMiddleware recording
# suppressed tool calls). Double-underscore prefix marks it as a
# runtime-internal channel; user code must not depend on the key name.
if journal is not None:
runtime_ctx["__run_journal"] = journal
_install_runtime_context(config, runtime_ctx)
runtime = Runtime(context=cast(Any, runtime_ctx), store=store)
config.setdefault("configurable", {})["__pregel_runtime"] = runtime
# Inject RunJournal as a LangChain callback handler.
# on_llm_end captures token usage; on_chain_start/end captures lifecycle.
if journal is not None:
config.setdefault("callbacks", []).append(journal)
# Inject Langfuse trace-attribute metadata so the langchain CallbackHandler
# can lift session_id / user_id / trace_name / tags onto the root trace.
# Shared helper with ``DeerFlowClient.stream`` so both entry points stay
# in sync; caller-provided metadata wins via setdefault inside the helper.
inject_langfuse_metadata(
config,
thread_id=thread_id,
user_id=resolve_runtime_user_id(runtime),
assistant_id=record.assistant_id,
model_name=record.model_name,
environment=os.environ.get("DEER_FLOW_ENV") or os.environ.get("ENVIRONMENT"),
deerflow_trace_id=deerflow_trace_id,
)
# Resolve after runtime context installation so context/configurable reflect
# the agent name that this run will actually execute.
config.setdefault("run_name", resolve_root_run_name(config, record.assistant_id))
initial_runnable_config = RunnableConfig(**config)
def _continuation_runnable_config() -> RunnableConfig:
continuation_config = dict(config)
configurable = dict(continuation_config.get("configurable", {}) or {})
configurable["checkpoint_ns"] = ""
configurable.pop("checkpoint_id", None)
configurable.pop("checkpoint_map", None)
continuation_config["configurable"] = configurable
return RunnableConfig(**continuation_config)
if ctx.app_config is not None and _agent_factory_supports_app_config(agent_factory):
agent = agent_factory(config=initial_runnable_config, app_config=ctx.app_config)
else:
agent = agent_factory(config=initial_runnable_config)
accessor = CheckpointStateAccessor.bind(
agent,
checkpointer,
store=store,
mode=mode,
)
# Capture the pre-run rollback point (materialized state + raw pending
# writes) before this run mutates the thread. Raw checkpoint blobs
# cannot reconstruct Delta-channel messages (their checkpoints omit
# channel_values), so rollback forks the pre-run lineage through the
# graph and needs the materialized messages up front. Any capture
# failure disables rollback: restoring an empty or partial message
# history would silently truncate the thread.
if checkpointer is not None:
# A previous successful run may still be persisting duration
# metadata after its active admission slot is released. Share its
# checkpoint lock so the rollback snapshot and any resume rewrite
# are one uninterrupted read/write sequence against the head.
async with _checkpoint_thread_lock(thread_id):
try:
rollback_point = await _capture_rollback_point(accessor, checkpointer, checkpoint_config)
except Exception:
snapshot_capture_failed = True
logger.warning("Could not capture pre-run checkpoint snapshot for run %s", run_id, exc_info=True)
if rollback_point is not None:
pre_run_checkpoint_id = rollback_point.config.get("configurable", {}).get("checkpoint_id")
pre_existing_message_ids = _collect_pre_existing_message_ids({"messages": list(rollback_point.messages)})
# Resuming from an older checkpoint is a fork, and a delta fork
# materializes the abandoned sibling's writes back into state
# (#4458). Rewrite it as a linear head write *after* the rollback
# point is captured, so cancel-with-rollback still restores the
# real pre-run head rather than the rolled-back one.
resumed_messages = await _linearize_delta_checkpoint_resume(
accessor=accessor,
checkpointer=checkpointer,
config=config,
thread_id=thread_id,
run_id=run_id,
)
if resumed_messages is not None:
# The graph now starts from the selected state, so the
# current-run message boundary is that state, not the head we
# captured for rollback.
pre_existing_message_ids = _collect_pre_existing_message_ids({"messages": list(resumed_messages)})
initial_runnable_config = RunnableConfig(**config)
runtime_ctx[CURRENT_RUN_PRE_EXISTING_MESSAGE_IDS_KEY] = frozenset(pre_existing_message_ids)
_install_runtime_context(config, runtime_ctx)
# Capture the effective (resolved) model name from the agent's metadata.
# _resolve_model_name in agent.py may return the default model if the
# requested name is not in the allowlist — this update ensures the
# persisted model_name reflects the actual model used.
if record.model_name is not None:
resolved = getattr(agent, "metadata", {}) or {}
if isinstance(resolved, dict):
effective = resolved.get("model_name")
if effective and effective != record.model_name:
await run_manager.update_model_name(record.run_id, effective)
# 4. Attach checkpointer and store
if checkpointer is not None:
agent.checkpointer = checkpointer
if store is not None:
agent.store = store
# 5. Set interrupt nodes
if interrupt_before:
agent.interrupt_before_nodes = interrupt_before
if interrupt_after:
agent.interrupt_after_nodes = interrupt_after
logger.info("Run %s: streaming with modes %s (requested: %s)", run_id, lg_modes, requested_modes)
# Buffer subagent step events and persist them in batches (#3779) instead
# of one low-frequency put() per step on the hot stream loop. Flushed in
# the finally block so buffered steps survive abort/exception paths too.
subagent_events = _SubagentEventBuffer(event_store, thread_id, run_id)
goal_evaluator_model: Any | None = None
def _get_goal_evaluator_model() -> Any:
nonlocal goal_evaluator_model
if goal_evaluator_model is None:
goal_evaluator_model = create_goal_evaluator_model(
model_name=record.model_name,
app_config=ctx.app_config,
)
return goal_evaluator_model
async def _stream_once(input_payload: Any, stream_config: RunnableConfig) -> None:
nonlocal llm_error_fallback_message
file_tool_chunk_batcher = _LargeFileToolChunkBatcher() if "values" in requested_modes else None
try:
async with _checkpoint_thread_lock(thread_id):
if len(lg_modes) == 1 and not stream_subgraphs:
# Single mode, no subgraphs: astream yields raw chunks
single_mode = lg_modes[0]
async for chunk in agent.astream(input_payload, config=stream_config, stream_mode=single_mode):
if record.abort_event.is_set():
logger.info("Run %s abort requested — stopping", run_id)
break
llm_error_fallback_message = llm_error_fallback_message or _extract_llm_error_fallback_message(chunk, pre_existing_message_ids)
sse_event = _lg_mode_to_sse_event(single_mode)
await bridge.publish(run_id, sse_event, serialize(chunk, mode=single_mode))
if single_mode == "custom":
await subagent_events.add(chunk)
return
# Multiple modes or subgraphs: astream yields tuples
async for item in agent.astream(
input_payload,
config=stream_config,
stream_mode=lg_modes,
subgraphs=stream_subgraphs,
):
if record.abort_event.is_set():
logger.info("Run %s abort requested — stopping", run_id)
break
mode, chunk, namespace = _unpack_stream_item(item, lg_modes, stream_subgraphs)
if mode is None:
continue
if not namespace:
# Only root-graph frames may decide the parent run's error
# fallback: a delegated subagent's marked fallback is the
# executor's to map (task_failed), not this run's.
llm_error_fallback_message = llm_error_fallback_message or _extract_llm_error_fallback_message(chunk, pre_existing_message_ids)
await _publish_stream_item(
bridge=bridge,
run_id=run_id,
mode=mode,
chunk=chunk,
namespace=namespace,
file_tool_chunk_batcher=file_tool_chunk_batcher,
subagent_events=subagent_events,
)
finally:
stream_error = sys.exception()
if file_tool_chunk_batcher is not None:
try:
for publish_chunk in file_tool_chunk_batcher.finish():
await bridge.publish(run_id, "messages", serialize(publish_chunk, mode="messages"))
except Exception:
if stream_error is None:
raise
logger.debug("Could not flush pending file-tool chunks for run %s", run_id, exc_info=True)
# 7. Stream the requested turn, then optionally continue hidden goal turns.
# Clear any stale stop_reason before the first (user-visible) turn only.
# Continuation turns preserve a cap reason from the user turn: a run that
# hits a cap during the user turn IS capped even if hidden goal-evaluator
# turns complete cleanly afterward (#4176 review).
if isinstance(runtime.context, dict):
runtime.context.pop("stop_reason", None)
await _stream_once(graph_input, initial_runnable_config)
while not record.abort_event.is_set() and not llm_error_fallback_message and (journal is None or not journal.had_llm_error_fallback):
continuation_input = await _prepare_goal_continuation_input(
bridge=bridge,
accessor=accessor,
checkpointer=checkpointer,
thread_id=thread_id,
run_id=run_id,
model_name=record.model_name,
app_config=ctx.app_config,
evaluator_model_factory=_get_goal_evaluator_model,
abort_event=record.abort_event,
user_id=resolve_runtime_user_id(runtime),
deerflow_trace_id=deerflow_trace_id,
)
if continuation_input is None or record.abort_event.is_set():
break
await _stream_once(continuation_input, _continuation_runnable_config())
# 8. Final status
if record.abort_event.is_set():
await _finish_cancellation(record.abort_action)
elif llm_error_fallback_message or (journal is not None and journal.had_llm_error_fallback):
error_msg = llm_error_fallback_message
if error_msg is None and journal is not None:
error_msg = journal.llm_error_fallback_message
error_msg = error_msg or "LLM provider failed after retries"
await _ensure_finalizing_before_edit_failure(run_manager, record)
cancel_action = await run_manager.set_status_if_not_cancelled(
run_id,
RunStatus.error,
error=error_msg,
**terminal_status_kwargs,
)
if cancel_action is not None:
await _finish_cancellation(cancel_action)
else:
runtime_context = runtime.context if isinstance(runtime.context, dict) else None
# Guard middlewares that hard-stop a run by stripping tool_calls
# stamp stop_reason into runtime.context so the worker can surface
# it on the run record:
# loop_detection -> "loop_capped"
# token_budget -> "token_capped"
# safety_finish_reason -> "safety_capped"
# subagent_limit -> "subagent_limit_capped"
# model_length_finish_reason -> "model_length_capped"
#
# If more guards grow stop_reason semantics, consider a publish/
# collect pattern (e.g. each guard middleware publishes its cap
# reason to a dedicated runtime.context channel, and the worker
# collects the most severe / first / all reasons) instead of each
# guard writing directly to the same key.
stop_reason = runtime_context.get("stop_reason") if runtime_context is not None else None
produced_output_paths = await _produced_output_paths(
pre_run_workspace_snapshot,
thread_id=thread_id,
user_id=workspace_changes_user_id,
)
delivery_content = _delivery_content_with_outputs(
journal.get_delivery_content() if journal is not None else _empty_delivery_content(),
produced_output_paths,
)
delivery_error = _delivery_error(delivery_content)
cancel_action = await run_manager.set_status_if_not_cancelled(
run_id,
RunStatus.error if delivery_error else RunStatus.success,
error=delivery_error,
stop_reason=stop_reason,
**terminal_status_kwargs,
)
if cancel_action is not None:
await _finish_cancellation(cancel_action)
except asyncio.CancelledError:
await _finish_cancellation(record.abort_action)
except Exception as exc:
error_msg = f"{exc}"
logger.exception("Run %s failed: %s", run_id, error_msg)
await _ensure_finalizing_before_edit_failure(run_manager, record)
cancel_action = await run_manager.set_status_if_not_cancelled(
run_id,
RunStatus.error,
error=error_msg,
**terminal_status_kwargs,
)
if cancel_action is not None:
await _finish_cancellation(cancel_action)
else:
await bridge.publish(
run_id,
"error",
{
"message": error_msg,
"name": type(exc).__name__,
},
)
finally:
if record.ownership_lost:
logger.warning(
"Skipping durable finalization for run %s because this worker no longer owns its lease",
run_id,
)
if not record.ownership_lost and _is_edit_replay_run(record) and record.status != RunStatus.success:
if not record.finalizing:
await run_manager.set_finalizing(run_id, True)
try:
if not checkpoint_rollback_completed:
checkpoint_rollback_completed = await _rollback_to_pre_run_checkpoint(
accessor=accessor,
checkpointer=checkpointer,
thread_id=thread_id,
run_id=run_id,
rollback_point=rollback_point,
snapshot_capture_failed=snapshot_capture_failed,
)
if checkpoint_rollback_completed:
await _publish_restored_checkpoint_values(
bridge=bridge,
run_id=run_id,
accessor=accessor,
thread_id=thread_id,
)
logger.info("Run %s edit replay restored pre-run checkpoint %s", run_id, pre_run_checkpoint_id)
except Exception:
logger.warning("Run %s edit replay rollback failed", run_id, exc_info=True)
# Persist any subagent step events still buffered (#3779) — including on
# abort/exception paths, where the stream loop broke before its own flush.
if not record.ownership_lost and subagent_events is not None:
await subagent_events.flush()
if not record.ownership_lost and event_store is not None and pre_run_workspace_snapshot is not None:
try:
await record_workspace_changes(
event_store,
thread_id,
run_id,
pre_run_workspace_snapshot,
user_id=workspace_changes_user_id,
)
except Exception:
logger.warning("Failed to record workspace changes for run %s", run_id, exc_info=True)
# Flush buffered journal events before the terminal receipt. The
# receipt uses a run-scoped idempotent write shared with recovery, then
# the staged terminal status is persisted. This ordering closes the
# crash window where a terminal run could otherwise outlive its receipt.
# A fenced worker leaves receipt recovery to the peer that claimed it.
if not record.ownership_lost and journal is not None:
try:
await journal.flush()
except Exception:
logger.warning("Failed to flush journal for run %s", run_id, exc_info=True)
if delivery_content is None:
if produced_output_paths is None:
produced_output_paths = await _produced_output_paths(
pre_run_workspace_snapshot,
thread_id=thread_id,
user_id=workspace_changes_user_id,
)
delivery_content = _delivery_content_with_outputs(journal.get_delivery_content(), produced_output_paths)
receipt_persisted = await _persist_delivery_receipt(
event_store,
thread_id=thread_id,
run_id=run_id,
content=delivery_content,
)
if produced_output_paths and record.status == RunStatus.success and not receipt_persisted:
await run_manager.set_status(
run_id,
RunStatus.error,
error=_DELIVERY_RECEIPT_FAILED_ERROR,
persist=False,
)
if not record.ownership_lost and event_store is not None:
try:
# Even after bounded receipt retries are exhausted, persist the
# real worker outcome. Leaving a successful row inflight would
# let lease recovery rewrite it as an error with a synthetic
# zero receipt.
if record.abort_event.is_set():
await run_manager.persist_current_status(run_id)
else:
cancel_action = await run_manager.set_status_if_not_cancelled(
run_id,
record.status,
error=record.error,
stop_reason=record.stop_reason,
)
if cancel_action is not None:
await _finish_cancellation(cancel_action)
await run_manager.persist_current_status(run_id)
except Exception:
logger.warning("Failed to persist terminal status for run %s after delivery receipt attempts", run_id, exc_info=True)
if not record.ownership_lost and journal is not None and persist_completion:
try:
# Persist token usage + convenience fields to RunStore
completion = journal.get_completion_data()
await run_manager.update_run_completion(run_id, status=record.status.value, **completion)
except Exception:
logger.warning("Failed to persist run completion for %s (non-fatal)", run_id, exc_info=True)
if started and not record.ownership_lost and checkpointer is not None and record.status == RunStatus.interrupted and not _is_edit_replay_run(record):
try:
await run_manager.wait_for_prior_finalizing(thread_id, run_id)
if not await run_manager.has_later_started_run(thread_id, run_id):
await _ensure_interrupted_title(checkpointer=checkpointer, thread_id=thread_id, app_config=ctx.app_config, graph_input=graph_input)
except Exception:
logger.debug("Failed to generate interrupted title for thread %s (non-fatal)", thread_id)
# Sync title from checkpoint to threads_meta.display_name
if started and not record.ownership_lost and checkpointer is not None and thread_store is not None:
try:
ckpt_config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
ckpt_tuple = await checkpointer.aget_tuple(ckpt_config)
if ckpt_tuple is not None:
ckpt = getattr(ckpt_tuple, "checkpoint", {}) or {}
title = ckpt.get("channel_values", {}).get("title")
if title:
await thread_store.update_display_name(thread_id, title)
except Exception:
logger.debug("Failed to sync title for thread %s (non-fatal)", thread_id)
# Persist run duration to checkpoint metadata so history reads
# don't need to correlate runs and events.
if started and not record.ownership_lost and checkpointer is not None and record.status == RunStatus.success:
try:
created = datetime.fromisoformat(record.created_at.replace("Z", "+00:00"))
updated = datetime.fromisoformat(record.updated_at.replace("Z", "+00:00"))
# Match legacy history semantics: turn_duration is the whole
# RunRecord lifetime in integer seconds, including admission
# delay. Persist zero for sub-second successful turns.
duration = max(0, int((updated - created).total_seconds()))
await _persist_run_duration(
checkpointer=checkpointer,
thread_id=thread_id,
run_id=run_id,
duration_seconds=duration,
)
except Exception:
logger.debug("Failed to persist run duration for thread %s run %s (non-fatal)", thread_id, run_id)
# Update threads_meta status based on run outcome
if started and not record.ownership_lost and thread_store is not None:
try:
final_status = "idle" if record.status == RunStatus.success else record.status.value
await thread_store.update_status(thread_id, final_status)
except Exception:
logger.debug("Failed to update thread_meta status for %s (non-fatal)", thread_id)
if not record.ownership_lost and ctx.on_run_completed is not None:
try:
await ctx.on_run_completed(record)
except Exception:
logger.warning("Run completion hook failed for %s (non-fatal)", run_id, exc_info=True)
if record.finalizing:
await run_manager.set_finalizing(run_id, False)
await bridge.publish_end(run_id)
asyncio.create_task(bridge.cleanup(run_id, delay=60))
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _checkpoint_id(checkpoint_tuple: Any) -> str | None:
config = getattr(checkpoint_tuple, "config", {}) or {}
configurable = config.get("configurable", {}) if isinstance(config, dict) else {}
checkpoint_id = configurable.get("checkpoint_id") if isinstance(configurable, dict) else None
if isinstance(checkpoint_id, str):
return checkpoint_id
checkpoint = getattr(checkpoint_tuple, "checkpoint", {}) or {}
if isinstance(checkpoint, dict) and isinstance(checkpoint.get("id"), str):
return checkpoint["id"]
return None
def _goal_instance_matches(left: GoalState | None, right: GoalState | None) -> bool:
if not left or not right:
return False
same_status = left.get("status") == right.get("status") == "active"
same_objective = left.get("objective") == right.get("objective")
same_created_at = left.get("created_at") == right.get("created_at")
return same_status and same_objective and same_created_at
async def _materialized_checkpoint_messages(accessor: CheckpointStateAccessor, thread_id: str) -> list[Any]:
"""Read ``messages`` through the mode-matched accessor.
Raw ``channel_values`` reads see a sentinel in delta mode; only a
materialized read reconstructs the list. Raw checkpoint tuples remain
valid for tuple-level metadata (checkpoint id, ``pending_writes``).
"""
snapshot = await accessor.aget({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}})
values = getattr(snapshot, "values", None) or {}
messages = values.get("messages") if isinstance(values, dict) else None
return list(messages) if isinstance(messages, list) else []
def _read_checkpoint_goal(checkpoint_tuple: Any) -> GoalState | None:
checkpoint = getattr(checkpoint_tuple, "checkpoint", {}) or {}
channel_values = checkpoint.get("channel_values", {}) if isinstance(checkpoint, dict) else {}
raw_goal = channel_values.get("goal") if isinstance(channel_values, dict) else None
return copy.deepcopy(raw_goal) if isinstance(raw_goal, dict) else None
def _has_durable_goal_turn_receipt(checkpoint_tuple: Any, messages: list[Any]) -> bool:
"""Return true when a completed visible assistant turn is safely checkpointed.
``pending_writes`` is the durability signal: a ``CheckpointTuple`` carries no
``tasks`` field (those live on a ``StateSnapshot``), so the presence of any
queued writes is what tells us the turn is still in flight.
"""
if _checkpoint_id(checkpoint_tuple) is None:
return False
if getattr(checkpoint_tuple, "pending_writes", None):
return False
visible_messages = []
for message in messages:
if _is_visible_message(message) and message_to_text(message).strip():
visible_messages.append(message)
if not visible_messages:
return False
return _message_type(visible_messages[-1]) == "ai"
def _stand_down_reason(goal: GoalState, evaluation: GoalEvaluation, no_progress_count: int) -> str | None:
if evaluation["satisfied"]:
return None
if evaluation["blocker"] != "goal_not_met_yet":
return f"blocked:{evaluation['blocker']}"
# Default caps mirror should_continue_goal so the two gate functions agree on
# a goal dict that is missing these fields.
if int(goal.get("continuation_count", 0)) >= int(goal.get("max_continuations", DEFAULT_MAX_GOAL_CONTINUATIONS)):
return "max_continuations_reached"
if no_progress_count >= int(goal.get("max_no_progress_continuations", DEFAULT_MAX_NO_PROGRESS_CONTINUATIONS)):
return "no_progress_detected"
return None
async def _persist_goal_evaluation(
*,
bridge: StreamBridge,
checkpointer: Any,
thread_id: str,
run_id: str,
goal: GoalState,
evaluation: GoalEvaluation,
no_progress_count: int,
continuation_count: int | None = None,
stand_down_reason: str | None = None,
evidence_signature: str = "",
) -> GoalState | None:
try:
async with goal_thread_lock(thread_id):
checkpoint_tuple = await _call_checkpointer_method(
checkpointer,
"aget_tuple",
"get_tuple",
{"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}},
)
if checkpoint_tuple is None:
return None
current_goal = _read_checkpoint_goal(checkpoint_tuple)
if current_goal is None or not _goal_instance_matches(goal, current_goal):
return None
# Defensive: compute continuation_count from the fresh current_goal
# inside the lock. The caller computed it from a possibly-stale goal
# snapshot; a racing continuation may have already bumped the count.
if continuation_count is not None:
current_count = int(current_goal.get("continuation_count", 0))
continuation_count = max(continuation_count, current_count + 1)
expected_checkpoint_id = _checkpoint_id(checkpoint_tuple)
updated_goal = attach_goal_evaluation(
current_goal,
evaluation,
run_id=run_id,
continuation_count=continuation_count,
no_progress_count=no_progress_count,
stand_down_reason=stand_down_reason,
evidence_signature=evidence_signature,
)
values = await write_thread_goal(
checkpointer,
thread_id,
updated_goal,
as_node="goal_evaluator",
expected_checkpoint_id=expected_checkpoint_id,
)
await bridge.publish(run_id, "values", serialize(values, mode="values"))
return updated_goal
except GoalWriteConflict:
return None
except Exception:
logger.warning("Could not persist goal evaluation for thread %s", thread_id, exc_info=True)
return None
async def _reread_goal_and_checkpoint(checkpointer: Any, thread_id: str) -> tuple[GoalState | None, Any]:
"""Re-read the goal and latest checkpoint together for a concurrency re-check."""
goal = await read_thread_goal(checkpointer, thread_id)
checkpoint_tuple = await _call_checkpointer_method(
checkpointer,
"aget_tuple",
"get_tuple",
{"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}},
)
return goal, checkpoint_tuple
async def _prepare_goal_continuation_input(
*,
bridge: StreamBridge,
accessor: CheckpointStateAccessor,
checkpointer: Any,
thread_id: str,
run_id: str,
model_name: str | None,
app_config: AppConfig | None,
evaluator_model_factory: Any | None = None,
abort_event: asyncio.Event | None = None,
user_id: str | None = None,
deerflow_trace_id: str | None = None,
) -> dict[str, Any] | None:
"""Evaluate the active goal and return a hidden continuation input if needed.
NOTE: The re-reads below catch a racing user message or ``/goal clear``
before we queue a continuation. Goal writes then serialize per thread and
pass the checkpoint id they read from, so stale evaluator writes stand down
instead of clobbering a newer goal change.
"""
if checkpointer is None:
return None
if abort_event is not None and abort_event.is_set():
return None
try:
goal = await read_thread_goal(checkpointer, thread_id)
except Exception:
logger.warning("Could not read goal for thread %s after run %s", thread_id, run_id, exc_info=True)
return None
if not goal or goal.get("status") != "active":
return None
async def _persist(
goal: GoalState,
evaluation: GoalEvaluation,
no_progress_count: int,
*,
stand_down_reason: str | None = None,
continuation_count: int | None = None,
) -> GoalState | None:
"""Record the evaluation against the still-current goal instance."""
return await _persist_goal_evaluation(
bridge=bridge,
checkpointer=checkpointer,
thread_id=thread_id,
run_id=run_id,
goal=goal,
evaluation=evaluation,
no_progress_count=no_progress_count,
continuation_count=continuation_count,
stand_down_reason=stand_down_reason,
evidence_signature=evidence_signature,
)
try:
checkpoint_tuple = await _call_checkpointer_method(
checkpointer,
"aget_tuple",
"get_tuple",
{"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}},
)
if checkpoint_tuple is None:
return None
checkpoint_id_before = _checkpoint_id(checkpoint_tuple)
messages = await _materialized_checkpoint_messages(accessor, thread_id)
conversation_signature_before = visible_conversation_signature(messages)
evidence_signature = latest_visible_assistant_signature(messages)
if not _has_durable_goal_turn_receipt(checkpoint_tuple, messages):
evaluation = GoalEvaluation(
satisfied=False,
blocker="run_failed",
reason="No durable assistant end-of-turn receipt was available.",
evidence_summary="",
)
no_progress_count = compute_no_progress_count(goal, evaluation, evidence_signature=evidence_signature)
await _persist(goal, evaluation, no_progress_count, stand_down_reason="no_durable_end_of_turn")
return None
if abort_event is not None and abort_event.is_set():
return None
evaluator_model = evaluator_model_factory() if evaluator_model_factory is not None else None
evaluation = await evaluate_goal_completion(
goal,
messages,
model=evaluator_model,
model_name=model_name,
app_config=app_config,
thread_id=thread_id,
user_id=user_id,
deerflow_trace_id=deerflow_trace_id,
)
if abort_event is not None and abort_event.is_set():
return None
except Exception:
logger.warning("Goal evaluator failed for thread %s after run %s", thread_id, run_id, exc_info=True)
return None
no_progress_count = compute_no_progress_count(goal, evaluation, evidence_signature=evidence_signature)
# Re-check that neither the goal nor the visible conversation changed while the
# evaluator ran — a user message or /goal clear racing the evaluation must win.
try:
current_goal, current_checkpoint_tuple = await _reread_goal_and_checkpoint(checkpointer, thread_id)
except Exception:
logger.warning("Could not re-check goal state for thread %s after evaluation", thread_id, exc_info=True)
return None
if not _goal_instance_matches(goal, current_goal) or current_checkpoint_tuple is None:
return None
checkpoint_changed = _checkpoint_id(current_checkpoint_tuple) != checkpoint_id_before
messages_changed = visible_conversation_signature(await _materialized_checkpoint_messages(accessor, thread_id)) != conversation_signature_before
if checkpoint_changed or messages_changed:
await _persist(current_goal, evaluation, no_progress_count, stand_down_reason="thread_changed_after_evaluation")
return None
if evaluation["satisfied"]:
try:
async with goal_thread_lock(thread_id):
latest_checkpoint_tuple = await _call_checkpointer_method(
checkpointer,
"aget_tuple",
"get_tuple",
{"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}},
)
if latest_checkpoint_tuple is None:
return None
latest_goal = _read_checkpoint_goal(latest_checkpoint_tuple)
if latest_goal is None or not _goal_instance_matches(goal, latest_goal):
return None
values = await write_thread_goal(
checkpointer,
thread_id,
None,
as_node="goal_evaluator",
expected_checkpoint_id=_checkpoint_id(latest_checkpoint_tuple),
)
await bridge.publish(run_id, "values", serialize(values, mode="values"))
except GoalWriteConflict:
return None
except Exception:
logger.warning("Could not clear satisfied goal for thread %s", thread_id, exc_info=True)
return None
stand_down_reason = _stand_down_reason(goal, evaluation, no_progress_count)
if stand_down_reason is not None or not should_continue_goal(goal, evaluation, no_progress_count=no_progress_count):
await _persist(goal, evaluation, no_progress_count, stand_down_reason=stand_down_reason)
return None
next_count = int(goal.get("continuation_count", 0)) + 1
updated_goal = await _persist(goal, evaluation, no_progress_count, continuation_count=next_count)
if updated_goal is None:
return None
# Final guard: the persist above bumped the checkpoint id, so only the visible
# conversation signature is meaningful for detecting a racing user turn here.
try:
latest_goal, latest_checkpoint_tuple = await _reread_goal_and_checkpoint(checkpointer, thread_id)
except Exception:
logger.warning("Could not verify queued goal continuation for thread %s", thread_id, exc_info=True)
return None
if not _goal_instance_matches(updated_goal, latest_goal) or latest_checkpoint_tuple is None:
return None
if visible_conversation_signature(await _materialized_checkpoint_messages(accessor, thread_id)) != conversation_signature_before:
# Do not pass continuation_count here: the persist above already
# committed it (as next_count). Re-passing next_count would make
# _persist_goal_evaluation's race guard (#4088) see that same write as
# a "current_count" bump and add another +1 on top of it, silently
# double-counting this single continuation attempt against the
# continuation budget even though it is being stood down, not
# delivered. Omitting it leaves the already-committed count untouched,
# matching every other stand-down call site in this function.
await _persist(
latest_goal,
evaluation,
no_progress_count,
stand_down_reason="thread_changed_before_continuation",
)
return None
logger.info(
"Run %s continuing thread %s for active goal (%d/%d)",
run_id,
thread_id,
updated_goal.get("continuation_count", next_count),
updated_goal.get("max_continuations", 0),
)
return {"messages": [make_goal_continuation_message(updated_goal, evaluation)]}
def _is_edit_replay_run(record: RunRecord) -> bool:
metadata = record.metadata or {}
return metadata.get("replay_kind") == "edit"
async def _ensure_finalizing_before_edit_failure(run_manager: RunManager, record: RunRecord) -> None:
if _is_edit_replay_run(record) and not record.finalizing:
await run_manager.set_finalizing(record.run_id, True)
async def _publish_restored_checkpoint_values(
*,
bridge: StreamBridge,
run_id: str,
accessor: CheckpointStateAccessor | None,
thread_id: str,
) -> None:
if accessor is None:
return
snapshot = await accessor.aget({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}})
values = getattr(snapshot, "values", None)
if isinstance(values, dict):
await bridge.publish(run_id, "values", serialize(values, mode="values"))
@dataclass(frozen=True)
class RollbackPoint:
"""Materialized pre-run state used to restore the thread after cancellation.
Raw checkpoint blobs cannot reconstruct Delta-channel messages (their
checkpoints omit the materialized value), so rollback preserves those
messages plus delta mode's materialized non-message state in addition to
the raw pending writes.
"""
config: dict[str, Any]
state_values: dict[str, Any]
messages: tuple[Any, ...]
metadata: dict[str, Any]
pending_writes: tuple[tuple[str, str, Any], ...]
async def _capture_rollback_point(
accessor: CheckpointStateAccessor,
checkpointer: Any,
read_config: dict[str, Any],
) -> RollbackPoint | None:
"""Materialize the pre-run checkpoint state and its raw pending writes.
Returns ``None`` when the thread has no checkpoint yet; the caller keeps
the existing delete/reset rollback contract for that case.
"""
snapshot = await accessor.aget(read_config)
snapshot_config = getattr(snapshot, "config", None) or {}
configurable = snapshot_config.get("configurable") or {}
if not configurable.get("checkpoint_id"):
return None
checkpoint_tuple = await _call_checkpointer_method(checkpointer, "aget_tuple", "get_tuple", snapshot_config)
raw_values = getattr(snapshot, "values", None) or {}
messages = raw_values.get("messages") if isinstance(raw_values, dict) else None
state_values = copy.deepcopy({key: value for key, value in raw_values.items() if key != "messages"}) if accessor.mode == "delta" and isinstance(raw_values, dict) else {}
return RollbackPoint(
config={
"configurable": {
"thread_id": configurable.get("thread_id"),
"checkpoint_ns": configurable.get("checkpoint_ns") or "",
"checkpoint_id": configurable.get("checkpoint_id"),
}
},
state_values=state_values,
messages=tuple(messages or ()),
metadata=dict(getattr(snapshot, "metadata", None) or {}),
pending_writes=tuple(getattr(checkpoint_tuple, "pending_writes", ()) or ()),
)
def _complete_state_replacement_values(
*,
mutation_graph: Any,
selected_values: dict[str, Any],
current_values: dict[str, Any],
run_id: str,
operation: str,
) -> dict[str, Any]:
"""Build a whole-state replacement through the graph's effective schema."""
writable_fields = graph_writable_channels(mutation_graph)
reducer_fields = graph_reducer_channels(mutation_graph)
if writable_fields is None or reducer_fields is None:
raise RuntimeError(f"Run {run_id} could not inspect the state schema for {operation}")
replacement_values: dict[str, Any] = {}
for field_name in writable_fields:
if field_name in selected_values:
replacement = copy.deepcopy(selected_values[field_name])
elif field_name in current_values:
# LangGraph has no public "unset channel" update. A fresh channel
# exposes its schema default when one exists (for example [] / {});
# optional and otherwise-unconstructible channels reset to None.
channel = mutation_graph.channels.get(field_name)
replacement = copy.deepcopy(channel.get()) if channel is not None and channel.is_available() else None
else:
continue
replacement_values[field_name] = Overwrite(replacement) if field_name in reducer_fields else replacement
return replacement_values
async def _linearize_delta_checkpoint_resume(
*,
accessor: CheckpointStateAccessor,
checkpointer: Any,
config: dict[str, Any],
thread_id: str,
run_id: str,
) -> list[Any] | None:
"""Replace a delta-mode checkpoint fork with an equivalent linear write.
Resuming from an older checkpoint forks the lineage, and in ``delta`` mode
the fork's state cannot be materialized correctly: the delta history walk
collects **every** ``pending_writes`` entry stored on each on-path
ancestor, but a shared parent also carries the writes of the sibling child
that was abandoned. Those writes are replayed into the fork, so the run
starts from a message list that still contains the answer it was supposed
to replace — regenerating in a branched thread surfaced this as the old
assistant message reappearing beside the new one after a reload (#4458).
Reproduced on postgres, sqlite, and the in-memory saver; ``full`` mode is
unaffected because its checkpoints carry complete ``channel_values`` and
need no replay.
The upstream contract (`BaseCheckpointSaver.get_delta_channel_history` and
the savers overriding it) is where write-to-child ownership belongs, so
this does not reimplement it. Instead the fork is expressed as what it
means: materialize the requested checkpoint's state and write it with
replace semantics on the **current head**, which has no other children,
then run linearly. Every materialized channel is restored; channels that
exist only on the newer head are reset to their schema default (or
``None`` when the channel has no constructible default). The abandoned
turn stays in checkpoint history as the rewritten head's ancestry.
Returns the materialized messages when the resume was linearized, or
``None`` when there was nothing to do (full mode, no checkpoint selector,
a non-root namespace, or a selector that already names the head). Failures
propagate: silently falling back to the fork would persist the corrupted
history this exists to prevent. The worker call site holds
``_checkpoint_thread_lock`` across rollback capture and this rewrite; do
not reacquire that non-reentrant lock inside this helper.
"""
if checkpointer is None or accessor.mode != "delta":
return None
configurable = config.get("configurable")
if not isinstance(configurable, dict):
return None
checkpoint_id = configurable.get("checkpoint_id")
if not isinstance(checkpoint_id, str) or not checkpoint_id:
return None
if configurable.get("checkpoint_ns"):
# Subgraph namespaces have their own lineage; the Gateway only selects
# root checkpoints, so leave anything else untouched.
return None
head_config: dict[str, Any] = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
head = await accessor.aget(head_config)
if _checkpoint_id(head) == checkpoint_id:
# Selecting the head is already linear — no sibling can exist yet.
return None
source_config: dict[str, Any] = {"configurable": {"thread_id": thread_id, "checkpoint_ns": "", "checkpoint_id": checkpoint_id}}
snapshot = await accessor.aget(source_config)
values = getattr(snapshot, "values", None) or {}
messages = values.get("messages") if isinstance(values, dict) else None
if not isinstance(messages, list):
raise RuntimeError(f"Run {run_id} could not materialize resume checkpoint {checkpoint_id}")
# Write through the thread's effective schema so every application and
# middleware channel can be restored. Reducer channels need Overwrite to
# replace their already-aggregated value instead of merging it again.
mutation_graph = build_state_mutation_graph("checkpoint_resume", accessor.mode, graph_state_schema(getattr(accessor, "graph", None)))
selected_values = dict(values)
head_values = getattr(head, "values", None) or {}
head_values = dict(head_values) if isinstance(head_values, dict) else {}
replacement_values = _complete_state_replacement_values(
mutation_graph=mutation_graph,
selected_values=selected_values,
current_values=head_values,
run_id=run_id,
operation="checkpoint resume",
)
mutation_accessor = CheckpointStateAccessor.bind(mutation_graph, checkpointer, mode=accessor.mode)
await mutation_accessor.aupdate(head_config, replacement_values, as_node="checkpoint_resume")
configurable.pop("checkpoint_id", None)
configurable.pop("checkpoint_map", None)
logger.info("Run %s linearized a delta-mode resume of checkpoint %s onto thread %s", run_id, checkpoint_id, thread_id)
return list(messages)
async def _rollback_to_pre_run_checkpoint(
*,
accessor: CheckpointStateAccessor | None,
checkpointer: Any,
thread_id: str,
run_id: str,
rollback_point: RollbackPoint | None,
snapshot_capture_failed: bool,
) -> bool:
"""Restore the complete pre-run state and report whether it completed.
Full mode forks the captured pre-run checkpoint and overwrites messages;
all other channels inherit from that parent. Delta mode cannot safely fork
once the cancelled path has attached writes to the same parent, so it
replaces every captured channel on the current head instead. Both writes
use a state-only mutation graph whose synthetic ``rollback_restore`` node
finishes immediately and schedules no agent work.
"""
if checkpointer is None:
logger.info("Run %s rollback requested but no checkpointer is configured", run_id)
return False
if snapshot_capture_failed:
logger.warning("Run %s rollback skipped: pre-run checkpoint capture failed", run_id)
return False
if rollback_point is None:
await _call_checkpointer_method(checkpointer, "adelete_thread", "delete_thread", thread_id)
logger.info("Run %s rollback reset thread %s to empty state", run_id, thread_id)
return True
configurable = rollback_point.config.get("configurable", {})
if not configurable.get("checkpoint_id"):
logger.warning("Run %s rollback skipped: pre-run checkpoint has no checkpoint id", run_id)
return False
if accessor is None:
# Unreachable in practice: a rollback point can only be captured
# through the bound accessor. Stay fail-closed.
logger.warning("Run %s rollback skipped: agent accessor unavailable", run_id)
return False
# Compile with the thread's effective schema so middleware-contributed
# channels survive (the base ThreadState fallback would silently drop
# them).
mutation_graph = build_state_mutation_graph("rollback_restore", accessor.mode, graph_state_schema(getattr(accessor, "graph", None)))
mutation_accessor = CheckpointStateAccessor.bind(mutation_graph, checkpointer, mode=accessor.mode)
if accessor.mode == "delta":
# A delta rollback fork has the same write-ownership problem as a
# checkpoint resume: the captured parent now carries writes from the
# cancelled sibling. Restore linearly on the current head instead.
restore_config: dict[str, Any] = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
current = await accessor.aget(restore_config)
raw_current_values = getattr(current, "values", None) or {}
current_values = dict(raw_current_values) if isinstance(raw_current_values, dict) else {}
selected_values = copy.deepcopy(rollback_point.state_values)
selected_values["messages"] = list(rollback_point.messages)
replacement_values = _complete_state_replacement_values(
mutation_graph=mutation_graph,
selected_values=selected_values,
current_values=current_values,
run_id=run_id,
operation="rollback",
)
else:
restore_config = rollback_point.config
replacement_values = {"messages": Overwrite(list(rollback_point.messages))}
restored_config = await mutation_accessor.aupdate(
restore_config,
replacement_values,
as_node="rollback_restore",
)
if not isinstance(restored_config, dict):
raise RuntimeError(f"Run {run_id} rollback restore returned invalid config: expected dict")
restored_configurable = restored_config.get("configurable", {})
if not isinstance(restored_configurable, dict):
raise RuntimeError(f"Run {run_id} rollback restore returned invalid config payload")
restored_checkpoint_id = restored_configurable.get("checkpoint_id")
if not restored_checkpoint_id:
raise RuntimeError(f"Run {run_id} rollback restore did not return checkpoint_id")
pending_writes = rollback_point.pending_writes
if not pending_writes:
return True
writes_by_task: dict[str, list[tuple[str, Any]]] = {}
for item in pending_writes:
if not isinstance(item, (tuple, list)) or len(item) != 3:
raise RuntimeError(f"Run {run_id} rollback failed: pending_write is not a 3-tuple: {item!r}")
task_id, channel, value = item
if not isinstance(channel, str):
raise RuntimeError(f"Run {run_id} rollback failed: pending_write has non-string channel: task_id={task_id!r}, channel={channel!r}")
writes_by_task.setdefault(str(task_id), []).append((channel, value))
for task_id, writes in writes_by_task.items():
await _call_checkpointer_method(
checkpointer,
"aput_writes",
"put_writes",
restored_config,
writes,
task_id=task_id,
)
return True
def _new_checkpoint_marker() -> dict[str, str]:
marker = empty_checkpoint()
return {"id": marker["id"], "ts": marker["ts"]}
def _bump_channel_version(checkpointer: Any, current_version: Any) -> Any:
"""Return a strictly-different next version for a checkpoint channel.
DB-backed LangGraph savers (PostgresSaver / v4 SqliteSaver blob layout)
persist channel blobs keyed by ``channel_versions[<channel>]``, so the
new value MUST differ from the prior value. We delegate to the
checkpointer's ``get_next_version`` when available — that is the canonical
versioning scheme each saver picks (int, monotonic float, or
UUID-shaped string). When the checkpointer doesn't expose it (or it
returns ``None``/an unchanged value), fall back to a defensive bump that
still guarantees inequality.
"""
get_next_version = getattr(checkpointer, "get_next_version", None)
if callable(get_next_version):
try:
next_version = get_next_version(current_version, None)
except Exception:
next_version = None
if next_version is not None and next_version != current_version:
return next_version
# fall through to defensive bump
if isinstance(current_version, bool):
# ``bool`` is a subclass of ``int``; treat True/False as 1/0 instead of
# adding to the boolean itself, which would produce an int anyway but
# via a path that surprises readers.
return int(current_version) + 1
if isinstance(current_version, int):
return current_version + 1
if isinstance(current_version, float):
# Match LangGraph's default float versioning (monotonic increment).
return current_version + 1.0
if isinstance(current_version, str):
try:
return str(int(current_version) + 1)
except ValueError:
return f"{current_version}.1"
return 1
def _checkpoint_identity(ckpt_tuple: Any | None, checkpoint: dict[str, Any]) -> str | None:
tuple_config = getattr(ckpt_tuple, "config", {}) or {}
tuple_configurable = tuple_config.get("configurable", {}) if isinstance(tuple_config, dict) else {}
if isinstance(tuple_configurable, dict):
checkpoint_id = tuple_configurable.get("checkpoint_id")
if isinstance(checkpoint_id, str) and checkpoint_id:
return checkpoint_id
checkpoint_id = checkpoint.get("id")
return checkpoint_id if isinstance(checkpoint_id, str) and checkpoint_id else None
def _checkpoint_namespace(ckpt_tuple: Any | None) -> str:
tuple_config = getattr(ckpt_tuple, "config", {}) or {}
tuple_configurable = tuple_config.get("configurable", {}) if isinstance(tuple_config, dict) else {}
checkpoint_ns = tuple_configurable.get("checkpoint_ns", "") if isinstance(tuple_configurable, dict) else ""
return checkpoint_ns if isinstance(checkpoint_ns, str) else ""
def _graph_input_messages(graph_input: Any | None) -> list[Any]:
if not isinstance(graph_input, dict):
return []
messages = graph_input.get("messages")
if isinstance(messages, list):
return messages
if isinstance(messages, tuple):
return list(messages)
return []
def _title_generation_state(channel_values: dict[str, Any], graph_input: Any | None) -> dict[str, Any]:
state = dict(channel_values)
messages = state.get("messages")
if not messages:
fallback_messages = _graph_input_messages(graph_input)
if fallback_messages:
state["messages"] = fallback_messages
return state
def valid_duration_entry(run_id: Any, duration_seconds: Any) -> bool:
"""Check that (run_id, duration_seconds) is a well-formed duration entry."""
return isinstance(run_id, str) and bool(run_id) and isinstance(duration_seconds, int) and not isinstance(duration_seconds, bool)
async def persist_run_durations(
*,
checkpointer: Any,
thread_id: str,
durations: dict[str, int],
) -> bool:
"""Merge validated run durations into a metadata-only checkpoint.
Durations accumulate so the history fast path can serve every known turn
from the latest checkpoint. Per-entry overhead is negligible (~50 bytes
per run_id) compared to the messages channel blob written on every graph
checkpoint, so no pruning is needed.
"""
updates = {run_id: max(0, duration_seconds) for run_id, duration_seconds in durations.items() if valid_duration_entry(run_id, duration_seconds)}
if not updates:
return False
ckpt_config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
async with _checkpoint_thread_lock(thread_id):
for _attempt in range(3):
ckpt_tuple = await _call_checkpointer_method(checkpointer, "aget_tuple", "get_tuple", ckpt_config)
if ckpt_tuple is None:
return False
checkpoint = dict(getattr(ckpt_tuple, "checkpoint", {}) or {})
metadata = dict(getattr(ckpt_tuple, "metadata", {}) or {})
raw_run_durations = metadata.get("run_durations")
run_durations = {key: value for key, value in raw_run_durations.items() if valid_duration_entry(key, value)} if isinstance(raw_run_durations, dict) else {}
changed_durations = {run_id: duration for run_id, duration in updates.items() if run_durations.get(run_id) != duration}
if not changed_durations:
return False
run_durations.update(changed_durations)
parent_checkpoint_id = _checkpoint_identity(ckpt_tuple, checkpoint)
latest_tuple = await _call_checkpointer_method(checkpointer, "aget_tuple", "get_tuple", ckpt_config)
latest_checkpoint = dict(getattr(latest_tuple, "checkpoint", {}) or {}) if latest_tuple is not None else {}
if _checkpoint_identity(latest_tuple, latest_checkpoint) != parent_checkpoint_id:
continue
checkpoint.update(_new_checkpoint_marker())
metadata["source"] = "update"
prev_step = metadata.get("step")
metadata["step"] = (prev_step + 1) if isinstance(prev_step, int) else 1
metadata["run_durations"] = run_durations
metadata["writes"] = {"runtime_run_duration": {"run_ids": sorted(changed_durations)}}
checkpoint_ns = _checkpoint_namespace(ckpt_tuple)
write_config = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": checkpoint_ns,
"checkpoint_id": parent_checkpoint_id,
}
}
await _call_checkpointer_method(
checkpointer,
"aput",
"put",
write_config,
checkpoint,
metadata,
{},
)
return True
return False
async def _persist_run_duration(
*,
checkpointer: Any,
thread_id: str,
run_id: str,
duration_seconds: int,
) -> None:
"""Persist one completed run duration in the thread checkpoint metadata."""
await persist_run_durations(
checkpointer=checkpointer,
thread_id=thread_id,
durations={run_id: duration_seconds},
)
async def _ensure_interrupted_title(*, checkpointer: Any, thread_id: str, app_config: AppConfig | None, graph_input: Any | None = None) -> str | None:
"""Persist a local fallback title for interrupted first-turn runs.
Returns the title that is now persisted (existing or newly written), or
``None`` when no checkpoint is available or no title text can be derived.
Idempotent: re-invoking against a checkpoint that already carries a title
short-circuits without writing a new checkpoint.
"""
from deerflow.agents.middlewares.title_middleware import TitleMiddleware
middleware = TitleMiddleware(app_config=app_config) if app_config is not None else TitleMiddleware()
ckpt_config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
for _attempt in range(3):
ckpt_tuple = await _call_checkpointer_method(checkpointer, "aget_tuple", "get_tuple", ckpt_config)
checkpoint = copy.deepcopy(getattr(ckpt_tuple, "checkpoint", {}) or {}) if ckpt_tuple is not None else empty_checkpoint()
channel_values = dict(checkpoint.get("channel_values", {}) or {})
existing_title = channel_values.get("title")
if existing_title:
return existing_title
result = middleware._generate_title_result(_title_generation_state(channel_values, graph_input), allow_partial_exchange=True)
title = result.get("title") if isinstance(result, dict) else None
if not title:
return None
# ``empty_checkpoint()`` creates a fresh id every time; only real tuples
# carry an identity stable enough for the stale-snapshot comparison.
base_identity = _checkpoint_identity(ckpt_tuple, checkpoint) if ckpt_tuple is not None else None
latest_tuple = await _call_checkpointer_method(checkpointer, "aget_tuple", "get_tuple", ckpt_config)
latest_checkpoint = copy.deepcopy(getattr(latest_tuple, "checkpoint", {}) or {}) if latest_tuple is not None else empty_checkpoint()
latest_identity = _checkpoint_identity(latest_tuple, latest_checkpoint) if latest_tuple is not None else None
if base_identity is None:
if latest_identity is not None:
continue
elif latest_identity != base_identity:
continue
checkpoint = latest_checkpoint
channel_values = dict(checkpoint.get("channel_values", {}) or {})
existing_title = channel_values.get("title")
if existing_title:
return existing_title
channel_values["title"] = title
marker = _new_checkpoint_marker()
checkpoint.update({"id": marker["id"], "ts": marker["ts"], "channel_values": channel_values})
# Bump ``channel_versions["title"]`` and declare the bump in ``new_versions``
# so DB-backed savers (SqliteSaver v4 / PostgresSaver) actually persist the
# new blob — those savers strip inline ``channel_values`` from ``put`` and
# only write blobs for channels listed in ``new_versions``. The legacy
# single-table sqlite saver ignores ``new_versions`` and inlines the
# snapshot, so this path is correct for both layouts. Mirrors
# ``_rollback_to_pre_run_checkpoint`` in the same file.
channel_versions = dict(checkpoint.get("channel_versions", {}) or {})
next_title_version = _bump_channel_version(checkpointer, channel_versions.get("title"))
channel_versions["title"] = next_title_version
checkpoint["channel_versions"] = channel_versions
metadata = dict(getattr(latest_tuple, "metadata", {}) or {})
metadata["source"] = "update"
prev_step = metadata.get("step")
metadata["step"] = (prev_step + 1) if isinstance(prev_step, int) else 1
metadata["writes"] = {"runtime_interrupt_title": {"title": title}}
checkpoint_ns = _checkpoint_namespace(latest_tuple)
# Parent to the checkpoint this write was derived from - a parentless
# raw write would sever Delta-channel replay ancestry (and truncate
# full-mode history walks).
write_config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": checkpoint_ns, "checkpoint_id": latest_identity}}
await _call_checkpointer_method(
checkpointer,
"aput",
"put",
write_config,
checkpoint,
metadata,
{"title": next_title_version},
)
return title
return None
def _lg_mode_to_sse_event(mode: str) -> str:
"""Map LangGraph internal stream_mode name to SSE event name.
LangGraph's ``astream(stream_mode="messages")`` produces message
tuples. The SSE protocol calls this ``messages-tuple`` when the
client explicitly requests it, but the default SSE event name used
by LangGraph Platform is simply ``"messages"``.
"""
# All LG modes map 1:1 to SSE event names — "messages" stays "messages"
return mode
def _error_fallback_message_from_metadata(metadata: dict[str, Any], content: Any) -> str:
detail = metadata.get("error_detail")
if isinstance(detail, str) and detail.strip():
return detail.strip()
reason = metadata.get("error_reason")
if isinstance(reason, str) and reason.strip():
return reason.strip()
if isinstance(content, str) and content.strip():
return content.strip()[:2000]
return "LLM provider failed after retries"
def _message_id(obj: Any) -> str | None:
"""Best-effort extraction of a stable message id from a message-like object."""
msg_id = getattr(obj, "id", None)
if isinstance(msg_id, str) and msg_id:
return msg_id
if isinstance(obj, dict):
raw = obj.get("id")
if isinstance(raw, str) and raw:
return raw
return None
def _try_extract_from_message(obj: Any, pre_existing_ids: set[str] | None = None) -> str | None:
"""Try to extract fallback marker from a single message object or dict.
Messages whose id appears in ``pre_existing_ids`` are skipped — those are
history checkpointed by a *prior* run on this thread and any fallback
marker on them was already accounted for when that earlier run finished.
Without this filter, a single past run that ended with a fallback marker
would mark every subsequent run on the same thread as ``error``, because
LangGraph replays the full message history through ``stream_mode="values"``.
"""
if pre_existing_ids:
msg_id = _message_id(obj)
if msg_id is not None and msg_id in pre_existing_ids:
return None
additional_kwargs = getattr(obj, "additional_kwargs", None)
if isinstance(additional_kwargs, dict) and additional_kwargs.get("deerflow_error_fallback"):
return _error_fallback_message_from_metadata(additional_kwargs, getattr(obj, "content", None))
if isinstance(obj, dict):
nested_kwargs = obj.get("additional_kwargs")
if isinstance(nested_kwargs, dict) and nested_kwargs.get("deerflow_error_fallback"):
return _error_fallback_message_from_metadata(nested_kwargs, obj.get("content"))
return None
def _extract_llm_error_fallback_message(value: Any, pre_existing_ids: set[str] | None = None) -> str | None:
"""Find LLM fallback markers in streamed LangGraph chunks.
Error fallback messages returned by model-call middleware are not guaranteed
to pass through LLM end callbacks, but they do appear in graph state chunks.
Messages whose id appears in ``pre_existing_ids`` are ignored — they are
history from prior runs on the same thread (LangGraph replays the full
messages channel in ``stream_mode="values"`` chunks), and any error
fallback in that history was already resolved when its run finished.
"""
# Fast path: large state chunks produced by stream_mode="values" have a
# top-level "messages" list. Scanning only that list avoids expensive deep
# recursion into large state dicts.
if isinstance(value, dict):
messages = value.get("messages")
if isinstance(messages, (list, tuple)):
for msg in messages:
result = _try_extract_from_message(msg, pre_existing_ids)
if result is not None:
return result
# Fallback marker is attached to an AI message in the messages
# channel; it will never appear elsewhere in a values chunk.
return None
# No top-level "messages" — this is likely an "updates" chunk (small
# dict keyed by node name). Fall through to deep walk, which is cheap
# for these payloads.
# Deep walk for updates / messages / tuple / list modes. Payloads are
# small, so full recursion is acceptable here.
seen: set[int] = set()
def walk(obj: Any) -> str | None:
oid = id(obj)
if oid in seen:
return None
seen.add(oid)
result = _try_extract_from_message(obj, pre_existing_ids)
if result is not None:
return result
if isinstance(obj, dict):
for item in obj.values():
result = walk(item)
if result is not None:
return result
return None
if isinstance(obj, (list, tuple, set)):
for item in obj:
result = walk(item)
if result is not None:
return result
return None
return walk(value)
def _collect_pre_existing_message_ids(values: Any) -> set[str]:
"""Collect stable message IDs from graph-materialized channel values."""
if not isinstance(values, dict):
return set()
messages = values.get("messages")
if not isinstance(messages, (list, tuple)):
return set()
return {message_id for message in messages if (message_id := _message_id(message)) is not None}
def _unpack_stream_item(
item: Any,
lg_modes: list[str],
stream_subgraphs: bool,
) -> tuple[str | None, Any, tuple[str, ...]]:
"""Unpack a multi-mode or subgraph stream item into (mode, chunk, namespace).
``namespace`` is the subgraph namespace tuple LangGraph prefixes onto each
frame when ``subgraphs=True``; it is empty for root-graph frames. Delegated
subagent graphs inherit the parent's checkpoint namespace (see
``subagents/executor.py``), so their frames arrive here with a non-empty
namespace and must not be mistaken for root frames.
Returns ``(None, None, ())`` if the item cannot be parsed.
"""
if stream_subgraphs:
if isinstance(item, tuple) and len(item) == 3:
ns, mode, chunk = item
namespace = tuple(str(part) for part in ns) if isinstance(ns, (list, tuple)) else (str(ns),)
return str(mode), chunk, namespace
if isinstance(item, tuple) and len(item) == 2:
mode, chunk = item
return str(mode), chunk, ()
return None, None, ()
if isinstance(item, tuple) and len(item) == 2:
mode, chunk = item
return str(mode), chunk, ()
# Fallback: single-element output from first mode
return lg_modes[0] if lg_modes else None, item, ()
def _compose_sse_event(sse_event: str, namespace: tuple[str, ...]) -> str:
"""Namespace-qualified SSE event name, LangGraph Platform style.
Root frames keep the bare event name; subgraph frames become
``mode|ns1|ns2`` so clients can tell them apart. The LangGraph SDK parses
exactly this shape (``event.split("|").slice(1)``) and routes
subagent-namespaced values away from the thread view.
"""
if not namespace:
return sse_event
return "|".join((sse_event, *namespace))
async def _publish_stream_item(
*,
bridge: Any,
run_id: str,
mode: str,
chunk: Any,
namespace: tuple[str, ...],
file_tool_chunk_batcher: Any,
subagent_events: Any,
) -> None:
"""Publish one stream frame, preserving the subgraph namespace.
A subgraph frame published under a bare event name impersonates the root
graph: a delegated subagent's ``values`` snapshot then replaces the whole
thread view in SDK clients and its token chunks flood the parent message
stream (#4399). Subgraph frames therefore keep their namespace in the event
name and bypass the root-only consumers (file-tool chunk batcher, subagent
event persistence — task_* lifecycle events are root frames already).
"""
sse_event = _compose_sse_event(_lg_mode_to_sse_event(mode), namespace)
if namespace:
await bridge.publish(run_id, sse_event, serialize(chunk, mode=mode))
return
if file_tool_chunk_batcher is not None and mode != "messages":
pending_chunks = file_tool_chunk_batcher.finish() if mode == "values" else file_tool_chunk_batcher.flush()
for publish_chunk in pending_chunks:
await bridge.publish(run_id, "messages", serialize(publish_chunk, mode="messages"))
chunks_to_publish = file_tool_chunk_batcher.push(chunk) if mode == "messages" and file_tool_chunk_batcher is not None else [chunk]
for publish_chunk in chunks_to_publish:
await bridge.publish(run_id, sse_event, serialize(publish_chunk, mode=mode))
if mode == "custom":
await subagent_events.add(chunk)