mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-10 14:08:52 +00:00
fix(gateway): preserve exact history attribution beyond event page limits (#4953)
* fix(gateway): preserve exact history run attribution * fix(gateway): make history migration authoritative * docs(runtime): keep history contract within guidance budget * fix(runtime): fence final run duration write
This commit is contained in:
parent
ff0a6768c2
commit
e8410cebfc
@ -66,7 +66,7 @@ from deerflow.runtime.goal import (
|
|||||||
)
|
)
|
||||||
from deerflow.runtime.journal import build_branch_history_seed_events
|
from deerflow.runtime.journal import build_branch_history_seed_events
|
||||||
from deerflow.runtime.runs.manager import ConflictError
|
from deerflow.runtime.runs.manager import ConflictError
|
||||||
from deerflow.runtime.runs.worker import valid_duration_entry
|
from deerflow.runtime.runs.worker import RUN_MESSAGE_IDS_METADATA_KEY, valid_duration_entry, valid_run_message_id_entry
|
||||||
from deerflow.runtime.secret_context import redact_metadata_secrets
|
from deerflow.runtime.secret_context import redact_metadata_secrets
|
||||||
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
|
||||||
@ -1355,6 +1355,96 @@ def _checkpoint_run_durations(metadata: Any) -> dict[str, int]:
|
|||||||
return {run_id: duration_seconds for run_id, duration_seconds in raw_durations.items() if valid_duration_entry(run_id, duration_seconds)}
|
return {run_id: duration_seconds for run_id, duration_seconds in raw_durations.items() if valid_duration_entry(run_id, duration_seconds)}
|
||||||
|
|
||||||
|
|
||||||
|
def _checkpoint_run_message_ids(metadata: Any) -> dict[str, str]:
|
||||||
|
raw_message_run_ids = metadata.get(RUN_MESSAGE_IDS_METADATA_KEY) if isinstance(metadata, dict) else None
|
||||||
|
if not isinstance(raw_message_run_ids, dict):
|
||||||
|
return {}
|
||||||
|
return {message_id: run_id for message_id, run_id in raw_message_run_ids.items() if valid_run_message_id_entry(message_id, run_id)}
|
||||||
|
|
||||||
|
|
||||||
|
async def _load_run_durations(
|
||||||
|
*,
|
||||||
|
run_manager: Any,
|
||||||
|
thread_id: str,
|
||||||
|
user_id: str | None,
|
||||||
|
run_ids: set[str],
|
||||||
|
) -> dict[str, int]:
|
||||||
|
"""Batch-hydrate the requested runs and compute their latest durations."""
|
||||||
|
if not run_ids:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
from app.gateway.routers.thread_runs import compute_run_durations
|
||||||
|
|
||||||
|
runs = await run_manager.list_by_thread(
|
||||||
|
thread_id,
|
||||||
|
user_id=user_id,
|
||||||
|
limit=max(100, len(run_ids)),
|
||||||
|
)
|
||||||
|
known_run_ids = {run.run_id for run in runs}
|
||||||
|
for run_id in sorted(run_ids - known_run_ids):
|
||||||
|
run = await run_manager.get(run_id, user_id=user_id)
|
||||||
|
if run is not None:
|
||||||
|
runs.append(run)
|
||||||
|
known_run_ids.add(run_id)
|
||||||
|
|
||||||
|
computed_durations = compute_run_durations(runs)
|
||||||
|
return {run_id: duration for run_id, duration in computed_durations.items() if run_id in run_ids}
|
||||||
|
|
||||||
|
|
||||||
|
async def _persist_run_history_metadata_background(
|
||||||
|
*,
|
||||||
|
request: Request,
|
||||||
|
checkpointer: Any,
|
||||||
|
thread_id: str,
|
||||||
|
user_id: str | None,
|
||||||
|
duration_run_ids: set[str],
|
||||||
|
message_run_ids: dict[str, str],
|
||||||
|
audited_message_ids: set[str],
|
||||||
|
) -> None:
|
||||||
|
"""Best-effort history migration behind durable checkpoint admission."""
|
||||||
|
from deerflow.runtime.runs.worker import persist_run_history_metadata
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with reserve_checkpoint_write(request, thread_id, user_id=user_id):
|
||||||
|
from app.gateway.deps import get_run_event_store, get_run_manager
|
||||||
|
|
||||||
|
authoritative_message_run_ids = dict(message_run_ids)
|
||||||
|
authoritative_duration_run_ids = set(duration_run_ids)
|
||||||
|
if audited_message_ids:
|
||||||
|
exact_after_admission = await get_run_event_store(request).find_latest_ai_message_run_ids(
|
||||||
|
thread_id,
|
||||||
|
audited_message_ids,
|
||||||
|
user_id=user_id,
|
||||||
|
)
|
||||||
|
for message_id in audited_message_ids:
|
||||||
|
exact_run_id = exact_after_admission.get(message_id)
|
||||||
|
if valid_run_message_id_entry(message_id, exact_run_id):
|
||||||
|
if authoritative_message_run_ids.get(message_id) != exact_run_id:
|
||||||
|
authoritative_duration_run_ids.add(exact_run_id)
|
||||||
|
authoritative_message_run_ids[message_id] = exact_run_id
|
||||||
|
|
||||||
|
authoritative_durations = await _load_run_durations(
|
||||||
|
run_manager=get_run_manager(request),
|
||||||
|
thread_id=thread_id,
|
||||||
|
user_id=user_id,
|
||||||
|
run_ids=authoritative_duration_run_ids,
|
||||||
|
)
|
||||||
|
|
||||||
|
await persist_run_history_metadata(
|
||||||
|
checkpointer=checkpointer,
|
||||||
|
thread_id=thread_id,
|
||||||
|
durations=authoritative_durations,
|
||||||
|
message_run_ids=authoritative_message_run_ids,
|
||||||
|
)
|
||||||
|
except ConflictError:
|
||||||
|
# A live run or another checkpoint writer owns the thread. The mapping
|
||||||
|
# is a read-through optimization, so the next history request can retry
|
||||||
|
# instead of racing a user-visible state mutation.
|
||||||
|
logger.debug("Skipped run-history metadata migration for busy thread %s", sanitize_log_param(thread_id))
|
||||||
|
except Exception:
|
||||||
|
logger.warning("Failed to persist run-history metadata for thread %s", sanitize_log_param(thread_id), exc_info=True)
|
||||||
|
|
||||||
|
|
||||||
@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(
|
async def get_thread_history(
|
||||||
@ -1408,8 +1498,10 @@ async def get_thread_history(
|
|||||||
# carry the completed turns' durations in metadata, so the
|
# carry the completed turns' durations in metadata, so the
|
||||||
# messages channel stays unchanged.
|
# messages channel stays unchanged.
|
||||||
checkpoint_run_durations = _checkpoint_run_durations(metadata)
|
checkpoint_run_durations = _checkpoint_run_durations(metadata)
|
||||||
|
checkpoint_run_message_ids = _checkpoint_run_message_ids(metadata)
|
||||||
current_turn_run_id = None
|
current_turn_run_id = None
|
||||||
turn_run_ids: set[str] = set()
|
turn_run_ids: set[str] = set()
|
||||||
|
legacy_ai_message_ids: set[str] = set()
|
||||||
for msg in serialized_msgs:
|
for msg in serialized_msgs:
|
||||||
if msg.get("type") == "human":
|
if msg.get("type") == "human":
|
||||||
additional_kwargs = msg.get("additional_kwargs")
|
additional_kwargs = msg.get("additional_kwargs")
|
||||||
@ -1419,75 +1511,121 @@ async def get_thread_history(
|
|||||||
current_turn_run_id = run_id
|
current_turn_run_id = run_id
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if msg.get("type") not in {"ai", "tool"} or not current_turn_run_id:
|
message_type = msg.get("type")
|
||||||
|
if message_type not in {"ai", "tool"}:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
msg.setdefault("run_id", current_turn_run_id)
|
if message_type == "ai":
|
||||||
if msg.get("type") == "ai":
|
message_id = msg.get("id")
|
||||||
turn_run_ids.add(current_turn_run_id)
|
persisted_run_id = checkpoint_run_message_ids.get(message_id) if isinstance(message_id, str) else None
|
||||||
|
if persisted_run_id:
|
||||||
|
msg["run_id"] = persisted_run_id
|
||||||
|
elif not isinstance(msg.get("run_id"), str) or not msg.get("run_id"):
|
||||||
|
if current_turn_run_id:
|
||||||
|
msg["run_id"] = current_turn_run_id
|
||||||
|
if isinstance(message_id, str) and message_id:
|
||||||
|
legacy_ai_message_ids.add(message_id)
|
||||||
|
|
||||||
# Stamp each run's duration on its last AI message only,
|
run_id = msg.get("run_id")
|
||||||
# same as the live message endpoints — never every AI
|
if isinstance(run_id, str) and run_id:
|
||||||
# message in a multi-message turn (#4152).
|
turn_run_ids.add(run_id)
|
||||||
stamp_turn_duration_on_last_ai(serialized_msgs, checkpoint_run_durations)
|
elif current_turn_run_id:
|
||||||
|
msg.setdefault("run_id", current_turn_run_id)
|
||||||
|
|
||||||
# Runs referenced by this checkpoint's AI messages but
|
# Runs referenced by this checkpoint's AI messages but
|
||||||
# absent from checkpoint metadata are either legacy
|
# absent from duration metadata are either legacy
|
||||||
# (never migrated) or just completed. Correlate once via
|
# (never migrated) or just completed. Exact attribution
|
||||||
# event-store + run-manager, then upgrade by a
|
# has its own completeness condition: duration-only
|
||||||
# metadata-only checkpoint write.
|
# checkpoints written before #4949 still need their AI
|
||||||
|
# IDs correlated and persisted. Correlate once via the
|
||||||
|
# event store, then hydrate only the run rows whose
|
||||||
|
# durations are actually required.
|
||||||
|
resolved_run_durations = dict(checkpoint_run_durations)
|
||||||
missing_run_ids = turn_run_ids - set(checkpoint_run_durations)
|
missing_run_ids = turn_run_ids - set(checkpoint_run_durations)
|
||||||
if missing_run_ids:
|
if missing_run_ids or legacy_ai_message_ids:
|
||||||
from app.gateway.deps import get_run_event_store, get_run_manager
|
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
|
|
||||||
|
|
||||||
run_mgr = get_run_manager(request)
|
run_mgr = get_run_manager(request)
|
||||||
event_store = get_run_event_store(request)
|
event_store = get_run_event_store(request)
|
||||||
|
user_id = get_effective_user_id()
|
||||||
runs = await run_mgr.list_by_thread(thread_id)
|
ai_message_ids = set(legacy_ai_message_ids)
|
||||||
events = await event_store.list_messages(thread_id, limit=1000)
|
try:
|
||||||
|
msg_to_run = (
|
||||||
if runs:
|
await event_store.find_latest_ai_message_run_ids(
|
||||||
run_durations = compute_run_durations(runs)
|
thread_id,
|
||||||
msg_to_run = {}
|
ai_message_ids,
|
||||||
for event in events:
|
user_id=user_id,
|
||||||
content = event.get("content", {})
|
)
|
||||||
run_id = event.get("run_id")
|
if ai_message_ids
|
||||||
if isinstance(content, dict) and content.get("type") == "ai" and "id" in content and isinstance(run_id, str) and run_id:
|
else {}
|
||||||
msg_to_run[content["id"]] = run_id
|
)
|
||||||
|
except Exception:
|
||||||
current_turn_run_id = None
|
# A failed exact lookup must not masquerade as a
|
||||||
|
# successful boundary attribution. Removing the
|
||||||
|
# synthesized ids leaves the response incomplete
|
||||||
|
# rather than deterministically wrong. Durations
|
||||||
|
# backed by persisted mappings remain provable
|
||||||
|
# and should still survive this degraded path.
|
||||||
for msg in serialized_msgs:
|
for msg in serialized_msgs:
|
||||||
if msg.get("type") == "human":
|
if msg.get("type") == "ai" and msg.get("id") in ai_message_ids:
|
||||||
additional_kwargs = msg.get("additional_kwargs")
|
msg.pop("run_id", None)
|
||||||
if isinstance(additional_kwargs, dict):
|
stamp_turn_duration_on_last_ai(
|
||||||
run_id = additional_kwargs.get("run_id")
|
serialized_msgs,
|
||||||
if isinstance(run_id, str) and run_id:
|
checkpoint_run_durations,
|
||||||
current_turn_run_id = run_id
|
)
|
||||||
continue
|
raise
|
||||||
|
|
||||||
if msg.get("type") not in {"ai", "tool"}:
|
for msg in serialized_msgs:
|
||||||
continue
|
if msg.get("type") != "ai":
|
||||||
run_id = msg_to_run.get(msg.get("id")) or current_turn_run_id
|
continue
|
||||||
if run_id:
|
exact_run_id = msg_to_run.get(msg.get("id"))
|
||||||
msg["run_id"] = run_id
|
if exact_run_id:
|
||||||
|
msg["run_id"] = exact_run_id
|
||||||
|
|
||||||
stamp_turn_duration_on_last_ai(serialized_msgs, run_durations)
|
# Cache the complete audited attribution, including
|
||||||
|
# boundary fallbacks for IDs with no event. Without
|
||||||
|
# those negative-result entries, every history read
|
||||||
|
# would rescan the same pre-event-store prefix.
|
||||||
|
message_run_ids_to_persist = {
|
||||||
|
message_id: run_id
|
||||||
|
for msg in serialized_msgs
|
||||||
|
if msg.get("type") == "ai" and isinstance((message_id := msg.get("id")), str) and message_id in ai_message_ids and isinstance((run_id := msg.get("run_id")), str) and run_id
|
||||||
|
}
|
||||||
|
required_run_ids = {run_id for msg in serialized_msgs if msg.get("type") == "ai" and isinstance((run_id := msg.get("run_id")), str) and run_id and run_id not in checkpoint_run_durations}
|
||||||
|
run_durations = await _load_run_durations(
|
||||||
|
run_manager=run_mgr,
|
||||||
|
thread_id=thread_id,
|
||||||
|
user_id=user_id,
|
||||||
|
run_ids=required_run_ids,
|
||||||
|
)
|
||||||
|
resolved_run_durations.update(run_durations)
|
||||||
|
|
||||||
# Intentional, best-effort write-on-read migration:
|
# Intentional, best-effort write-on-read migration:
|
||||||
# persist legacy metadata after the response so the
|
# persist both exact attribution and duration after
|
||||||
# history request never waits on an active stream's
|
# the response so subsequent reads stay exact without
|
||||||
# same-thread checkpoint lock.
|
# waiting on an active stream's checkpoint lock.
|
||||||
|
if required_run_ids or message_run_ids_to_persist:
|
||||||
background_tasks.add_task(
|
background_tasks.add_task(
|
||||||
persist_run_durations,
|
_persist_run_history_metadata_background,
|
||||||
|
request=request,
|
||||||
checkpointer=checkpointer,
|
checkpointer=checkpointer,
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
durations=run_durations,
|
user_id=user_id,
|
||||||
|
duration_run_ids=required_run_ids,
|
||||||
|
message_run_ids=message_run_ids_to_persist,
|
||||||
|
audited_message_ids=ai_message_ids,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Stamp only after exact attribution is final. Stamping
|
||||||
|
# the synthesized boundary first can leave its duration
|
||||||
|
# attached to a message whose run ID is later corrected.
|
||||||
|
stamp_turn_duration_on_last_ai(
|
||||||
|
serialized_msgs,
|
||||||
|
resolved_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", sanitize_log_param(thread_id), exc_info=True)
|
||||||
|
|
||||||
values["messages"] = serialized_msgs
|
values["messages"] = serialized_msgs
|
||||||
|
|
||||||
@ -1496,7 +1634,7 @@ async def get_thread_history(
|
|||||||
next_tasks = list(snapshot.next or ())
|
next_tasks = list(snapshot.next or ())
|
||||||
|
|
||||||
# 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", "run_durations")}
|
user_meta = {k: v for k, v in metadata.items() if k not in ("created_at", "updated_at", "step", "source", "writes", "parents", "run_durations", RUN_MESSAGE_IDS_METADATA_KEY)}
|
||||||
# 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"]
|
||||||
|
|||||||
@ -20,6 +20,39 @@ Checkpointer storage runs in one of two channel modes, selected by `checkpoint_c
|
|||||||
|
|
||||||
**Run rollback flow** (`runtime/runs/worker.py`): `_capture_rollback_point` materializes the complete pre-run state via the accessor and captures raw `pending_writes` via `aget_tuple` into an immutable `RollbackPoint` before the run starts — capture failure disables rollback (fail-closed), never restores partial state. In `full` mode, cancel-with-rollback forks from the pre-run checkpoint via the mutation graph and inherits non-message channels from that parent. In `delta` mode, forking is unsafe once the cancelled path has attached sibling writes to the pre-run checkpoint, so rollback replaces every captured channel on the current head, using `Overwrite` for reducers and schema defaults for current-head-only channels. Both modes reattach only the captured pre-run pending writes to the restored checkpoint. Edit replay runs (`metadata.replay_kind="edit"`) also restore the pre-run checkpoint on failed, timed-out, or interrupted completion and publish the restored `values` snapshot to the stream before `end`, so clients do not remain on a transient edited branch when the replay did not produce a successful replacement.
|
**Run rollback flow** (`runtime/runs/worker.py`): `_capture_rollback_point` materializes the complete pre-run state via the accessor and captures raw `pending_writes` via `aget_tuple` into an immutable `RollbackPoint` before the run starts — capture failure disables rollback (fail-closed), never restores partial state. In `full` mode, cancel-with-rollback forks from the pre-run checkpoint via the mutation graph and inherits non-message channels from that parent. In `delta` mode, forking is unsafe once the cancelled path has attached sibling writes to the pre-run checkpoint, so rollback replaces every captured channel on the current head, using `Overwrite` for reducers and schema defaults for current-head-only channels. Both modes reattach only the captured pre-run pending writes to the restored checkpoint. Edit replay runs (`metadata.replay_kind="edit"`) also restore the pre-run checkpoint on failed, timed-out, or interrupted completion and publish the restored `values` snapshot to the stream before `end`, so clients do not remain on a transient edited branch when the replay did not produce a successful replacement.
|
||||||
|
|
||||||
|
**Targeted run-event attribution** (`runtime/events/store/`):
|
||||||
|
`RunEventStore.find_latest_ai_message_run_ids()` has a complete-or-error
|
||||||
|
contract. Its default implementation walks `list_messages()` backward in
|
||||||
|
1000-row pages, preserves the first page's high-watermark through the exclusive
|
||||||
|
`before_seq` cursor, and raises when a full page has no safe progressing `seq`.
|
||||||
|
Memory and database stores use that bounded path; the JSONL store overrides it
|
||||||
|
with one complete thread-log read because each JSONL page would otherwise
|
||||||
|
rescan every run file. The default and JSONL paths share the public
|
||||||
|
`normalize_message_ids()` and `match_ai_message_run_id()` helpers from
|
||||||
|
`events/store/base.py`. Database owner filtering is inherited on every page.
|
||||||
|
Callers may use a missing key as proof that no valid AI event exists only after
|
||||||
|
an ordinary return, never after an exception. A caller that crosses a run or
|
||||||
|
checkpoint-write admission boundary must repeat the complete audit after
|
||||||
|
admission; a pre-admission exact hit can be superseded by a later event just as
|
||||||
|
a pre-admission miss can become an exact hit.
|
||||||
|
|
||||||
|
Gateway `POST /api/threads/{id}/history` uses that lookup to migrate legacy AI
|
||||||
|
messages. An exhaustive miss preserves the human-boundary fallback; an
|
||||||
|
incomplete lookup removes unproven synthesized IDs. Its metadata-only
|
||||||
|
write-on-read cache stores `run_message_ids` for every audited AI ID (including
|
||||||
|
exhaustive misses) plus required `run_durations`; duration presence alone does
|
||||||
|
not prove attribution. Historical `body.before` reads write the audit to the
|
||||||
|
head, and the merge may retain IDs no longer in materialized history, which
|
||||||
|
readers ignore. Migration must acquire the durable `checkpoint_write`
|
||||||
|
reservation, then repeat the whole message audit and batch-reload required run
|
||||||
|
rows before persisting. Post-admission exact hits replace foreground exact or
|
||||||
|
boundary mappings, and recomputed final durations replace foreground snapshots.
|
||||||
|
Successful workers keep their durable run row active through the final duration
|
||||||
|
checkpoint write, so a peer migration cannot enter during terminalization.
|
||||||
|
The first `RunManager.list_by_thread()` hydration page uses a 100-row floor or
|
||||||
|
the number of required IDs, whichever is larger; missing exact runs use targeted
|
||||||
|
`get()` calls.
|
||||||
|
|
||||||
**Where things live**:
|
**Where things live**:
|
||||||
- `runtime/checkpoint_mode.py` — mode + snapshot-frequency freeze, marker injection, delta detection, compatibility gate, both error types
|
- `runtime/checkpoint_mode.py` — mode + snapshot-frequency freeze, marker injection, delta detection, compatibility gate, both error types
|
||||||
- `runtime/checkpoint_state.py` — `CheckpointStateAccessor`, `build_state_mutation_graph`, `RollbackPoint`
|
- `runtime/checkpoint_state.py` — `CheckpointStateAccessor`, `build_state_mutation_graph`, `RollbackPoint`
|
||||||
|
|||||||
@ -16,6 +16,31 @@ import abc
|
|||||||
|
|
||||||
from deerflow.runtime.user_context import AUTO, _AutoSentinel
|
from deerflow.runtime.user_context import AUTO, _AutoSentinel
|
||||||
|
|
||||||
|
_AI_MESSAGE_RUN_LOOKUP_PAGE_SIZE = 1000
|
||||||
|
|
||||||
|
|
||||||
|
class IncompleteMessageRunLookupError(RuntimeError):
|
||||||
|
"""Raised when a store cannot prove that a targeted lookup is complete."""
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_message_ids(message_ids: set[str]) -> set[str]:
|
||||||
|
"""Return the non-empty string IDs that can participate in a lookup."""
|
||||||
|
return {message_id for message_id in message_ids if isinstance(message_id, str) and message_id}
|
||||||
|
|
||||||
|
|
||||||
|
def match_ai_message_run_id(event: object, message_ids: set[str]) -> tuple[str, str] | None:
|
||||||
|
"""Return a target AI message ID and its valid run ID, if present."""
|
||||||
|
if not isinstance(event, dict) or event.get("category") != "message":
|
||||||
|
return None
|
||||||
|
content = event.get("content")
|
||||||
|
run_id = event.get("run_id")
|
||||||
|
if not isinstance(content, dict) or content.get("type") != "ai" or not isinstance(run_id, str) or not run_id:
|
||||||
|
return None
|
||||||
|
message_id = content.get("id")
|
||||||
|
if not isinstance(message_id, str) or message_id not in message_ids:
|
||||||
|
return None
|
||||||
|
return message_id, run_id
|
||||||
|
|
||||||
|
|
||||||
class RunEventStore(abc.ABC):
|
class RunEventStore(abc.ABC):
|
||||||
"""Run event stream storage interface.
|
"""Run event stream storage interface.
|
||||||
@ -27,6 +52,8 @@ class RunEventStore(abc.ABC):
|
|||||||
4. list_events() returns all events for the specified run
|
4. list_events() returns all events for the specified run
|
||||||
5. Returned dicts contain the required RunEvent envelope fields; backends
|
5. Returned dicts contain the required RunEvent envelope fields; backends
|
||||||
may add documented fields such as DbRunEventStore.user_id
|
may add documented fields such as DbRunEventStore.user_id
|
||||||
|
6. find_latest_ai_message_run_ids() returns the newest valid AI message
|
||||||
|
event for each requested ID and performs no storage work for empty input
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
@ -92,6 +119,70 @@ class RunEventStore(abc.ABC):
|
|||||||
user-scoped backends must apply it according to their isolation model.
|
user-scoped backends must apply it according to their isolation model.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
async def find_latest_ai_message_run_ids(
|
||||||
|
self,
|
||||||
|
thread_id: str,
|
||||||
|
message_ids: set[str],
|
||||||
|
*,
|
||||||
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
) -> dict[str, str]:
|
||||||
|
"""Map target message IDs to their newest valid AI event's run ID.
|
||||||
|
|
||||||
|
Only ``category="message"`` events whose structured content has
|
||||||
|
``type="ai"`` and whose ``run_id`` is a non-empty string qualify. An
|
||||||
|
empty target set must return immediately without storage work. The
|
||||||
|
default implementation pages backward in bounded windows. It raises
|
||||||
|
:class:`IncompleteMessageRunLookupError` instead of returning a
|
||||||
|
partial result when a full page lacks a safe, progressing ``seq``
|
||||||
|
cursor; callers may only treat an ordinary return as an exhaustive
|
||||||
|
lookup for unresolved IDs.
|
||||||
|
|
||||||
|
``user_id`` follows the same explicit-caller semantics as
|
||||||
|
:meth:`list_messages`.
|
||||||
|
"""
|
||||||
|
pending = normalize_message_ids(message_ids)
|
||||||
|
if not pending:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
result: dict[str, str] = {}
|
||||||
|
before_seq: int | None = None
|
||||||
|
while pending:
|
||||||
|
page = await self.list_messages(
|
||||||
|
thread_id,
|
||||||
|
limit=_AI_MESSAGE_RUN_LOOKUP_PAGE_SIZE,
|
||||||
|
before_seq=before_seq,
|
||||||
|
user_id=user_id,
|
||||||
|
)
|
||||||
|
if not page:
|
||||||
|
break
|
||||||
|
|
||||||
|
for event in reversed(page):
|
||||||
|
match = match_ai_message_run_id(event, pending)
|
||||||
|
if match is None:
|
||||||
|
continue
|
||||||
|
message_id, run_id = match
|
||||||
|
result[message_id] = run_id
|
||||||
|
pending.remove(message_id)
|
||||||
|
if not pending:
|
||||||
|
break
|
||||||
|
|
||||||
|
if not pending or len(page) < _AI_MESSAGE_RUN_LOOKUP_PAGE_SIZE:
|
||||||
|
break
|
||||||
|
|
||||||
|
seqs: list[int] = []
|
||||||
|
for event in page:
|
||||||
|
seq = event.get("seq") if isinstance(event, dict) else None
|
||||||
|
if not isinstance(seq, int) or isinstance(seq, bool):
|
||||||
|
raise IncompleteMessageRunLookupError("Run event lookup could not form a safe backward cursor from a full page")
|
||||||
|
seqs.append(seq)
|
||||||
|
|
||||||
|
next_before_seq = min(seqs)
|
||||||
|
if before_seq is not None and next_before_seq >= before_seq:
|
||||||
|
raise IncompleteMessageRunLookupError("Run event lookup could not form a safe backward cursor because seq did not progress")
|
||||||
|
before_seq = next_before_seq
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
async def list_events(
|
async def list_events(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@ -30,7 +30,7 @@ from datetime import UTC, datetime
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from deerflow.runtime.events.store.base import RunEventStore
|
from deerflow.runtime.events.store.base import RunEventStore, match_ai_message_run_id, normalize_message_ids
|
||||||
from deerflow.runtime.user_context import AUTO, _AutoSentinel
|
from deerflow.runtime.user_context import AUTO, _AutoSentinel
|
||||||
from deerflow.utils.thread_id import validate_thread_id
|
from deerflow.utils.thread_id import validate_thread_id
|
||||||
|
|
||||||
@ -280,6 +280,34 @@ class JsonlRunEventStore(RunEventStore):
|
|||||||
else:
|
else:
|
||||||
return messages[-limit:]
|
return messages[-limit:]
|
||||||
|
|
||||||
|
async def find_latest_ai_message_run_ids(
|
||||||
|
self,
|
||||||
|
thread_id: str,
|
||||||
|
message_ids: set[str],
|
||||||
|
*,
|
||||||
|
user_id: str | None | _AutoSentinel = AUTO,
|
||||||
|
) -> dict[str, str]:
|
||||||
|
pending = normalize_message_ids(message_ids)
|
||||||
|
if not pending:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
# Keep the one-pass view stable against this backend's supported
|
||||||
|
# single-process writers. Without the write lock, reading run files one
|
||||||
|
# by one can mix events from opposite sides of a concurrent append.
|
||||||
|
async with self._get_write_lock(thread_id):
|
||||||
|
events = await asyncio.to_thread(self._read_thread_events, thread_id)
|
||||||
|
result: dict[str, str] = {}
|
||||||
|
for event in reversed(events):
|
||||||
|
match = match_ai_message_run_id(event, pending)
|
||||||
|
if match is None:
|
||||||
|
continue
|
||||||
|
message_id, run_id = match
|
||||||
|
result[message_id] = run_id
|
||||||
|
pending.remove(message_id)
|
||||||
|
if not pending:
|
||||||
|
break
|
||||||
|
return result
|
||||||
|
|
||||||
async def list_events(self, thread_id, run_id, *, event_types=None, task_id=None, limit=500, after_seq=None):
|
async def list_events(self, thread_id, run_id, *, event_types=None, task_id=None, limit=500, after_seq=None):
|
||||||
events = await asyncio.to_thread(self._read_run_events, thread_id, run_id)
|
events = await asyncio.to_thread(self._read_run_events, thread_id, run_id)
|
||||||
if event_types is not None:
|
if event_types is not None:
|
||||||
|
|||||||
@ -549,6 +549,26 @@ class RunManager:
|
|||||||
except Exception:
|
except Exception:
|
||||||
logger.warning("Failed to persist run progress for %s", run_id, exc_info=True)
|
logger.warning("Failed to persist run progress for %s", run_id, exc_info=True)
|
||||||
|
|
||||||
|
async def update_finalizing_progress(self, run_id: str, **kwargs) -> None:
|
||||||
|
"""Persist final fields while the durable row is deliberately active."""
|
||||||
|
should_persist = False
|
||||||
|
async with self._lock:
|
||||||
|
record = self._runs.get(run_id)
|
||||||
|
if record is not None and not record.ownership_lost:
|
||||||
|
should_persist = record.status not in (RunStatus.pending, RunStatus.running)
|
||||||
|
if should_persist:
|
||||||
|
for key, value in kwargs.items():
|
||||||
|
if hasattr(record, key) and value is not None:
|
||||||
|
setattr(record, key, value)
|
||||||
|
record.updated_at = _now_iso()
|
||||||
|
if should_persist and self._store is not None:
|
||||||
|
try:
|
||||||
|
# The local status is already staged as terminal, but the store
|
||||||
|
# row intentionally remains running until checkpoint finalization.
|
||||||
|
await self._store.update_run_progress(run_id, **kwargs)
|
||||||
|
except Exception:
|
||||||
|
logger.warning("Failed to persist finalizing progress for %s", run_id, exc_info=True)
|
||||||
|
|
||||||
async def create(
|
async def create(
|
||||||
self,
|
self,
|
||||||
thread_id: str,
|
thread_id: str,
|
||||||
|
|||||||
@ -626,6 +626,7 @@ async def run_agent(
|
|||||||
# checkpoint failures / cancellation while waiting did not write an empty
|
# checkpoint failures / cancellation while waiting did not write an empty
|
||||||
# completion snapshot into RunStore.
|
# completion snapshot into RunStore.
|
||||||
persist_completion = False
|
persist_completion = False
|
||||||
|
completion_data: dict[str, Any] | None = None
|
||||||
# Buffers subagent step events for batched persistence (#3779); assigned once
|
# Buffers subagent step events for batched persistence (#3779); assigned once
|
||||||
# streaming starts and flushed in the finally block. Pre-bound to None so the
|
# 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.
|
# finally is safe even if an exception fires before streaming begins.
|
||||||
@ -1226,6 +1227,37 @@ async def run_agent(
|
|||||||
persist=False,
|
persist=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if not record.ownership_lost and journal is not None and persist_completion:
|
||||||
|
try:
|
||||||
|
# Advance the final completion fields and timestamp without
|
||||||
|
# terminalizing the durable row. That active row continues to
|
||||||
|
# fence peer checkpoint writers through the duration write.
|
||||||
|
completion_data = journal.get_completion_data()
|
||||||
|
await run_manager.update_finalizing_progress(run_id, **completion_data)
|
||||||
|
except Exception:
|
||||||
|
logger.warning("Failed to persist finalizing run progress for %s (non-fatal)", run_id, exc_info=True)
|
||||||
|
|
||||||
|
# Keep the durable run row active through its final duration checkpoint
|
||||||
|
# write. A peer Gateway admits history migration from the durable row,
|
||||||
|
# not this worker's staged terminal status; terminalizing first would
|
||||||
|
# let that migration read an unfinished lifetime and race this write.
|
||||||
|
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)
|
||||||
|
|
||||||
if not record.ownership_lost and event_store is not None:
|
if not record.ownership_lost and event_store is not None:
|
||||||
try:
|
try:
|
||||||
# Even after bounded receipt retries are exhausted, persist the
|
# Even after bounded receipt retries are exhausted, persist the
|
||||||
@ -1250,8 +1282,8 @@ async def run_agent(
|
|||||||
if not record.ownership_lost and journal is not None and persist_completion:
|
if not record.ownership_lost and journal is not None and persist_completion:
|
||||||
try:
|
try:
|
||||||
# Persist token usage + convenience fields to RunStore
|
# Persist token usage + convenience fields to RunStore
|
||||||
completion = journal.get_completion_data()
|
completion_data = completion_data or journal.get_completion_data()
|
||||||
await run_manager.update_run_completion(run_id, status=record.status.value, **completion)
|
await run_manager.update_run_completion(run_id, status=record.status.value, **completion_data)
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.warning("Failed to persist run completion for %s (non-fatal)", run_id, exc_info=True)
|
logger.warning("Failed to persist run completion for %s (non-fatal)", run_id, exc_info=True)
|
||||||
|
|
||||||
@ -1276,25 +1308,6 @@ 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 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
|
# Update threads_meta status based on run outcome
|
||||||
if started and not record.ownership_lost and thread_store is not None:
|
if started and not record.ownership_lost and thread_store is not None:
|
||||||
try:
|
try:
|
||||||
@ -2085,21 +2098,36 @@ def valid_duration_entry(run_id: Any, duration_seconds: Any) -> bool:
|
|||||||
return isinstance(run_id, str) and bool(run_id) and isinstance(duration_seconds, int) and not isinstance(duration_seconds, bool)
|
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(
|
RUN_MESSAGE_IDS_METADATA_KEY = "run_message_ids"
|
||||||
|
|
||||||
|
|
||||||
|
def valid_run_message_id_entry(message_id: Any, run_id: Any) -> bool:
|
||||||
|
"""Check that a persisted legacy message-to-run attribution is well formed."""
|
||||||
|
return isinstance(message_id, str) and bool(message_id) and isinstance(run_id, str) and bool(run_id)
|
||||||
|
|
||||||
|
|
||||||
|
async def persist_run_history_metadata(
|
||||||
*,
|
*,
|
||||||
checkpointer: Any,
|
checkpointer: Any,
|
||||||
thread_id: str,
|
thread_id: str,
|
||||||
durations: dict[str, int],
|
durations: dict[str, int] | None = None,
|
||||||
|
message_run_ids: dict[str, str] | None = None,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Merge validated run durations into a metadata-only checkpoint.
|
"""Merge validated run history indexes into a metadata-only checkpoint.
|
||||||
|
|
||||||
Durations accumulate so the history fast path can serve every known turn
|
Durations accumulate so the history fast path can serve every known turn
|
||||||
from the latest checkpoint. Per-entry overhead is negligible (~50 bytes
|
from the latest checkpoint. Legacy AI-message attributions are persisted
|
||||||
per run_id) compared to the messages channel blob written on every graph
|
alongside them for every audited AI ID, including boundary fallbacks whose
|
||||||
checkpoint, so no pruning is needed.
|
event lookup was exhaustively empty. The full mapping is deliberate: it is
|
||||||
|
both the exact-attribution cache and the negative-result coverage proof.
|
||||||
|
While the materialized message set at the head remains unchanged, later
|
||||||
|
reads query only uncached IDs. This metadata-only merge retains existing
|
||||||
|
entries, so compaction timing or historical migration may leave stale IDs;
|
||||||
|
reads ignore them because they only consult IDs in the materialized history.
|
||||||
"""
|
"""
|
||||||
updates = {run_id: max(0, duration_seconds) for run_id, duration_seconds in durations.items() if valid_duration_entry(run_id, duration_seconds)}
|
duration_updates = {run_id: max(0, duration_seconds) for run_id, duration_seconds in (durations or {}).items() if valid_duration_entry(run_id, duration_seconds)}
|
||||||
if not updates:
|
message_run_id_updates = {message_id: run_id for message_id, run_id in (message_run_ids or {}).items() if valid_run_message_id_entry(message_id, run_id)}
|
||||||
|
if not duration_updates and not message_run_id_updates:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
ckpt_config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
|
ckpt_config = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
|
||||||
@ -2113,11 +2141,15 @@ async def persist_run_durations(
|
|||||||
metadata = dict(getattr(ckpt_tuple, "metadata", {}) or {})
|
metadata = dict(getattr(ckpt_tuple, "metadata", {}) or {})
|
||||||
raw_run_durations = metadata.get("run_durations")
|
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 {}
|
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}
|
raw_message_run_ids = metadata.get(RUN_MESSAGE_IDS_METADATA_KEY)
|
||||||
if not changed_durations:
|
run_message_ids = {message_id: run_id for message_id, run_id in raw_message_run_ids.items() if valid_run_message_id_entry(message_id, run_id)} if isinstance(raw_message_run_ids, dict) else {}
|
||||||
|
changed_durations = {run_id: duration for run_id, duration in duration_updates.items() if run_durations.get(run_id) != duration}
|
||||||
|
changed_message_run_ids = {message_id: run_id for message_id, run_id in message_run_id_updates.items() if run_message_ids.get(message_id) != run_id}
|
||||||
|
if not changed_durations and not changed_message_run_ids:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
run_durations.update(changed_durations)
|
run_durations.update(changed_durations)
|
||||||
|
run_message_ids.update(changed_message_run_ids)
|
||||||
parent_checkpoint_id = _checkpoint_identity(ckpt_tuple, checkpoint)
|
parent_checkpoint_id = _checkpoint_identity(ckpt_tuple, checkpoint)
|
||||||
latest_tuple = await _call_checkpointer_method(checkpointer, "aget_tuple", "get_tuple", ckpt_config)
|
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 {}
|
latest_checkpoint = dict(getattr(latest_tuple, "checkpoint", {}) or {}) if latest_tuple is not None else {}
|
||||||
@ -2129,7 +2161,16 @@ async def persist_run_durations(
|
|||||||
prev_step = metadata.get("step")
|
prev_step = metadata.get("step")
|
||||||
metadata["step"] = (prev_step + 1) if isinstance(prev_step, int) else 1
|
metadata["step"] = (prev_step + 1) if isinstance(prev_step, int) else 1
|
||||||
metadata["run_durations"] = run_durations
|
metadata["run_durations"] = run_durations
|
||||||
metadata["writes"] = {"runtime_run_duration": {"run_ids": sorted(changed_durations)}}
|
if run_message_ids:
|
||||||
|
metadata[RUN_MESSAGE_IDS_METADATA_KEY] = run_message_ids
|
||||||
|
else:
|
||||||
|
metadata.pop(RUN_MESSAGE_IDS_METADATA_KEY, None)
|
||||||
|
metadata["writes"] = {
|
||||||
|
"runtime_run_duration": {
|
||||||
|
"run_ids": sorted(changed_durations),
|
||||||
|
"message_ids": sorted(changed_message_run_ids),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
checkpoint_ns = _checkpoint_namespace(ckpt_tuple)
|
checkpoint_ns = _checkpoint_namespace(ckpt_tuple)
|
||||||
write_config = {
|
write_config = {
|
||||||
@ -2152,6 +2193,20 @@ async def persist_run_durations(
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
async def persist_run_durations(
|
||||||
|
*,
|
||||||
|
checkpointer: Any,
|
||||||
|
thread_id: str,
|
||||||
|
durations: dict[str, int],
|
||||||
|
) -> bool:
|
||||||
|
"""Merge validated run durations into a metadata-only checkpoint."""
|
||||||
|
return await persist_run_history_metadata(
|
||||||
|
checkpointer=checkpointer,
|
||||||
|
thread_id=thread_id,
|
||||||
|
durations=durations,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def _persist_run_duration(
|
async def _persist_run_duration(
|
||||||
*,
|
*,
|
||||||
checkpointer: Any,
|
checkpointer: Any,
|
||||||
|
|||||||
@ -9,9 +9,12 @@ from langgraph.checkpoint.base import empty_checkpoint, uuid6
|
|||||||
from langgraph.checkpoint.memory import InMemorySaver
|
from langgraph.checkpoint.memory import InMemorySaver
|
||||||
|
|
||||||
import deerflow.runtime.runs.worker as worker
|
import deerflow.runtime.runs.worker as worker
|
||||||
|
from deerflow.runtime import ConflictError, ThreadOperationKind
|
||||||
|
from deerflow.runtime.events.store.memory import MemoryRunEventStore
|
||||||
from deerflow.runtime.goal import goal_thread_lock
|
from deerflow.runtime.goal import goal_thread_lock
|
||||||
from deerflow.runtime.runs.manager import RunManager, RunStartOutcome
|
from deerflow.runtime.runs.manager import RunManager, RunStartOutcome
|
||||||
from deerflow.runtime.runs.schemas import RunStatus
|
from deerflow.runtime.runs.schemas import RunStatus
|
||||||
|
from deerflow.runtime.runs.store.memory import MemoryRunStore
|
||||||
from deerflow.runtime.runs.worker import RunContext, _persist_run_duration, run_agent
|
from deerflow.runtime.runs.worker import RunContext, _persist_run_duration, run_agent
|
||||||
|
|
||||||
|
|
||||||
@ -312,6 +315,75 @@ async def test_agent_stream_serializes_with_duration_checkpoint_write() -> None:
|
|||||||
assert finished_during_stream is False
|
assert finished_during_stream is False
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_successful_run_stays_durably_active_through_final_duration_write(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
"""A peer migration cannot enter before the terminal duration is stored."""
|
||||||
|
run_store = MemoryRunStore()
|
||||||
|
owner = RunManager(store=run_store, worker_id="duration-owner")
|
||||||
|
peer = RunManager(store=run_store, worker_id="duration-peer")
|
||||||
|
record = await owner.create_or_reject("duration-finalization-admission")
|
||||||
|
checkpointer = InMemorySaver()
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
observed_store_statuses: list[str] = []
|
||||||
|
persist_duration = worker._persist_run_duration
|
||||||
|
|
||||||
|
async def assert_active_then_persist(**kwargs) -> None:
|
||||||
|
stored = await run_store.get(record.run_id)
|
||||||
|
assert stored is not None
|
||||||
|
observed_store_statuses.append(stored["status"])
|
||||||
|
with pytest.raises(ConflictError):
|
||||||
|
async with peer.reserve_thread_operation(
|
||||||
|
record.thread_id,
|
||||||
|
kind=ThreadOperationKind.checkpoint_write,
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
await persist_duration(**kwargs)
|
||||||
|
|
||||||
|
class DummyAgent:
|
||||||
|
async def astream(self, graph_input, config=None, stream_mode=None, subgraphs=False):
|
||||||
|
yield {"messages": []}
|
||||||
|
|
||||||
|
monkeypatch.setattr(worker, "_persist_run_duration", assert_active_then_persist)
|
||||||
|
|
||||||
|
await run_agent(
|
||||||
|
SimpleNamespace(
|
||||||
|
publish=AsyncMock(),
|
||||||
|
publish_end=AsyncMock(),
|
||||||
|
cleanup=AsyncMock(),
|
||||||
|
),
|
||||||
|
owner,
|
||||||
|
record,
|
||||||
|
ctx=RunContext(
|
||||||
|
checkpointer=checkpointer,
|
||||||
|
event_store=MemoryRunEventStore(),
|
||||||
|
),
|
||||||
|
agent_factory=lambda *, config: DummyAgent(),
|
||||||
|
graph_input={},
|
||||||
|
config={},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert observed_store_statuses == [RunStatus.running.value]
|
||||||
|
stored = await run_store.get(record.run_id)
|
||||||
|
assert stored is not None
|
||||||
|
assert stored["status"] == RunStatus.success.value
|
||||||
|
latest = await checkpointer.aget_tuple({"configurable": {"thread_id": record.thread_id, "checkpoint_ns": ""}})
|
||||||
|
assert latest is not None
|
||||||
|
assert record.run_id in latest.metadata["run_durations"]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.anyio
|
@pytest.mark.anyio
|
||||||
async def test_agent_stream_allows_graph_goal_state_access() -> None:
|
async def test_agent_stream_allows_graph_goal_state_access() -> None:
|
||||||
"""A graph node may acquire the goal lock while a run is streaming."""
|
"""A graph node may acquire the goal lock while a run is streaming."""
|
||||||
|
|||||||
@ -4,6 +4,8 @@ Uses a helper to create the store for each backend type.
|
|||||||
Memory tests run directly; DB and JSONL tests create stores inside each test.
|
Memory tests run directly; DB and JSONL tests create stores inside each test.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from deerflow.runtime.events.store.memory import MemoryRunEventStore
|
from deerflow.runtime.events.store.memory import MemoryRunEventStore
|
||||||
@ -14,6 +16,73 @@ def store():
|
|||||||
return MemoryRunEventStore()
|
return MemoryRunEventStore()
|
||||||
|
|
||||||
|
|
||||||
|
async def _assert_find_latest_ai_message_run_ids_contract(store, *, allow_empty_run_id: bool) -> None:
|
||||||
|
assert await store.find_latest_ai_message_run_ids("t1", set()) == {}
|
||||||
|
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="trace-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="trace",
|
||||||
|
content={"type": "ai", "id": "target"},
|
||||||
|
)
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="string-content-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content="target",
|
||||||
|
)
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="human-run",
|
||||||
|
event_type="human_message",
|
||||||
|
category="message",
|
||||||
|
content={"type": "human", "id": "target"},
|
||||||
|
)
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="old-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "target"},
|
||||||
|
)
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="other-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "other"},
|
||||||
|
)
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="decoy-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "decoy", "note": "target"},
|
||||||
|
)
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="new-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "target"},
|
||||||
|
)
|
||||||
|
if allow_empty_run_id:
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "target"},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert await store.find_latest_ai_message_run_ids("t1", {"target", "other", "missing"}, user_id=None) == {
|
||||||
|
"target": "new-run",
|
||||||
|
"other": "other-run",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
# -- Basic write and query --
|
# -- Basic write and query --
|
||||||
|
|
||||||
|
|
||||||
@ -165,6 +234,121 @@ class TestListMessages:
|
|||||||
assert [m["seq"] for m in await store.list_messages("t1", after_seq=5, limit=5)] == [7, 9]
|
assert [m["seq"] for m in await store.list_messages("t1", after_seq=5, limit=5)] == [7, 9]
|
||||||
|
|
||||||
|
|
||||||
|
class TestFindLatestAiMessageRunIds:
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_memory_contract(self, store):
|
||||||
|
await _assert_find_latest_ai_message_run_ids_contract(store, allow_empty_run_id=True)
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_memory_stops_after_all_targets_are_found(self, store):
|
||||||
|
from deerflow.runtime.events.store import base as event_store_base
|
||||||
|
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="old-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "old"},
|
||||||
|
)
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="new-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "target"},
|
||||||
|
)
|
||||||
|
|
||||||
|
store.list_messages = AsyncMock(wraps=store.list_messages)
|
||||||
|
with patch.object(event_store_base, "match_ai_message_run_id", wraps=event_store_base.match_ai_message_run_id) as match_event:
|
||||||
|
assert await store.find_latest_ai_message_run_ids("t1", {"target"}, user_id=None) == {"target": "new-run"}
|
||||||
|
assert match_event.call_count == 1
|
||||||
|
store.list_messages.assert_awaited_once_with("t1", limit=1000, before_seq=None, user_id=None)
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_memory_pages_in_bounded_windows_and_keeps_initial_high_watermark(self, store):
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="old-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "target"},
|
||||||
|
)
|
||||||
|
await store.put_batch(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"thread_id": "t1",
|
||||||
|
"run_id": "noise-run",
|
||||||
|
"event_type": "llm.ai.response",
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": f"noise-{index}"},
|
||||||
|
}
|
||||||
|
for index in range(1000)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
original_list_messages = store.list_messages
|
||||||
|
calls: list[dict] = []
|
||||||
|
|
||||||
|
async def list_messages(*args, **kwargs):
|
||||||
|
page = await original_list_messages(*args, **kwargs)
|
||||||
|
calls.append(kwargs)
|
||||||
|
if len(calls) == 1:
|
||||||
|
# This duplicate is newer than the first page's snapshot. A
|
||||||
|
# backward cursor must not let it replace the older answer
|
||||||
|
# while resolving the rest of that same lookup.
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="concurrent-new-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "target"},
|
||||||
|
)
|
||||||
|
return page
|
||||||
|
|
||||||
|
store.list_messages = AsyncMock(side_effect=list_messages)
|
||||||
|
|
||||||
|
assert await store.find_latest_ai_message_run_ids("t1", {"target"}, user_id=None) == {"target": "old-run"}
|
||||||
|
assert len(calls) == 2
|
||||||
|
assert all(call["limit"] == 1000 for call in calls)
|
||||||
|
assert calls[0].get("before_seq") is None
|
||||||
|
assert calls[1]["before_seq"] == 2
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
@pytest.mark.parametrize("malformed_page", ["missing-seq", "non-progressing-seq"])
|
||||||
|
async def test_default_lookup_raises_instead_of_looping_on_unsafe_cursor(self, store, malformed_page):
|
||||||
|
from deerflow.runtime.events.store.base import RunEventStore
|
||||||
|
|
||||||
|
calls = 0
|
||||||
|
|
||||||
|
async def list_messages(*_args, **_kwargs):
|
||||||
|
nonlocal calls
|
||||||
|
calls += 1
|
||||||
|
if malformed_page == "missing-seq":
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": f"noise-{index}"},
|
||||||
|
"run_id": "noise-run",
|
||||||
|
}
|
||||||
|
for index in range(1000)
|
||||||
|
]
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": f"noise-{index}"},
|
||||||
|
"run_id": "noise-run",
|
||||||
|
"seq": index + 1,
|
||||||
|
}
|
||||||
|
for index in range(1000)
|
||||||
|
]
|
||||||
|
|
||||||
|
store.list_messages = AsyncMock(side_effect=list_messages)
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="safe backward cursor"):
|
||||||
|
await RunEventStore.find_latest_ai_message_run_ids(store, "t1", {"missing"}, user_id=None)
|
||||||
|
assert calls == (1 if malformed_page == "missing-seq" else 2)
|
||||||
|
|
||||||
|
|
||||||
# -- list_events --
|
# -- list_events --
|
||||||
|
|
||||||
|
|
||||||
@ -365,6 +549,154 @@ class TestDbRunEventStore:
|
|||||||
|
|
||||||
await close_engine()
|
await close_engine()
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_find_latest_ai_message_run_ids_contract_and_owner_filter(self, tmp_path):
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
from deerflow.persistence.engine import close_engine, get_session_factory, init_engine
|
||||||
|
from deerflow.runtime.events.store.db import DbRunEventStore
|
||||||
|
from deerflow.runtime.user_context import reset_current_user, set_current_user
|
||||||
|
|
||||||
|
url = f"sqlite+aiosqlite:///{tmp_path / 'test.db'}"
|
||||||
|
await init_engine("sqlite", url=url, sqlite_dir=str(tmp_path))
|
||||||
|
try:
|
||||||
|
store = DbRunEventStore(get_session_factory())
|
||||||
|
await _assert_find_latest_ai_message_run_ids_contract(store, allow_empty_run_id=True)
|
||||||
|
owner_a_token = set_current_user(SimpleNamespace(id="owner-a"))
|
||||||
|
try:
|
||||||
|
await store.put(
|
||||||
|
thread_id="owned-thread",
|
||||||
|
run_id="owner-a-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "shared-id"},
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
reset_current_user(owner_a_token)
|
||||||
|
|
||||||
|
owner_b_token = set_current_user(SimpleNamespace(id="owner-b"))
|
||||||
|
try:
|
||||||
|
await store.put(
|
||||||
|
thread_id="owned-thread",
|
||||||
|
run_id="owner-b-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "shared-id"},
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
reset_current_user(owner_b_token)
|
||||||
|
|
||||||
|
assert await store.find_latest_ai_message_run_ids("owned-thread", {"shared-id"}, user_id="owner-a") == {"shared-id": "owner-a-run"}
|
||||||
|
assert await store.find_latest_ai_message_run_ids("owned-thread", {"shared-id"}, user_id="owner-b") == {"shared-id": "owner-b-run"}
|
||||||
|
finally:
|
||||||
|
await close_engine()
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_find_latest_ai_message_run_ids_handles_large_target_sets_and_special_ids(self, tmp_path):
|
||||||
|
from deerflow.persistence.engine import close_engine, get_session_factory, init_engine
|
||||||
|
from deerflow.runtime.events.store.db import DbRunEventStore
|
||||||
|
|
||||||
|
url = f"sqlite+aiosqlite:///{tmp_path / 'test.db'}"
|
||||||
|
await init_engine("sqlite", url=url, sqlite_dir=str(tmp_path))
|
||||||
|
try:
|
||||||
|
store = DbRunEventStore(get_session_factory())
|
||||||
|
target_ids = {f"id-{index:03d}" for index in range(201)}
|
||||||
|
special_id = 'message-%_/"-雪'
|
||||||
|
target_ids.add(special_id)
|
||||||
|
await store.put_batch(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"thread_id": "t1",
|
||||||
|
"run_id": "first-chunk-run",
|
||||||
|
"event_type": "llm.ai.response",
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": "id-000"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"thread_id": "t1",
|
||||||
|
"run_id": "last-chunk-run",
|
||||||
|
"event_type": "llm.ai.response",
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": "id-200"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"thread_id": "t1",
|
||||||
|
"run_id": "special-run",
|
||||||
|
"event_type": "llm.ai.response",
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": special_id},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert await store.find_latest_ai_message_run_ids("t1", target_ids, user_id=None) == {
|
||||||
|
"id-000": "first-chunk-run",
|
||||||
|
"id-200": "last-chunk-run",
|
||||||
|
special_id: "special-run",
|
||||||
|
}
|
||||||
|
finally:
|
||||||
|
await close_engine()
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_find_latest_ai_message_run_ids_pages_db_with_owner_scoped_high_watermark(self, tmp_path):
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
from deerflow.persistence.engine import close_engine, get_session_factory, init_engine
|
||||||
|
from deerflow.runtime.events.store.db import DbRunEventStore
|
||||||
|
from deerflow.runtime.user_context import reset_current_user, set_current_user
|
||||||
|
|
||||||
|
url = f"sqlite+aiosqlite:///{tmp_path / 'test.db'}"
|
||||||
|
await init_engine("sqlite", url=url, sqlite_dir=str(tmp_path))
|
||||||
|
owner_token = set_current_user(SimpleNamespace(id="owner-a"))
|
||||||
|
try:
|
||||||
|
store = DbRunEventStore(get_session_factory())
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="old-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "target"},
|
||||||
|
)
|
||||||
|
await store.put_batch(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"thread_id": "t1",
|
||||||
|
"run_id": "noise-run",
|
||||||
|
"event_type": "llm.ai.response",
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": f"noise-{index}"},
|
||||||
|
}
|
||||||
|
for index in range(1000)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
original_list_messages = store.list_messages
|
||||||
|
calls: list[dict] = []
|
||||||
|
|
||||||
|
async def list_messages(*args, **kwargs):
|
||||||
|
page = await original_list_messages(*args, **kwargs)
|
||||||
|
calls.append(kwargs)
|
||||||
|
if len(calls) == 1:
|
||||||
|
await store.put(
|
||||||
|
thread_id="t1",
|
||||||
|
run_id="concurrent-new-run",
|
||||||
|
event_type="llm.ai.response",
|
||||||
|
category="message",
|
||||||
|
content={"type": "ai", "id": "target"},
|
||||||
|
)
|
||||||
|
return page
|
||||||
|
|
||||||
|
store.list_messages = AsyncMock(side_effect=list_messages)
|
||||||
|
|
||||||
|
assert await store.find_latest_ai_message_run_ids("t1", {"target"}, user_id="owner-a") == {"target": "old-run"}
|
||||||
|
assert len(calls) == 2
|
||||||
|
assert all(call["limit"] == 1000 and call["user_id"] == "owner-a" for call in calls)
|
||||||
|
assert calls[0].get("before_seq") is None
|
||||||
|
assert calls[1]["before_seq"] == 2
|
||||||
|
finally:
|
||||||
|
reset_current_user(owner_token)
|
||||||
|
await close_engine()
|
||||||
|
|
||||||
@pytest.mark.anyio
|
@pytest.mark.anyio
|
||||||
async def test_put_if_absent_is_idempotent(self, tmp_path):
|
async def test_put_if_absent_is_idempotent(self, tmp_path):
|
||||||
from deerflow.persistence.engine import close_engine, get_session_factory, init_engine
|
from deerflow.persistence.engine import close_engine, get_session_factory, init_engine
|
||||||
@ -749,6 +1081,44 @@ class TestJsonlRunEventStore:
|
|||||||
messages = await s.list_messages("t1")
|
messages = await s.list_messages("t1")
|
||||||
assert len(messages) == 1
|
assert len(messages) == 1
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_find_latest_ai_message_run_ids_contract(self, tmp_path):
|
||||||
|
from deerflow.runtime.events.store.jsonl import JsonlRunEventStore
|
||||||
|
|
||||||
|
store = JsonlRunEventStore(base_dir=tmp_path / "jsonl")
|
||||||
|
await _assert_find_latest_ai_message_run_ids_contract(store, allow_empty_run_id=False)
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_find_latest_ai_message_run_ids_reads_thread_once_and_ignores_empty_run(self, tmp_path):
|
||||||
|
from deerflow.runtime.events.store.jsonl import JsonlRunEventStore
|
||||||
|
|
||||||
|
store = JsonlRunEventStore(base_dir=tmp_path / "jsonl")
|
||||||
|
events = [
|
||||||
|
{
|
||||||
|
"thread_id": "t1",
|
||||||
|
"run_id": "valid-run",
|
||||||
|
"event_type": "llm.ai.response",
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": "target"},
|
||||||
|
"seq": 1,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"thread_id": "t1",
|
||||||
|
"run_id": "",
|
||||||
|
"event_type": "llm.ai.response",
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": "target"},
|
||||||
|
"seq": 2,
|
||||||
|
},
|
||||||
|
]
|
||||||
|
with patch.object(store, "_read_thread_events", return_value=events) as read_thread_events:
|
||||||
|
assert await store.find_latest_ai_message_run_ids("t1", {"target"}, user_id=None) == {"target": "valid-run"}
|
||||||
|
read_thread_events.assert_called_once_with("t1")
|
||||||
|
|
||||||
|
with patch.object(store, "_read_thread_events", side_effect=AssertionError("empty input must not read")) as read_thread_events:
|
||||||
|
assert await store.find_latest_ai_message_run_ids("t1", set()) == {}
|
||||||
|
read_thread_events.assert_not_called()
|
||||||
|
|
||||||
@pytest.mark.anyio
|
@pytest.mark.anyio
|
||||||
async def test_put_if_absent_is_idempotent(self, tmp_path):
|
async def test_put_if_absent_is_idempotent(self, tmp_path):
|
||||||
from deerflow.runtime.events.store.jsonl import JsonlRunEventStore
|
from deerflow.runtime.events.store.jsonl import JsonlRunEventStore
|
||||||
|
|||||||
@ -1355,6 +1355,7 @@ def test_get_thread_history_returns_iso_for_legacy_checkpoint_metadata() -> None
|
|||||||
|
|
||||||
def test_get_thread_history_associates_tool_messages_from_checkpoint_turn() -> None:
|
def test_get_thread_history_associates_tool_messages_from_checkpoint_turn() -> None:
|
||||||
app, _store, checkpointer = _build_thread_app()
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=AsyncMock(return_value={}))
|
||||||
thread_id = "history-tool-run"
|
thread_id = "history-tool-run"
|
||||||
messages = [
|
messages = [
|
||||||
HumanMessage(id="human-1", content="Use a tool", additional_kwargs={"run_id": "run-1"}),
|
HumanMessage(id="human-1", content="Use a tool", additional_kwargs={"run_id": "run-1"}),
|
||||||
@ -1414,11 +1415,14 @@ def test_get_thread_history_fast_path_skips_runs_already_in_checkpoint_metadata(
|
|||||||
"checkpoint-partial",
|
"checkpoint-partial",
|
||||||
messages,
|
messages,
|
||||||
step=1,
|
step=1,
|
||||||
metadata={"run_durations": {"run-migrated": 4}},
|
metadata={
|
||||||
|
"run_durations": {"run-migrated": 4},
|
||||||
|
"run_message_ids": {"ai-1": "run-migrated"},
|
||||||
|
},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
async def list_by_thread(_: str) -> list[SimpleNamespace]:
|
async def list_by_thread(_: str, *, user_id=None, limit: int = 100) -> list[SimpleNamespace]:
|
||||||
return [
|
return [
|
||||||
SimpleNamespace(
|
SimpleNamespace(
|
||||||
run_id="run-pending",
|
run_id="run-pending",
|
||||||
@ -1427,14 +1431,19 @@ def test_get_thread_history_fast_path_skips_runs_already_in_checkpoint_metadata(
|
|||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
list_messages_calls: list[str] = []
|
lookup_calls: list[set[str]] = []
|
||||||
|
|
||||||
async def list_messages(thread: str, *, limit: int) -> list[dict]:
|
async def find_latest_ai_message_run_ids(thread: str, message_ids: set[str], *, user_id=None) -> dict[str, str]:
|
||||||
list_messages_calls.append(thread)
|
assert thread == thread_id
|
||||||
return []
|
assert message_ids == {"ai-2"}
|
||||||
|
lookup_calls.append(message_ids)
|
||||||
|
return {}
|
||||||
|
|
||||||
app.state.run_manager = SimpleNamespace(list_by_thread=list_by_thread)
|
app.state.run_manager = SimpleNamespace(
|
||||||
app.state.run_event_store = SimpleNamespace(list_messages=list_messages)
|
list_by_thread=list_by_thread,
|
||||||
|
reserve_thread_operation=app.state.run_manager.reserve_thread_operation,
|
||||||
|
)
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
with TestClient(app) as client:
|
with TestClient(app) as client:
|
||||||
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
||||||
@ -1443,9 +1452,475 @@ def test_get_thread_history_fast_path_skips_runs_already_in_checkpoint_metadata(
|
|||||||
history_messages = response.json()[0]["values"]["messages"]
|
history_messages = response.json()[0]["values"]["messages"]
|
||||||
assert history_messages[1]["additional_kwargs"]["turn_duration"] == 4
|
assert history_messages[1]["additional_kwargs"]["turn_duration"] == 4
|
||||||
assert history_messages[3]["additional_kwargs"]["turn_duration"] == 6
|
assert history_messages[3]["additional_kwargs"]["turn_duration"] == 6
|
||||||
# The fallback still runs (run-pending was missing), but it is the only
|
# The missing ID is checked once for the response and once after write
|
||||||
# reason it ran — proven by it firing exactly once, not skipped entirely.
|
# admission; the already migrated ID is absent from both lookups.
|
||||||
assert list_messages_calls == [thread_id]
|
assert lookup_calls == [{"ai-2"}, {"ai-2"}]
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_thread_history_backfills_exact_mapping_when_durations_already_exist() -> None:
|
||||||
|
"""Duration metadata alone does not prove exact message attribution.
|
||||||
|
|
||||||
|
A pre-#4949 checkpoint can already carry every run duration while lacking
|
||||||
|
``run_message_ids``. The history read must still consult the event index;
|
||||||
|
otherwise the synthesized human-boundary run becomes permanent.
|
||||||
|
"""
|
||||||
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
thread_id = "history-duration-without-attribution"
|
||||||
|
messages = [
|
||||||
|
HumanMessage(id="human-1", content="Question", additional_kwargs={"run_id": "boundary-run"}),
|
||||||
|
AIMessage(id="ai-1", content="Answer"),
|
||||||
|
]
|
||||||
|
asyncio.run(
|
||||||
|
_write_checkpoint(
|
||||||
|
checkpointer,
|
||||||
|
thread_id,
|
||||||
|
"00000000-0000-6000-8000-000000000010",
|
||||||
|
messages,
|
||||||
|
step=1,
|
||||||
|
metadata={"run_durations": {"boundary-run": 3, "exact-run": 7}},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
lookup_calls: list[set[str]] = []
|
||||||
|
|
||||||
|
async def find_latest_ai_message_run_ids(_: str, message_ids: set[str], *, user_id=None) -> dict[str, str]:
|
||||||
|
lookup_calls.append(message_ids)
|
||||||
|
return {"ai-1": "exact-run"}
|
||||||
|
|
||||||
|
async def list_by_thread(_: str, *, user_id=None, limit: int = 100) -> list[SimpleNamespace]:
|
||||||
|
return []
|
||||||
|
|
||||||
|
app.state.run_manager = SimpleNamespace(
|
||||||
|
list_by_thread=list_by_thread,
|
||||||
|
reserve_thread_operation=app.state.run_manager.reserve_thread_operation,
|
||||||
|
)
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
|
with TestClient(app) as client:
|
||||||
|
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
||||||
|
|
||||||
|
assert response.status_code == 200, response.text
|
||||||
|
ai_message = response.json()[0]["values"]["messages"][1]
|
||||||
|
assert ai_message["run_id"] == "exact-run"
|
||||||
|
assert ai_message["additional_kwargs"]["turn_duration"] == 7
|
||||||
|
assert lookup_calls == [{"ai-1"}, {"ai-1"}]
|
||||||
|
|
||||||
|
latest = asyncio.run(checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}))
|
||||||
|
assert latest is not None
|
||||||
|
assert latest.metadata["run_message_ids"] == {"ai-1": "exact-run"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_thread_history_preserves_boundary_fallback_after_complete_partial_lookup() -> None:
|
||||||
|
"""A complete lookup may legitimately find no event for old messages.
|
||||||
|
|
||||||
|
Pre-event-store checkpoints still rely on the human turn boundary. A
|
||||||
|
partial result therefore corrects the IDs it can prove and preserves that
|
||||||
|
compatibility fallback for IDs that are definitively absent.
|
||||||
|
"""
|
||||||
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
thread_id = "history-partial-exact-attribution"
|
||||||
|
messages = [
|
||||||
|
HumanMessage(id="human-1", content="First", additional_kwargs={"run_id": "boundary-1"}),
|
||||||
|
AIMessage(id="ai-1", content="First answer"),
|
||||||
|
HumanMessage(id="human-2", content="Second", additional_kwargs={"run_id": "boundary-2"}),
|
||||||
|
AIMessage(id="ai-2", content="Second answer"),
|
||||||
|
]
|
||||||
|
asyncio.run(_write_checkpoint(checkpointer, thread_id, "00000000-0000-6000-8000-000000000011", messages, step=1))
|
||||||
|
|
||||||
|
lookup_calls: list[set[str]] = []
|
||||||
|
|
||||||
|
async def find_latest_ai_message_run_ids(_: str, message_ids: set[str], *, user_id=None) -> dict[str, str]:
|
||||||
|
lookup_calls.append(message_ids)
|
||||||
|
if message_ids == {"ai-1", "ai-2"}:
|
||||||
|
return {"ai-1": "exact-1"}
|
||||||
|
assert message_ids == {"ai-2"}
|
||||||
|
return {}
|
||||||
|
|
||||||
|
async def list_by_thread(_: str, *, user_id=None, limit: int = 100) -> list[SimpleNamespace]:
|
||||||
|
return [
|
||||||
|
SimpleNamespace(
|
||||||
|
run_id="exact-1",
|
||||||
|
created_at="2026-07-05T00:00:00+00:00",
|
||||||
|
updated_at="2026-07-05T00:00:07+00:00",
|
||||||
|
),
|
||||||
|
SimpleNamespace(
|
||||||
|
run_id="boundary-2",
|
||||||
|
created_at="2026-07-05T00:00:00+00:00",
|
||||||
|
updated_at="2026-07-05T00:00:03+00:00",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
app.state.run_manager = SimpleNamespace(
|
||||||
|
list_by_thread=list_by_thread,
|
||||||
|
reserve_thread_operation=app.state.run_manager.reserve_thread_operation,
|
||||||
|
)
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
|
with TestClient(app) as client:
|
||||||
|
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
||||||
|
|
||||||
|
assert response.status_code == 200, response.text
|
||||||
|
first_ai, second_ai = [message for message in response.json()[0]["values"]["messages"] if message["type"] == "ai"]
|
||||||
|
assert first_ai["run_id"] == "exact-1"
|
||||||
|
assert first_ai["additional_kwargs"]["turn_duration"] == 7
|
||||||
|
assert second_ai["run_id"] == "boundary-2"
|
||||||
|
assert second_ai["additional_kwargs"]["turn_duration"] == 3
|
||||||
|
|
||||||
|
latest = asyncio.run(checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}))
|
||||||
|
assert latest is not None
|
||||||
|
assert latest.metadata["run_message_ids"] == {"ai-1": "exact-1", "ai-2": "boundary-2"}
|
||||||
|
assert latest.metadata["run_durations"] == {"boundary-2": 3, "exact-1": 7}
|
||||||
|
assert lookup_calls == [{"ai-1", "ai-2"}, {"ai-1", "ai-2"}]
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_thread_history_removes_synthesized_boundary_when_exact_lookup_is_incomplete() -> None:
|
||||||
|
"""Unsafe pagination removes only attribution it cannot prove."""
|
||||||
|
from deerflow.runtime.events.store.base import IncompleteMessageRunLookupError
|
||||||
|
|
||||||
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
thread_id = "history-incomplete-exact-attribution"
|
||||||
|
messages = [
|
||||||
|
HumanMessage(id="human-1", content="Proven question", additional_kwargs={"run_id": "proven-run"}),
|
||||||
|
AIMessage(id="ai-1", content="Proven answer"),
|
||||||
|
HumanMessage(id="human-2", content="Legacy question", additional_kwargs={"run_id": "boundary-run"}),
|
||||||
|
AIMessage(id="ai-2", content="Legacy answer"),
|
||||||
|
]
|
||||||
|
asyncio.run(
|
||||||
|
_write_checkpoint(
|
||||||
|
checkpointer,
|
||||||
|
thread_id,
|
||||||
|
"00000000-0000-6000-8000-000000000012",
|
||||||
|
messages,
|
||||||
|
step=1,
|
||||||
|
metadata={
|
||||||
|
"run_durations": {"proven-run": 4},
|
||||||
|
"run_message_ids": {"ai-1": "proven-run"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=AsyncMock(side_effect=IncompleteMessageRunLookupError("Run event lookup could not form a safe backward cursor")))
|
||||||
|
|
||||||
|
with TestClient(app) as client:
|
||||||
|
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
||||||
|
|
||||||
|
assert response.status_code == 200, response.text
|
||||||
|
first_ai, second_ai = [message for message in response.json()[0]["values"]["messages"] if message["type"] == "ai"]
|
||||||
|
assert first_ai["run_id"] == "proven-run"
|
||||||
|
assert first_ai["additional_kwargs"]["turn_duration"] == 4
|
||||||
|
assert "run_id" not in second_ai
|
||||||
|
assert "turn_duration" not in (second_ai.get("additional_kwargs") or {})
|
||||||
|
assert app.state.run_manager.reservations == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_thread_history_caches_complete_boundary_attribution() -> None:
|
||||||
|
"""A complete audit, including a negative event result, is a one-time scan."""
|
||||||
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
thread_id = "history-sparse-exact-attribution"
|
||||||
|
messages = [
|
||||||
|
HumanMessage(id="human-1", content="Question", additional_kwargs={"run_id": "boundary-run"}),
|
||||||
|
AIMessage(id="ai-1", content="Answer"),
|
||||||
|
]
|
||||||
|
asyncio.run(_write_checkpoint(checkpointer, thread_id, "00000000-0000-6000-8000-000000000013", messages, step=1))
|
||||||
|
|
||||||
|
lookup_calls: list[set[str]] = []
|
||||||
|
|
||||||
|
async def find_latest_ai_message_run_ids(_: str, message_ids: set[str], *, user_id=None) -> dict[str, str]:
|
||||||
|
assert message_ids == {"ai-1"}
|
||||||
|
lookup_calls.append(message_ids)
|
||||||
|
return {}
|
||||||
|
|
||||||
|
async def list_by_thread(_: str, *, user_id=None, limit: int = 100) -> 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",
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
app.state.run_manager = SimpleNamespace(
|
||||||
|
list_by_thread=list_by_thread,
|
||||||
|
reserve_thread_operation=app.state.run_manager.reserve_thread_operation,
|
||||||
|
)
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
|
with TestClient(app) as client:
|
||||||
|
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
||||||
|
second_response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
||||||
|
|
||||||
|
assert response.status_code == 200, response.text
|
||||||
|
ai_message = response.json()[0]["values"]["messages"][1]
|
||||||
|
assert ai_message["run_id"] == "boundary-run"
|
||||||
|
assert ai_message["additional_kwargs"]["turn_duration"] == 3
|
||||||
|
assert second_response.status_code == 200, second_response.text
|
||||||
|
assert lookup_calls == [{"ai-1"}, {"ai-1"}]
|
||||||
|
|
||||||
|
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}
|
||||||
|
assert latest.metadata["run_message_ids"] == {"ai-1": "boundary-run"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_thread_history_revalidates_boundary_fallback_after_reservation() -> None:
|
||||||
|
"""A run may flush its exact event before the metadata task is admitted.
|
||||||
|
|
||||||
|
The foreground lookup can exhaust the event log while the run's journal is
|
||||||
|
still buffered. If the background task acquires its checkpoint reservation
|
||||||
|
only after that run flushes and releases the thread, persisting the earlier
|
||||||
|
human-boundary fallback would make the temporary miss permanent.
|
||||||
|
"""
|
||||||
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
thread_id = "history-fallback-reservation-race"
|
||||||
|
messages = [
|
||||||
|
HumanMessage(
|
||||||
|
id="human-1",
|
||||||
|
content="Question",
|
||||||
|
additional_kwargs={"run_id": "boundary-run"},
|
||||||
|
),
|
||||||
|
AIMessage(id="ai-1", content="Answer"),
|
||||||
|
]
|
||||||
|
asyncio.run(
|
||||||
|
_write_checkpoint(
|
||||||
|
checkpointer,
|
||||||
|
thread_id,
|
||||||
|
"00000000-0000-6000-8000-000000000014",
|
||||||
|
messages,
|
||||||
|
step=1,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
event_visible = False
|
||||||
|
lookup_visibility: list[bool] = []
|
||||||
|
|
||||||
|
async def find_latest_ai_message_run_ids(
|
||||||
|
_: str,
|
||||||
|
message_ids: set[str],
|
||||||
|
*,
|
||||||
|
user_id=None,
|
||||||
|
) -> dict[str, str]:
|
||||||
|
assert message_ids == {"ai-1"}
|
||||||
|
lookup_visibility.append(event_visible)
|
||||||
|
return {"ai-1": "exact-run"} if event_visible else {}
|
||||||
|
|
||||||
|
class RunManager(_ThreadTestRunManager):
|
||||||
|
async def list_by_thread(
|
||||||
|
self,
|
||||||
|
_thread_id: str,
|
||||||
|
*,
|
||||||
|
user_id=None,
|
||||||
|
limit: int = 100,
|
||||||
|
) -> list[SimpleNamespace]:
|
||||||
|
runs = [
|
||||||
|
SimpleNamespace(
|
||||||
|
run_id="boundary-run",
|
||||||
|
created_at="2026-07-05T00:00:00+00:00",
|
||||||
|
updated_at="2026-07-05T00:00:03+00:00",
|
||||||
|
)
|
||||||
|
]
|
||||||
|
if event_visible:
|
||||||
|
runs.append(
|
||||||
|
SimpleNamespace(
|
||||||
|
run_id="exact-run",
|
||||||
|
created_at="2026-07-05T00:00:00+00:00",
|
||||||
|
updated_at="2026-07-05T00:00:07+00:00",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return runs
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def reserve_thread_operation(self, _thread_id: str, **kwargs):
|
||||||
|
nonlocal event_visible
|
||||||
|
self.reservations.append((_thread_id, kwargs))
|
||||||
|
event_visible = True
|
||||||
|
yield
|
||||||
|
|
||||||
|
app.state.run_manager = RunManager()
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
|
with TestClient(app) as client:
|
||||||
|
response = client.post(
|
||||||
|
f"/api/threads/{thread_id}/history",
|
||||||
|
json={"limit": 10},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 200, response.text
|
||||||
|
response_ai = response.json()[0]["values"]["messages"][1]
|
||||||
|
assert response_ai["run_id"] == "boundary-run"
|
||||||
|
assert response_ai["additional_kwargs"]["turn_duration"] == 3
|
||||||
|
assert lookup_visibility == [False, True]
|
||||||
|
|
||||||
|
latest = asyncio.run(checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}))
|
||||||
|
assert latest is not None
|
||||||
|
assert latest.metadata["run_message_ids"] == {"ai-1": "exact-run"}
|
||||||
|
assert latest.metadata["run_durations"]["exact-run"] == 7
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_thread_history_revalidates_exact_attribution_after_reservation() -> None:
|
||||||
|
"""A newer exact event must replace the foreground mapping before persistence."""
|
||||||
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
thread_id = "history-exact-reservation-race"
|
||||||
|
messages = [
|
||||||
|
HumanMessage(
|
||||||
|
id="human-1",
|
||||||
|
content="Question",
|
||||||
|
additional_kwargs={"run_id": "boundary-run"},
|
||||||
|
),
|
||||||
|
AIMessage(id="ai-1", content="Answer"),
|
||||||
|
]
|
||||||
|
asyncio.run(
|
||||||
|
_write_checkpoint(
|
||||||
|
checkpointer,
|
||||||
|
thread_id,
|
||||||
|
"00000000-0000-6000-8000-000000000015",
|
||||||
|
messages,
|
||||||
|
step=1,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
admitted = False
|
||||||
|
lookup_states: list[bool] = []
|
||||||
|
|
||||||
|
async def find_latest_ai_message_run_ids(
|
||||||
|
_: str,
|
||||||
|
message_ids: set[str],
|
||||||
|
*,
|
||||||
|
user_id=None,
|
||||||
|
) -> dict[str, str]:
|
||||||
|
assert message_ids == {"ai-1"}
|
||||||
|
lookup_states.append(admitted)
|
||||||
|
return {"ai-1": "new-run" if admitted else "old-run"}
|
||||||
|
|
||||||
|
class RunManager(_ThreadTestRunManager):
|
||||||
|
async def list_by_thread(
|
||||||
|
self,
|
||||||
|
_thread_id: str,
|
||||||
|
*,
|
||||||
|
user_id=None,
|
||||||
|
limit: int = 100,
|
||||||
|
) -> list[SimpleNamespace]:
|
||||||
|
runs = [
|
||||||
|
SimpleNamespace(
|
||||||
|
run_id="old-run",
|
||||||
|
created_at="2026-07-05T00:00:00+00:00",
|
||||||
|
updated_at="2026-07-05T00:00:04+00:00",
|
||||||
|
)
|
||||||
|
]
|
||||||
|
if admitted:
|
||||||
|
runs.append(
|
||||||
|
SimpleNamespace(
|
||||||
|
run_id="new-run",
|
||||||
|
created_at="2026-07-05T00:00:00+00:00",
|
||||||
|
updated_at="2026-07-05T00:00:08+00:00",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return runs
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def reserve_thread_operation(self, _thread_id: str, **kwargs):
|
||||||
|
nonlocal admitted
|
||||||
|
self.reservations.append((_thread_id, kwargs))
|
||||||
|
admitted = True
|
||||||
|
yield
|
||||||
|
|
||||||
|
app.state.run_manager = RunManager()
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
|
with TestClient(app) as client:
|
||||||
|
response = client.post(
|
||||||
|
f"/api/threads/{thread_id}/history",
|
||||||
|
json={"limit": 10},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 200, response.text
|
||||||
|
response_ai = response.json()[0]["values"]["messages"][1]
|
||||||
|
assert response_ai["run_id"] == "old-run"
|
||||||
|
assert response_ai["additional_kwargs"]["turn_duration"] == 4
|
||||||
|
assert lookup_states == [False, True]
|
||||||
|
|
||||||
|
latest = asyncio.run(checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}))
|
||||||
|
assert latest is not None
|
||||||
|
assert latest.metadata["run_message_ids"] == {"ai-1": "new-run"}
|
||||||
|
assert latest.metadata["run_durations"]["new-run"] == 8
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_thread_history_recomputes_duration_after_reservation() -> None:
|
||||||
|
"""A final run row must replace a stale foreground duration before persistence."""
|
||||||
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
thread_id = "history-duration-reservation-race"
|
||||||
|
messages = [
|
||||||
|
HumanMessage(
|
||||||
|
id="human-1",
|
||||||
|
content="Question",
|
||||||
|
additional_kwargs={"run_id": "run-1"},
|
||||||
|
),
|
||||||
|
AIMessage(id="ai-1", content="Answer"),
|
||||||
|
]
|
||||||
|
asyncio.run(
|
||||||
|
_write_checkpoint(
|
||||||
|
checkpointer,
|
||||||
|
thread_id,
|
||||||
|
"00000000-0000-6000-8000-000000000016",
|
||||||
|
messages,
|
||||||
|
step=1,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
admitted = False
|
||||||
|
lookup_states: list[bool] = []
|
||||||
|
|
||||||
|
async def find_latest_ai_message_run_ids(
|
||||||
|
_: str,
|
||||||
|
message_ids: set[str],
|
||||||
|
*,
|
||||||
|
user_id=None,
|
||||||
|
) -> dict[str, str]:
|
||||||
|
assert message_ids == {"ai-1"}
|
||||||
|
lookup_states.append(admitted)
|
||||||
|
return {"ai-1": "run-1"}
|
||||||
|
|
||||||
|
class RunManager(_ThreadTestRunManager):
|
||||||
|
async def list_by_thread(
|
||||||
|
self,
|
||||||
|
_thread_id: str,
|
||||||
|
*,
|
||||||
|
user_id=None,
|
||||||
|
limit: int = 100,
|
||||||
|
) -> list[SimpleNamespace]:
|
||||||
|
return [
|
||||||
|
SimpleNamespace(
|
||||||
|
run_id="run-1",
|
||||||
|
created_at="2026-07-05T00:00:00+00:00",
|
||||||
|
updated_at=("2026-07-05T00:00:09+00:00" if admitted else "2026-07-05T00:00:03+00:00"),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def reserve_thread_operation(self, _thread_id: str, **kwargs):
|
||||||
|
nonlocal admitted
|
||||||
|
self.reservations.append((_thread_id, kwargs))
|
||||||
|
admitted = True
|
||||||
|
yield
|
||||||
|
|
||||||
|
app.state.run_manager = RunManager()
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
|
with TestClient(app) as client:
|
||||||
|
response = client.post(
|
||||||
|
f"/api/threads/{thread_id}/history",
|
||||||
|
json={"limit": 10},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 200, response.text
|
||||||
|
response_ai = response.json()[0]["values"]["messages"][1]
|
||||||
|
assert response_ai["run_id"] == "run-1"
|
||||||
|
assert response_ai["additional_kwargs"]["turn_duration"] == 3
|
||||||
|
assert lookup_states == [False, True]
|
||||||
|
|
||||||
|
latest = asyncio.run(checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}))
|
||||||
|
assert latest is not None
|
||||||
|
assert latest.metadata["run_message_ids"] == {"ai-1": "run-1"}
|
||||||
|
assert latest.metadata["run_durations"]["run-1"] == 9
|
||||||
|
|
||||||
|
|
||||||
def test_get_thread_history_backfills_legacy_durations_with_exact_event_run_id() -> None:
|
def test_get_thread_history_backfills_legacy_durations_with_exact_event_run_id() -> None:
|
||||||
@ -1458,7 +1933,7 @@ def test_get_thread_history_backfills_legacy_durations_with_exact_event_run_id()
|
|||||||
]
|
]
|
||||||
asyncio.run(_write_checkpoint(checkpointer, thread_id, "00000000-0000-6000-8000-000000000001", messages, step=1))
|
asyncio.run(_write_checkpoint(checkpointer, thread_id, "00000000-0000-6000-8000-000000000001", messages, step=1))
|
||||||
|
|
||||||
async def list_by_thread(_: str) -> list[SimpleNamespace]:
|
async def list_by_thread(_: str, *, user_id=None, limit: int = 100) -> list[SimpleNamespace]:
|
||||||
return [
|
return [
|
||||||
SimpleNamespace(
|
SimpleNamespace(
|
||||||
run_id="boundary-run",
|
run_id="boundary-run",
|
||||||
@ -1472,15 +1947,23 @@ def test_get_thread_history_backfills_legacy_durations_with_exact_event_run_id()
|
|||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
async def list_messages(_: str, *, limit: int) -> list[dict]:
|
list_messages_calls: list[str] = []
|
||||||
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)
|
async def find_latest_ai_message_run_ids(thread: str, message_ids: set[str], *, user_id=None) -> dict[str, str]:
|
||||||
app.state.run_event_store = SimpleNamespace(list_messages=list_messages)
|
assert message_ids == {"ai-1"}
|
||||||
|
list_messages_calls.append(thread)
|
||||||
|
return {"ai-1": "exact-run"}
|
||||||
|
|
||||||
|
reservation_owner = app.state.run_manager
|
||||||
|
app.state.run_manager = SimpleNamespace(
|
||||||
|
list_by_thread=list_by_thread,
|
||||||
|
reserve_thread_operation=reservation_owner.reserve_thread_operation,
|
||||||
|
)
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
with TestClient(app) as client:
|
with TestClient(app) as client:
|
||||||
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
||||||
|
second_response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
||||||
|
|
||||||
assert response.status_code == 200, response.text
|
assert response.status_code == 200, response.text
|
||||||
entry = response.json()[0]
|
entry = response.json()[0]
|
||||||
@ -1489,10 +1972,200 @@ def test_get_thread_history_backfills_legacy_durations_with_exact_event_run_id()
|
|||||||
assert history_messages[1]["additional_kwargs"]["turn_duration"] == 7
|
assert history_messages[1]["additional_kwargs"]["turn_duration"] == 7
|
||||||
assert history_messages[2]["run_id"] == "boundary-run"
|
assert history_messages[2]["run_id"] == "boundary-run"
|
||||||
assert "run_durations" not in entry["metadata"]
|
assert "run_durations" not in entry["metadata"]
|
||||||
|
assert list_messages_calls == [thread_id, thread_id]
|
||||||
|
assert len(reservation_owner.reservations) == 1
|
||||||
|
reserved_thread_id, reservation_kwargs = reservation_owner.reservations[0]
|
||||||
|
assert reserved_thread_id == thread_id
|
||||||
|
assert reservation_kwargs["kind"] is ThreadOperationKind.checkpoint_write
|
||||||
|
assert isinstance(reservation_kwargs["user_id"], str)
|
||||||
|
|
||||||
|
assert second_response.status_code == 200, second_response.text
|
||||||
|
second_history_messages = second_response.json()[0]["values"]["messages"]
|
||||||
|
assert second_history_messages[1]["run_id"] == "exact-run"
|
||||||
|
assert second_history_messages[1]["additional_kwargs"]["turn_duration"] == 7
|
||||||
|
|
||||||
latest = asyncio.run(checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}))
|
latest = asyncio.run(checkpointer.aget_tuple({"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}))
|
||||||
assert latest is not None
|
assert latest is not None
|
||||||
assert latest.metadata["run_durations"] == {"boundary-run": 3, "exact-run": 7}
|
assert latest.metadata["run_durations"] == {"exact-run": 7}
|
||||||
|
assert latest.metadata["run_message_ids"] == {"ai-1": "exact-run"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_thread_history_finds_ai_event_beyond_ten_thousand_newer_events() -> None:
|
||||||
|
"""#4949: no arbitrary page cap may turn an old exact run into a boundary run."""
|
||||||
|
from deerflow.runtime.events.store.memory import MemoryRunEventStore
|
||||||
|
|
||||||
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
thread_id = "legacy-history-run-id-paginated"
|
||||||
|
messages = [
|
||||||
|
HumanMessage(id="human-1", content="Question", additional_kwargs={"run_id": "boundary-run"}),
|
||||||
|
AIMessage(id="ai-1", content="Answer"),
|
||||||
|
]
|
||||||
|
asyncio.run(_write_checkpoint(checkpointer, thread_id, "00000000-0000-6000-8000-000000000002", messages, step=1))
|
||||||
|
|
||||||
|
async def list_by_thread(_: str, *, user_id=None, limit: int = 100) -> 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",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
event_store = MemoryRunEventStore()
|
||||||
|
events = [
|
||||||
|
{
|
||||||
|
"thread_id": thread_id,
|
||||||
|
"run_id": "exact-run",
|
||||||
|
"event_type": "llm.ai.response",
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": "ai-1"},
|
||||||
|
},
|
||||||
|
*[
|
||||||
|
{
|
||||||
|
"thread_id": thread_id,
|
||||||
|
"run_id": "noise-run",
|
||||||
|
"event_type": "llm.ai.response",
|
||||||
|
"category": "message",
|
||||||
|
"content": {"type": "ai", "id": f"noise-{index}"},
|
||||||
|
}
|
||||||
|
for index in range(10_000)
|
||||||
|
],
|
||||||
|
]
|
||||||
|
asyncio.run(event_store.put_batch(events))
|
||||||
|
event_store.find_latest_ai_message_run_ids = AsyncMock(wraps=event_store.find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
|
app.state.run_manager = SimpleNamespace(
|
||||||
|
list_by_thread=list_by_thread,
|
||||||
|
reserve_thread_operation=app.state.run_manager.reserve_thread_operation,
|
||||||
|
)
|
||||||
|
app.state.run_event_store = event_store
|
||||||
|
|
||||||
|
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 history_messages[1]["run_id"] == "exact-run"
|
||||||
|
assert history_messages[1]["additional_kwargs"]["turn_duration"] == 7
|
||||||
|
assert event_store.find_latest_ai_message_run_ids.await_count == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_thread_history_sizes_initial_run_page_to_required_attributions() -> None:
|
||||||
|
"""A long thread should batch-hydrate its common migration path."""
|
||||||
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
thread_id = "legacy-history-run-page-sizing"
|
||||||
|
run_count = 101
|
||||||
|
messages = []
|
||||||
|
runs = []
|
||||||
|
message_run_ids: dict[str, str] = {}
|
||||||
|
for index in range(run_count):
|
||||||
|
boundary_run_id = f"boundary-{index}"
|
||||||
|
exact_run_id = f"exact-{index}"
|
||||||
|
message_id = f"ai-{index}"
|
||||||
|
messages.extend(
|
||||||
|
[
|
||||||
|
HumanMessage(id=f"human-{index}", content=f"Question {index}", additional_kwargs={"run_id": boundary_run_id}),
|
||||||
|
AIMessage(id=message_id, content=f"Answer {index}"),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
message_run_ids[message_id] = exact_run_id
|
||||||
|
runs.append(
|
||||||
|
SimpleNamespace(
|
||||||
|
run_id=exact_run_id,
|
||||||
|
created_at="2026-07-05T00:00:00+00:00",
|
||||||
|
updated_at="2026-07-05T00:00:05+00:00",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
asyncio.run(_write_checkpoint(checkpointer, thread_id, "00000000-0000-6000-8000-000000000020", messages, step=1))
|
||||||
|
|
||||||
|
list_limits: list[int] = []
|
||||||
|
|
||||||
|
async def list_by_thread(_: str, *, user_id=None, limit: int = 100) -> list[SimpleNamespace]:
|
||||||
|
list_limits.append(limit)
|
||||||
|
return runs[:limit]
|
||||||
|
|
||||||
|
async def get(run_id: str, *, user_id=None) -> SimpleNamespace | None:
|
||||||
|
return next((run for run in runs if run.run_id == run_id), None)
|
||||||
|
|
||||||
|
get_mock = AsyncMock(side_effect=get)
|
||||||
|
|
||||||
|
async def find_latest_ai_message_run_ids(_: str, message_ids: set[str], *, user_id=None) -> dict[str, str]:
|
||||||
|
assert message_ids == set(message_run_ids)
|
||||||
|
return message_run_ids
|
||||||
|
|
||||||
|
app.state.run_manager = SimpleNamespace(
|
||||||
|
list_by_thread=list_by_thread,
|
||||||
|
get=get_mock,
|
||||||
|
reserve_thread_operation=app.state.run_manager.reserve_thread_operation,
|
||||||
|
)
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
|
with TestClient(app) as client:
|
||||||
|
response = client.post(f"/api/threads/{thread_id}/history", json={"limit": 10})
|
||||||
|
|
||||||
|
assert response.status_code == 200, response.text
|
||||||
|
ai_messages = [message for message in response.json()[0]["values"]["messages"] if message["type"] == "ai"]
|
||||||
|
assert len(ai_messages) == run_count
|
||||||
|
assert ai_messages[-1]["additional_kwargs"]["turn_duration"] == 5
|
||||||
|
assert list_limits == [run_count, run_count]
|
||||||
|
get_mock.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_thread_history_fetches_exact_run_older_than_default_run_page() -> None:
|
||||||
|
"""The event index may resolve a run outside RunManager's newest-100 page."""
|
||||||
|
app, _store, checkpointer = _build_thread_app()
|
||||||
|
thread_id = "legacy-history-old-exact-run"
|
||||||
|
messages = [
|
||||||
|
HumanMessage(id="human-1", content="Question", additional_kwargs={"run_id": "boundary-run"}),
|
||||||
|
AIMessage(id="ai-1", content="Answer"),
|
||||||
|
]
|
||||||
|
asyncio.run(_write_checkpoint(checkpointer, thread_id, "00000000-0000-6000-8000-000000000003", messages, step=1))
|
||||||
|
|
||||||
|
boundary_run = SimpleNamespace(
|
||||||
|
run_id="boundary-run",
|
||||||
|
created_at="2026-07-05T00:00:00+00:00",
|
||||||
|
updated_at="2026-07-05T00:00:03+00:00",
|
||||||
|
)
|
||||||
|
exact_run = SimpleNamespace(
|
||||||
|
run_id="old-exact-run",
|
||||||
|
created_at="2026-06-01T00:00:00+00:00",
|
||||||
|
updated_at="2026-06-01T00:00:09+00:00",
|
||||||
|
)
|
||||||
|
get_calls: list[str] = []
|
||||||
|
|
||||||
|
async def list_by_thread(_: str, *, user_id=None, limit: int = 100) -> list[SimpleNamespace]:
|
||||||
|
assert limit == 100
|
||||||
|
return [boundary_run]
|
||||||
|
|
||||||
|
async def get(run_id: str, *, user_id=None) -> SimpleNamespace | None:
|
||||||
|
get_calls.append(run_id)
|
||||||
|
return exact_run if run_id == exact_run.run_id else None
|
||||||
|
|
||||||
|
async def find_latest_ai_message_run_ids(_: str, message_ids: set[str], *, user_id=None) -> dict[str, str]:
|
||||||
|
assert message_ids == {"ai-1"}
|
||||||
|
return {"ai-1": exact_run.run_id}
|
||||||
|
|
||||||
|
app.state.run_manager = SimpleNamespace(
|
||||||
|
list_by_thread=list_by_thread,
|
||||||
|
get=get,
|
||||||
|
reserve_thread_operation=app.state.run_manager.reserve_thread_operation,
|
||||||
|
)
|
||||||
|
app.state.run_event_store = SimpleNamespace(find_latest_ai_message_run_ids=find_latest_ai_message_run_ids)
|
||||||
|
|
||||||
|
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 history_messages[1]["run_id"] == exact_run.run_id
|
||||||
|
assert history_messages[1]["additional_kwargs"]["turn_duration"] == 9
|
||||||
|
assert get_calls == [exact_run.run_id, exact_run.run_id]
|
||||||
|
|
||||||
|
|
||||||
def test_get_thread_history_injects_turn_duration_once_per_run() -> None:
|
def test_get_thread_history_injects_turn_duration_once_per_run() -> None:
|
||||||
@ -1529,8 +2202,9 @@ def test_get_thread_history_injects_turn_duration_once_per_run() -> None:
|
|||||||
|
|
||||||
run_manager = AsyncMock()
|
run_manager = AsyncMock()
|
||||||
run_manager.list_by_thread = AsyncMock(return_value=[_run("run-1", 5), _run("run-2", 9)])
|
run_manager.list_by_thread = AsyncMock(return_value=[_run("run-1", 5), _run("run-2", 9)])
|
||||||
|
run_manager.reserve_thread_operation = _ThreadTestRunManager().reserve_thread_operation
|
||||||
event_store = MagicMock()
|
event_store = MagicMock()
|
||||||
event_store.list_messages = AsyncMock(return_value=[])
|
event_store.find_latest_ai_message_run_ids = AsyncMock(return_value={})
|
||||||
app.state.run_manager = run_manager
|
app.state.run_manager = run_manager
|
||||||
app.state.run_event_store = event_store
|
app.state.run_event_store = event_store
|
||||||
|
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user