fix(checkpoint): persist run duration in checkpoints for history reads (#4118)

* fix: persist run duration in checkpoints for history reads

* fix(checkpoint): harden run duration persistence

* fix(checkpoint): persist run durations in metadata

* fix(checkpoint): address review findings for run duration persistence

- Add valid_duration_entry() shared validation helper (worker.py)
- Rename _persist_run_durations -> persist_run_durations as public API
- Import public persist_run_durations and valid_duration_entry in threads.py
- Use BackgroundTasks for lazy backfill write to avoid blocking history reads
- Add TODO about O(runs) growth of run_durations in checkpoint metadata
- Document REGENERATE_HISTORY_RAW_SCAN_LIMIT doubling assumption

* fix(checkpoint): replace pruning TODO with justification

Accumulated run_durations overhead (~50 bytes/run_id) is negligible compared
to messages channel blobs; no pruning strategy is needed.

---------

Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
This commit is contained in:
Vanzeren 2026-07-14 21:37:47 +08:00 committed by GitHub
parent fabadae416
commit 13fd8e229a
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
6 changed files with 778 additions and 83 deletions

View File

@ -33,9 +33,18 @@ from deerflow.workspace_changes import get_workspace_changes_response
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/threads", tags=["runs"]) router = APIRouter(prefix="/api/threads", tags=["runs"])
REGENERATE_HISTORY_SCAN_LIMIT = 200 REGENERATE_HISTORY_SCAN_LIMIT = 200
# Doubled to keep ~200 effective checkpoints when duration-only checkpoints
# (one per successful run in steady state) consume roughly half of history.
REGENERATE_HISTORY_RAW_SCAN_LIMIT = REGENERATE_HISTORY_SCAN_LIMIT * 2
THREAD_MESSAGE_PAGE_SCAN_BATCH = 201 THREAD_MESSAGE_PAGE_SCAN_BATCH = 201
def _is_duration_only_checkpoint(checkpoint_tuple: Any) -> bool:
metadata = getattr(checkpoint_tuple, "metadata", None)
writes = metadata.get("writes") if isinstance(metadata, dict) else None
return isinstance(writes, dict) and "runtime_run_duration" in writes
def compute_run_durations(runs) -> dict[str, int]: def compute_run_durations(runs) -> dict[str, int]:
"""Map run_id -> duration in seconds from run timestamps.""" """Map run_id -> duration in seconds from run timestamps."""
from datetime import datetime from datetime import datetime
@ -379,7 +388,8 @@ async def _find_base_checkpoint_before_human(thread_id: str, human_message_id: s
checkpointer = get_checkpointer(request) checkpointer = get_checkpointer(request)
base_config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}} base_config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
try: try:
checkpoints = [item async for item in checkpointer.alist(base_config, limit=REGENERATE_HISTORY_SCAN_LIMIT)] raw_checkpoints = [item async for item in checkpointer.alist(base_config, limit=REGENERATE_HISTORY_RAW_SCAN_LIMIT)]
checkpoints = [item for item in raw_checkpoints if not _is_duration_only_checkpoint(item)]
except Exception as exc: except Exception as exc:
logger.exception("Failed to list checkpoints for regenerate thread %s", thread_id) logger.exception("Failed to list checkpoints for regenerate thread %s", thread_id)
raise HTTPException(status_code=500, detail="Failed to inspect checkpoint history") from exc raise HTTPException(status_code=500, detail="Failed to inspect checkpoint history") from exc

View File

@ -19,7 +19,7 @@ import uuid
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
from fastapi import APIRouter, HTTPException, Request from fastapi import APIRouter, BackgroundTasks, HTTPException, Request
from langgraph.checkpoint.base import empty_checkpoint, uuid6 from langgraph.checkpoint.base import empty_checkpoint, uuid6
from pydantic import BaseModel, Field, field_validator from pydantic import BaseModel, Field, field_validator
@ -44,6 +44,7 @@ from deerflow.runtime.goal import (
read_thread_goal, read_thread_goal,
write_thread_goal, write_thread_goal,
) )
from deerflow.runtime.runs.worker import valid_duration_entry
from deerflow.runtime.user_context import get_effective_user_id from deerflow.runtime.user_context import get_effective_user_id
from deerflow.utils.file_io import run_file_io from deerflow.utils.file_io import run_file_io
from deerflow.utils.time import coerce_iso, now_iso from deerflow.utils.time import coerce_iso, now_iso
@ -1043,9 +1044,36 @@ async def update_thread_state(thread_id: str, body: ThreadStateUpdateRequest, re
) )
def _ai_message_lacks_duration(message: dict[str, Any]) -> bool:
additional_kwargs = message.get("additional_kwargs")
return message.get("type") == "ai" and (not isinstance(additional_kwargs, dict) or "turn_duration" not in additional_kwargs)
def _checkpoint_run_durations(metadata: Any) -> dict[str, int]:
raw_durations = metadata.get("run_durations") if isinstance(metadata, dict) else None
if not isinstance(raw_durations, dict):
return {}
return {run_id: duration_seconds for run_id, duration_seconds in raw_durations.items() if valid_duration_entry(run_id, duration_seconds)}
def _set_message_turn_duration(message: dict[str, Any], run_id: str, run_durations: dict[str, int]) -> None:
if message.get("type") != "ai" or run_id not in run_durations:
return
additional_kwargs = message.get("additional_kwargs")
if not isinstance(additional_kwargs, dict):
additional_kwargs = {}
message["additional_kwargs"] = additional_kwargs
additional_kwargs.setdefault("turn_duration", run_durations[run_id])
@router.post("/{thread_id}/history", response_model=list[HistoryEntry]) @router.post("/{thread_id}/history", response_model=list[HistoryEntry])
@require_permission("threads", "read", owner_check=True) @require_permission("threads", "read", owner_check=True)
async def get_thread_history(thread_id: str, body: ThreadHistoryRequest, request: Request) -> list[HistoryEntry]: async def get_thread_history(
thread_id: str,
body: ThreadHistoryRequest,
request: Request,
background_tasks: BackgroundTasks,
) -> list[HistoryEntry]:
"""Get checkpoint history for a thread. """Get checkpoint history for a thread.
Messages are read from the checkpointer's channel values (the Messages are read from the checkpointer's channel values (the
@ -1089,56 +1117,76 @@ async def get_thread_history(thread_id: str, body: ThreadHistoryRequest, request
if messages: if messages:
serialized_msgs = serialize_channel_values_for_api({"messages": messages}).get("messages", []) serialized_msgs = serialize_channel_values_for_api({"messages": messages}).get("messages", [])
try: try:
from app.gateway.deps import get_run_event_store, get_run_manager # Human messages define turn boundaries. New checkpoints
from app.gateway.routers.thread_runs import compute_run_durations # carry the completed turns' durations in metadata, so the
# messages channel stays unchanged.
checkpoint_run_durations = _checkpoint_run_durations(metadata)
current_turn_run_id = None
for msg in serialized_msgs:
if msg.get("type") == "human":
additional_kwargs = msg.get("additional_kwargs")
if isinstance(additional_kwargs, dict):
run_id = additional_kwargs.get("run_id")
if isinstance(run_id, str) and run_id:
current_turn_run_id = run_id
continue
run_mgr = get_run_manager(request) if msg.get("type") not in {"ai", "tool"} or not current_turn_run_id:
event_store = get_run_event_store(request) continue
runs = await run_mgr.list_by_thread(thread_id) msg.setdefault("run_id", current_turn_run_id)
_set_message_turn_duration(msg, current_turn_run_id, checkpoint_run_durations)
# FIXME: Fetching limit=1000 silently drops durations for messages older than the cap on long threads. # Legacy checkpoints without duration metadata are
# We do this full fetch because raw LangGraph messages lack a native run_id link. # correlated once via event-store + run-manager, then
# upgraded by a metadata-only checkpoint write.
if any(_ai_message_lacks_duration(msg) for msg in serialized_msgs):
from app.gateway.deps import get_run_event_store, get_run_manager
from app.gateway.routers.thread_runs import compute_run_durations
from deerflow.runtime.runs.worker import persist_run_durations
events = await event_store.list_messages(thread_id, limit=1000) run_mgr = get_run_manager(request)
event_store = get_run_event_store(request)
if runs and serialized_msgs: runs = await run_mgr.list_by_thread(thread_id)
# 1. Map each run_id to its actual duration events = await event_store.list_messages(thread_id, limit=1000)
run_durations = compute_run_durations(runs)
# 2. Map every message id directly to its parent run_id if runs:
msg_to_run = {} run_durations = compute_run_durations(runs)
for e in events: msg_to_run = {}
content = e.get("content", {}) for event in events:
if isinstance(content, dict) and content.get("type") == "ai" and "id" in content: content = event.get("content", {})
msg_to_run[content["id"]] = e["run_id"] run_id = event.get("run_id")
if isinstance(content, dict) and content.get("type") == "ai" and "id" in content and isinstance(run_id, str) and run_id:
msg_to_run[content["id"]] = run_id
# 3. Attach the owning run_id to replayed messages. current_turn_run_id = None
# Raw LangGraph checkpoint messages do not carry a for msg in serialized_msgs:
# native run link. Message events are exact when if msg.get("type") == "human":
# present, but historical/runtime stores can miss additional_kwargs = msg.get("additional_kwargs")
# them; the user-input message already records the if isinstance(additional_kwargs, dict):
# run id for the whole turn, so use it as the run_id = additional_kwargs.get("run_id")
# fallback for following AI/tool messages. if isinstance(run_id, str) and run_id:
current_turn_run_id = None current_turn_run_id = run_id
for msg in serialized_msgs: continue
if msg.get("type") == "human":
additional_kwargs = msg.get("additional_kwargs")
if isinstance(additional_kwargs, dict):
run_id = additional_kwargs.get("run_id")
if isinstance(run_id, str) and run_id:
current_turn_run_id = run_id
continue
if msg.get("type") in {"ai", "tool"}: if msg.get("type") not in {"ai", "tool"}:
msg_id = msg.get("id") continue
run_id = msg_to_run.get(msg_id) or current_turn_run_id run_id = msg_to_run.get(msg.get("id")) or current_turn_run_id
if run_id: if run_id:
msg["run_id"] = run_id msg["run_id"] = run_id
if msg.get("type") == "ai" and run_id in run_durations: _set_message_turn_duration(msg, run_id, run_durations)
if "additional_kwargs" not in msg:
msg["additional_kwargs"] = {} # Intentional, best-effort write-on-read migration:
msg["additional_kwargs"]["turn_duration"] = run_durations[run_id] # persist legacy metadata after the response so the
# history request never waits on an active stream's
# same-thread checkpoint lock.
background_tasks.add_task(
persist_run_durations,
checkpointer=checkpointer,
thread_id=thread_id,
durations=run_durations,
)
except Exception: except Exception:
logger.warning("Failed to inject turn_duration for thread %s", thread_id, exc_info=True) logger.warning("Failed to inject turn_duration for thread %s", thread_id, exc_info=True)
@ -1152,7 +1200,7 @@ async def get_thread_history(thread_id: str, body: ThreadHistoryRequest, request
next_tasks = [t.name for t in tasks_raw if hasattr(t, "name")] next_tasks = [t.name for t in tasks_raw if hasattr(t, "name")]
# Strip LangGraph internal keys from metadata # Strip LangGraph internal keys from metadata
user_meta = {k: v for k, v in metadata.items() if k not in ("created_at", "updated_at", "step", "source", "writes", "parents")} user_meta = {k: v for k, v in metadata.items() if k not in ("created_at", "updated_at", "step", "source", "writes", "parents", "run_durations")}
# Keep step for ordering context # Keep step for ordering context
if "step" in metadata: if "step" in metadata:
user_meta["step"] = metadata["step"] user_meta["step"] = metadata["step"]

View File

@ -20,7 +20,12 @@ import copy
import inspect import inspect
import logging import logging
import os import os
import threading
import weakref
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from dataclasses import dataclass, field from dataclasses import dataclass, field
from datetime import datetime
from functools import lru_cache from functools import lru_cache
from typing import Any, Literal, cast from typing import Any, Literal, cast
@ -63,6 +68,28 @@ from .schemas import RunStatus
logger = logging.getLogger(__name__) 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
# Valid stream_mode values for LangGraph's graph.astream() # Valid stream_mode values for LangGraph's graph.astream()
_VALID_LG_MODES = {"values", "updates", "checkpoints", "tasks", "debug", "messages", "custom"} _VALID_LG_MODES = {"values", "updates", "checkpoints", "tasks", "debug", "messages", "custom"}
@ -454,39 +481,40 @@ async def run_agent(
async def _stream_once(input_payload: Any, stream_config: RunnableConfig) -> None: async def _stream_once(input_payload: Any, stream_config: RunnableConfig) -> None:
nonlocal llm_error_fallback_message nonlocal llm_error_fallback_message
if len(lg_modes) == 1 and not stream_subgraphs: async with _checkpoint_thread_lock(thread_id):
# Single mode, no subgraphs: astream yields raw chunks if len(lg_modes) == 1 and not stream_subgraphs:
single_mode = lg_modes[0] # Single mode, no subgraphs: astream yields raw chunks
async for chunk in agent.astream(input_payload, config=stream_config, stream_mode=single_mode): 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(): if record.abort_event.is_set():
logger.info("Run %s abort requested — stopping", run_id) logger.info("Run %s abort requested — stopping", run_id)
break break
mode, chunk = _unpack_stream_item(item, lg_modes, stream_subgraphs)
if mode is None:
continue
llm_error_fallback_message = llm_error_fallback_message or _extract_llm_error_fallback_message(chunk, pre_existing_message_ids) 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) sse_event = _lg_mode_to_sse_event(mode)
await bridge.publish(run_id, sse_event, serialize(chunk, mode=single_mode)) await bridge.publish(run_id, sse_event, serialize(chunk, mode=mode))
if single_mode == "custom": if mode == "custom":
await subagent_events.add(chunk) 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 = _unpack_stream_item(item, lg_modes, stream_subgraphs)
if mode is None:
continue
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(mode)
await bridge.publish(run_id, sse_event, serialize(chunk, mode=mode))
if mode == "custom":
await subagent_events.add(chunk)
# 7. Stream the requested turn, then optionally continue hidden goal turns. # 7. Stream the requested turn, then optionally continue hidden goal turns.
await _stream_once(graph_input, initial_runnable_config) await _stream_once(graph_input, initial_runnable_config)
@ -621,6 +649,25 @@ async def run_agent(
except Exception: except Exception:
logger.debug("Failed to sync title for thread %s (non-fatal)", thread_id) 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 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 # Update threads_meta status based on run outcome
if thread_store is not None: if thread_store is not None:
try: try:
@ -1142,6 +1189,93 @@ def _title_generation_state(channel_values: dict[str, Any], graph_input: Any | N
return state 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: 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. """Persist a local fallback title for interrupted first-turn runs.

View File

@ -0,0 +1,388 @@
import asyncio
import copy
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.base import empty_checkpoint, uuid6
from langgraph.checkpoint.memory import InMemorySaver
import deerflow.runtime.runs.worker as worker
from deerflow.runtime.goal import goal_thread_lock
from deerflow.runtime.runs.manager import RunManager
from deerflow.runtime.runs.worker import RunContext, _persist_run_duration, run_agent
class _YieldingSaver(InMemorySaver):
async def aget_tuple(self, config):
checkpoint_tuple = await super().aget_tuple(config)
await asyncio.sleep(0)
return checkpoint_tuple
async def aput(self, config, checkpoint, metadata, new_versions):
await asyncio.sleep(0)
return await super().aput(config, checkpoint, metadata, new_versions)
class _AdvancingSaver(InMemorySaver):
"""Inject a title checkpoint between duration read and write."""
def __init__(self) -> None:
super().__init__()
self._reads = 0
async def aget_tuple(self, config):
self._reads += 1
checkpoint_tuple = await super().aget_tuple(config)
if self._reads != 2 or checkpoint_tuple is None:
return checkpoint_tuple
checkpoint = copy.deepcopy(checkpoint_tuple.checkpoint)
checkpoint["id"] = str(uuid6())
channel_values = dict(checkpoint["channel_values"])
channel_values["title"] = "Concurrent title"
checkpoint["channel_values"] = channel_values
channel_versions = dict(checkpoint["channel_versions"])
channel_versions["title"] = 1
checkpoint["channel_versions"] = channel_versions
metadata = dict(checkpoint_tuple.metadata)
metadata.update({"step": metadata["step"] + 1, "source": "update", "writes": {"title": "Concurrent title"}})
await super().aput(checkpoint_tuple.config, checkpoint, metadata, {"title": 1})
return await super().aget_tuple(config)
async def _put_checkpoint(
checkpointer: InMemorySaver,
*,
thread_id: str,
checkpoint_id: str,
messages: list[object],
step: int,
parent_config: dict | None = None,
inherited_metadata: dict | None = None,
) -> dict:
checkpoint = empty_checkpoint()
checkpoint["id"] = checkpoint_id
checkpoint["channel_values"] = {"messages": messages}
checkpoint["channel_versions"] = {"messages": step}
config = parent_config or {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
metadata = dict(inherited_metadata or {})
metadata.update({"step": step, "source": "loop", "writes": {"test": {"messages": messages}}, "parents": {}})
return await checkpointer.aput(config, checkpoint, metadata, {"messages": step})
@pytest.mark.anyio
async def test_run_duration_survives_a_later_checkpoint() -> None:
checkpointer = InMemorySaver()
thread_id = "duration-survives"
messages = [
HumanMessage(id="human-1", content="Question", additional_kwargs={"run_id": "run-1"}),
AIMessage(id="ai-1", content="Answer"),
]
await _put_checkpoint(
checkpointer,
thread_id=thread_id,
checkpoint_id="00000000-0000-6000-8000-000000000001",
messages=messages,
step=1,
)
await _persist_run_duration(
checkpointer=checkpointer,
thread_id=thread_id,
run_id="run-1",
duration_seconds=7,
)
config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
duration_checkpoint = await checkpointer.aget_tuple(config)
assert duration_checkpoint is not None
persisted_messages = copy.deepcopy(duration_checkpoint.checkpoint["channel_values"]["messages"])
assert duration_checkpoint.metadata["run_durations"] == {"run-1": 7}
await _put_checkpoint(
checkpointer,
thread_id=thread_id,
checkpoint_id=str(uuid6()),
messages=persisted_messages,
step=3,
parent_config=duration_checkpoint.config,
inherited_metadata=duration_checkpoint.metadata,
)
latest = await checkpointer.aget_tuple(config)
assert latest is not None
assert latest.metadata["run_durations"] == {"run-1": 7}
@pytest.mark.anyio
async def test_run_duration_checkpoint_stores_duration_in_metadata_without_rewriting_messages() -> None:
checkpointer = InMemorySaver()
thread_id = "duration-metadata"
messages = [
HumanMessage(id="human-1", content="Question", additional_kwargs={"run_id": "run-1"}),
AIMessage(id="ai-1", content="Answer"),
]
await _put_checkpoint(
checkpointer,
thread_id=thread_id,
checkpoint_id="00000000-0000-6000-8000-000000000001",
messages=messages,
step=1,
)
await _persist_run_duration(
checkpointer=checkpointer,
thread_id=thread_id,
run_id="run-1",
duration_seconds=7,
)
latest = await checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}})
assert latest is not None
assert latest.metadata["run_durations"] == {"run-1": 7}
assert latest.checkpoint["channel_versions"]["messages"] == 1
assert "turn_duration" not in latest.checkpoint["channel_values"]["messages"][1].additional_kwargs
@pytest.mark.anyio
async def test_run_duration_retries_after_intervening_title_checkpoint() -> None:
checkpointer = _AdvancingSaver()
thread_id = "duration-title-race"
await _put_checkpoint(
checkpointer,
thread_id=thread_id,
checkpoint_id="00000000-0000-6000-8000-000000000001",
messages=[
HumanMessage(id="human-1", content="Question", additional_kwargs={"run_id": "run-1"}),
AIMessage(id="ai-1", content="Answer"),
],
step=1,
)
await _persist_run_duration(
checkpointer=checkpointer,
thread_id=thread_id,
run_id="run-1",
duration_seconds=7,
)
latest = await checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}})
assert latest is not None
assert latest.checkpoint["channel_values"]["title"] == "Concurrent title"
assert latest.metadata["run_durations"] == {"run-1": 7}
@pytest.mark.anyio
async def test_concurrent_run_duration_updates_preserve_both_turns() -> None:
checkpointer = _YieldingSaver()
thread_id = "duration-concurrent"
messages = [
HumanMessage(id="human-1", content="First", additional_kwargs={"run_id": "run-1"}),
AIMessage(id="ai-1", content="First answer"),
HumanMessage(id="human-2", content="Second", additional_kwargs={"run_id": "run-2"}),
AIMessage(id="ai-2", content="Second answer"),
]
await _put_checkpoint(
checkpointer,
thread_id=thread_id,
checkpoint_id="00000000-0000-6000-8000-000000000001",
messages=messages,
step=1,
)
await asyncio.gather(
_persist_run_duration(
checkpointer=checkpointer,
thread_id=thread_id,
run_id="run-1",
duration_seconds=3,
),
_persist_run_duration(
checkpointer=checkpointer,
thread_id=thread_id,
run_id="run-2",
duration_seconds=5,
),
)
latest = await checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}})
assert latest is not None
assert latest.metadata["run_durations"] == {"run-1": 3, "run-2": 5}
@pytest.mark.anyio
async def test_run_duration_checkpoint_preserves_parent_lineage() -> None:
checkpointer = InMemorySaver()
thread_id = "duration-parent"
parent_checkpoint_id = "00000000-0000-6000-8000-000000000001"
await _put_checkpoint(
checkpointer,
thread_id=thread_id,
checkpoint_id=parent_checkpoint_id,
messages=[
HumanMessage(id="human-1", content="Question", additional_kwargs={"run_id": "run-1"}),
AIMessage(id="ai-1", content="Answer"),
],
step=1,
)
await _persist_run_duration(
checkpointer=checkpointer,
thread_id=thread_id,
run_id="run-1",
duration_seconds=7,
)
history = [checkpoint async for checkpoint in checkpointer.alist({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}})]
assert len(history) == 2
assert history[0].config["configurable"]["checkpoint_id"] != parent_checkpoint_id
assert history[0].parent_config == {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": "",
"checkpoint_id": parent_checkpoint_id,
}
}
@pytest.mark.anyio
async def test_agent_stream_serializes_with_duration_checkpoint_write() -> None:
checkpointer = _YieldingSaver()
run_manager = RunManager()
record = await run_manager.create("duration-stream-lock")
await _put_checkpoint(
checkpointer,
thread_id=record.thread_id,
checkpoint_id="00000000-0000-6000-8000-000000000001",
messages=[
HumanMessage(
id="human-1",
content="Question",
additional_kwargs={"run_id": record.run_id},
),
AIMessage(id="ai-1", content="Answer"),
],
step=1,
)
bridge = SimpleNamespace(
publish=AsyncMock(),
publish_end=AsyncMock(),
cleanup=AsyncMock(),
)
duration_task = None
finished_during_stream = None
class DummyAgent:
async def astream(self, graph_input, config=None, stream_mode=None, subgraphs=False):
nonlocal duration_task, finished_during_stream
duration_task = asyncio.create_task(
_persist_run_duration(
checkpointer=checkpointer,
thread_id=record.thread_id,
run_id=record.run_id,
duration_seconds=9,
)
)
try:
await asyncio.wait_for(asyncio.shield(duration_task), timeout=0.05)
except TimeoutError:
finished_during_stream = False
else:
finished_during_stream = True
yield {"messages": []}
def factory(*, config):
return DummyAgent()
await run_agent(
bridge,
run_manager,
record,
ctx=RunContext(checkpointer=checkpointer),
agent_factory=factory,
graph_input={},
config={},
)
assert duration_task is not None
await duration_task
assert finished_during_stream is False
@pytest.mark.anyio
async def test_agent_stream_allows_graph_goal_state_access() -> None:
"""A graph node may acquire the goal lock while a run is streaming."""
checkpointer = InMemorySaver()
run_manager = RunManager()
record = await run_manager.create("duration-stream-goal-lock")
bridge = SimpleNamespace(
publish=AsyncMock(),
publish_end=AsyncMock(),
cleanup=AsyncMock(),
)
class DummyAgent:
async def astream(self, graph_input, config=None, stream_mode=None, subgraphs=False):
async with goal_thread_lock(record.thread_id):
yield {"messages": []}
def factory(*, config):
return DummyAgent()
await asyncio.wait_for(
run_agent(
bridge,
run_manager,
record,
ctx=RunContext(checkpointer=checkpointer),
agent_factory=factory,
graph_input={},
config={},
),
timeout=0.05,
)
assert record.status.value == "success"
@pytest.mark.anyio
async def test_successful_subsecond_run_persists_zero_duration(monkeypatch: pytest.MonkeyPatch) -> None:
checkpointer = InMemorySaver()
run_manager = RunManager()
record = await run_manager.create("duration-zero")
record.created_at = "2026-01-01T00:00:00+00:00"
record.updated_at = record.created_at
bridge = SimpleNamespace(
publish=AsyncMock(),
publish_end=AsyncMock(),
cleanup=AsyncMock(),
)
persist_duration = AsyncMock()
async def set_status(run_id, status, **kwargs):
record.status = status
class DummyAgent:
async def astream(self, graph_input, config=None, stream_mode=None, subgraphs=False):
yield {"messages": []}
monkeypatch.setattr(run_manager, "set_status", set_status)
monkeypatch.setattr(worker, "_persist_run_duration", persist_duration)
await run_agent(
bridge,
run_manager,
record,
ctx=RunContext(checkpointer=checkpointer),
agent_factory=lambda *, config: DummyAgent(),
graph_input={},
config={},
)
persist_duration.assert_awaited_once_with(
checkpointer=checkpointer,
thread_id=record.thread_id,
run_id=record.run_id,
duration_seconds=0,
)

View File

@ -11,7 +11,7 @@ from deerflow.runtime import RunStatus
from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY
def _checkpoint(checkpoint_id: str, messages: list[object]): def _checkpoint(checkpoint_id: str, messages: list[object], *, metadata: dict | None = None):
return SimpleNamespace( return SimpleNamespace(
config={ config={
"configurable": { "configurable": {
@ -22,6 +22,7 @@ def _checkpoint(checkpoint_id: str, messages: list[object]):
} }
}, },
checkpoint={"channel_values": {"messages": messages}}, checkpoint={"channel_values": {"messages": messages}},
metadata=metadata or {},
) )
@ -230,7 +231,7 @@ def test_prepare_regenerate_payload_requires_addressable_checkpoint_before_human
assert exc.value.status_code == 409 assert exc.value.status_code == 409
assert exc.value.detail == "Could not find an addressable checkpoint before the target user message" assert exc.value.detail == "Could not find an addressable checkpoint before the target user message"
assert checkpointer.alist_limits == [200] assert checkpointer.alist_limits == [400]
def test_prepare_regenerate_payload_reports_recent_checkpoint_scan_limit(): def test_prepare_regenerate_payload_reports_recent_checkpoint_scan_limit():
@ -258,4 +259,26 @@ def test_prepare_regenerate_payload_reports_recent_checkpoint_scan_limit():
assert exc.value.status_code == 409 assert exc.value.status_code == 409
assert exc.value.detail == "Could not locate target user message in recent checkpoint history (limit=200)" assert exc.value.detail == "Could not locate target user message in recent checkpoint history (limit=200)"
assert checkpointer.alist_limits == [200] assert checkpointer.alist_limits == [400]
def test_find_base_checkpoint_ignores_duration_only_checkpoints() -> None:
from app.gateway.routers.thread_runs import _find_base_checkpoint_before_human
human = HumanMessage(id="human-1", content="question")
duration_checkpoints = [
_checkpoint(
f"duration-{index}",
[],
metadata={"writes": {"runtime_run_duration": {"run_ids": [f"run-{index}"]}}},
)
for index in range(200)
]
base = _checkpoint("ckpt-base", [])
after_human = _checkpoint("ckpt-human", [human])
checkpointer = FakeCheckpointer([*duration_checkpoints, after_human, base])
result = asyncio.run(_find_base_checkpoint_before_human("thread-1", "human-1", _request(checkpointer, FakeEventStore([]))))
assert result is base
assert checkpointer.alist_limits == [400]

View File

@ -7,7 +7,7 @@ import pytest
from _router_auth_helpers import make_authed_test_app from _router_auth_helpers import make_authed_test_app
from fastapi import FastAPI, HTTPException from fastapi import FastAPI, HTTPException
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
from langchain_core.messages import AIMessage, HumanMessage from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from langgraph.checkpoint.base import empty_checkpoint from langgraph.checkpoint.base import empty_checkpoint
from langgraph.checkpoint.memory import InMemorySaver from langgraph.checkpoint.memory import InMemorySaver
from langgraph.store.memory import InMemoryStore from langgraph.store.memory import InMemoryStore
@ -75,21 +75,24 @@ async def _write_checkpoint(
messages: list[object], messages: list[object],
*, *,
step: int, step: int,
metadata: dict | None = None,
) -> dict: ) -> dict:
checkpoint = empty_checkpoint() checkpoint = empty_checkpoint()
checkpoint["id"] = checkpoint_id checkpoint["id"] = checkpoint_id
checkpoint["channel_values"] = {"messages": messages} checkpoint["channel_values"] = {"messages": messages}
checkpoint["channel_versions"] = {"messages": step} checkpoint["channel_versions"] = {"messages": step}
checkpoint_metadata = {
"step": step,
"source": "loop",
"writes": {"test": {"messages": messages}},
"parents": {},
"created_at": f"2026-07-05T00:00:0{step}+00:00",
}
checkpoint_metadata.update(metadata or {})
return await checkpointer.aput( return await checkpointer.aput(
{"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}, {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}},
checkpoint, checkpoint,
{ checkpoint_metadata,
"step": step,
"source": "loop",
"writes": {"test": {"messages": messages}},
"parents": {},
"created_at": f"2026-07-05T00:00:0{step}+00:00",
},
{"messages": step}, {"messages": step},
) )
@ -568,6 +571,95 @@ def test_get_thread_history_returns_iso_for_legacy_checkpoint_metadata() -> None
assert _ISO_TIMESTAMP_RE.match(entry["created_at"]), entry assert _ISO_TIMESTAMP_RE.match(entry["created_at"]), entry
def test_get_thread_history_associates_tool_messages_from_checkpoint_turn() -> None:
app, _store, checkpointer = _build_thread_app()
thread_id = "history-tool-run"
messages = [
HumanMessage(id="human-1", content="Use a tool", additional_kwargs={"run_id": "run-1"}),
AIMessage(
id="ai-1",
content="Calling tool",
tool_calls=[{"name": "lookup", "args": {}, "id": "call-1"}],
),
ToolMessage(id="tool-1", content="result", tool_call_id="call-1"),
AIMessage(id="ai-2", content="Done"),
]
asyncio.run(
_write_checkpoint(
checkpointer,
thread_id,
"checkpoint-tool-run",
messages,
step=1,
metadata={"run_durations": {"run-1": 4}},
)
)
with TestClient(app) as client:
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
assert response.status_code == 200, response.text
history_messages = response.json()[0]["values"]["messages"]
assert [message.get("run_id") for message in history_messages[1:]] == ["run-1", "run-1", "run-1"]
assert [message["additional_kwargs"]["turn_duration"] for message in history_messages if message["type"] == "ai"] == [4, 4]
def test_get_thread_history_backfills_legacy_durations_with_exact_event_run_id() -> None:
app, _store, checkpointer = _build_thread_app()
thread_id = "legacy-history-run-id"
messages = [
HumanMessage(id="human-1", content="Question", additional_kwargs={"run_id": "boundary-run"}),
AIMessage(id="ai-1", content="Answer"),
ToolMessage(id="tool-1", content="result", tool_call_id="call-1"),
]
asyncio.run(_write_checkpoint(checkpointer, thread_id, "00000000-0000-6000-8000-000000000001", messages, step=1))
async def list_by_thread(_: str) -> list[SimpleNamespace]:
return [
SimpleNamespace(
run_id="boundary-run",
created_at="2026-07-05T00:00:00+00:00",
updated_at="2026-07-05T00:00:03+00:00",
),
SimpleNamespace(
run_id="exact-run",
created_at="2026-07-05T00:00:00+00:00",
updated_at="2026-07-05T00:00:07+00:00",
),
]
async def list_messages(_: str, *, limit: int) -> list[dict]:
assert limit == 1000
return [{"content": {"type": "ai", "id": "ai-1"}, "run_id": "exact-run"}]
app.state.run_manager = SimpleNamespace(list_by_thread=list_by_thread)
app.state.run_event_store = SimpleNamespace(list_messages=list_messages)
with TestClient(app) as client:
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
assert response.status_code == 200, response.text
entry = response.json()[0]
history_messages = entry["values"]["messages"]
assert history_messages[1]["run_id"] == "exact-run"
assert history_messages[1]["additional_kwargs"]["turn_duration"] == 7
assert history_messages[2]["run_id"] == "boundary-run"
assert "run_durations" not in entry["metadata"]
latest = asyncio.run(checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}))
assert latest is not None
assert latest.metadata["run_durations"] == {"boundary-run": 3, "exact-run": 7}
def test_ai_message_lacks_duration_only_for_unannotated_ai_messages() -> None:
assert threads._ai_message_lacks_duration({"type": "ai"})
assert threads._ai_message_lacks_duration({"type": "ai", "additional_kwargs": []})
assert not threads._ai_message_lacks_duration({"type": "tool"})
assert not threads._ai_message_lacks_duration({"type": "ai", "additional_kwargs": {"turn_duration": 0}})
# ── branch threads from completed assistant turns ───────────────────────────── # ── branch threads from completed assistant turns ─────────────────────────────