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__)
router = APIRouter(prefix="/api/threads", tags=["runs"])
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
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]:
"""Map run_id -> duration in seconds from run timestamps."""
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)
base_config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
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:
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

View File

@ -19,7 +19,7 @@ import uuid
from pathlib import Path
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 pydantic import BaseModel, Field, field_validator
@ -44,6 +44,7 @@ from deerflow.runtime.goal import (
read_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.utils.file_io import run_file_io
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])
@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.
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:
serialized_msgs = serialize_channel_values_for_api({"messages": messages}).get("messages", [])
try:
from app.gateway.deps import get_run_event_store, get_run_manager
from app.gateway.routers.thread_runs import compute_run_durations
# Human messages define turn boundaries. New checkpoints
# 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)
event_store = get_run_event_store(request)
if msg.get("type") not in {"ai", "tool"} or not current_turn_run_id:
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.
# We do this full fetch because raw LangGraph messages lack a native run_id link.
# Legacy checkpoints without duration metadata are
# 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:
# 1. Map each run_id to its actual duration
run_durations = compute_run_durations(runs)
runs = await run_mgr.list_by_thread(thread_id)
events = await event_store.list_messages(thread_id, limit=1000)
# 2. Map every message id directly to its parent run_id
msg_to_run = {}
for e in events:
content = e.get("content", {})
if isinstance(content, dict) and content.get("type") == "ai" and "id" in content:
msg_to_run[content["id"]] = e["run_id"]
if runs:
run_durations = compute_run_durations(runs)
msg_to_run = {}
for event in events:
content = event.get("content", {})
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.
# Raw LangGraph checkpoint messages do not carry a
# native run link. Message events are exact when
# present, but historical/runtime stores can miss
# them; the user-input message already records the
# run id for the whole turn, so use it as the
# fallback for following AI/tool messages.
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
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
if msg.get("type") in {"ai", "tool"}:
msg_id = msg.get("id")
run_id = msg_to_run.get(msg_id) or current_turn_run_id
if msg.get("type") not in {"ai", "tool"}:
continue
run_id = msg_to_run.get(msg.get("id")) or current_turn_run_id
if run_id:
msg["run_id"] = run_id
if msg.get("type") == "ai" and run_id in run_durations:
if "additional_kwargs" not in msg:
msg["additional_kwargs"] = {}
msg["additional_kwargs"]["turn_duration"] = run_durations[run_id]
_set_message_turn_duration(msg, run_id, run_durations)
# Intentional, best-effort write-on-read migration:
# 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:
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")]
# 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
if "step" in metadata:
user_meta["step"] = metadata["step"]

View File

@ -20,7 +20,12 @@ import copy
import inspect
import logging
import os
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
@ -63,6 +68,28 @@ 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
# Valid stream_mode values for LangGraph's graph.astream()
_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:
nonlocal llm_error_fallback_message
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):
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 = _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(single_mode)
await bridge.publish(run_id, sse_event, serialize(chunk, mode=single_mode))
if single_mode == "custom":
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)
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.
await _stream_once(graph_input, initial_runnable_config)
@ -621,6 +649,25 @@ async def run_agent(
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 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 thread_store is not None:
try:
@ -1142,6 +1189,93 @@ def _title_generation_state(channel_values: dict[str, Any], graph_input: Any | N
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.

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
def _checkpoint(checkpoint_id: str, messages: list[object]):
def _checkpoint(checkpoint_id: str, messages: list[object], *, metadata: dict | None = None):
return SimpleNamespace(
config={
"configurable": {
@ -22,6 +22,7 @@ def _checkpoint(checkpoint_id: str, messages: list[object]):
}
},
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.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():
@ -258,4 +259,26 @@ def test_prepare_regenerate_payload_reports_recent_checkpoint_scan_limit():
assert exc.value.status_code == 409
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 fastapi import FastAPI, HTTPException
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.memory import InMemorySaver
from langgraph.store.memory import InMemoryStore
@ -75,21 +75,24 @@ async def _write_checkpoint(
messages: list[object],
*,
step: int,
metadata: dict | None = None,
) -> dict:
checkpoint = empty_checkpoint()
checkpoint["id"] = checkpoint_id
checkpoint["channel_values"] = {"messages": messages}
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(
{"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}},
checkpoint,
{
"step": step,
"source": "loop",
"writes": {"test": {"messages": messages}},
"parents": {},
"created_at": f"2026-07-05T00:00:0{step}+00:00",
},
checkpoint_metadata,
{"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
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 ─────────────────────────────