mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-09 13:39:26 +00:00
fix(journal): dedup llm.ai.response persistence on re-fired on_llm_end (#5187)
* fix(journal): dedup llm.ai.response persistence on re-fired on_llm_end LangChain may deliver on_llm_end more than once for the same run_id. RunJournal already dedups token accounting and the run summary (_record_message_summary) on that premise via _counted_message_llm_run_ids, but the durable llm.ai.response self._put() call was left unguarded. The event store is append-only and count_messages/list_messages read raw rows without read-time dedup, so a replayed callback persists a second llm.ai.response row for one logical response while the run's own message_count counts it once. This inflates count_messages, duplicates a message in list_messages pagination, and leaves the durable feed inconsistent with the run summary. Gate the persistence + summary block by the existing per-run_id guard so a replayed callback is a no-op, keeping the durable message feed and the run summary in agreement. Distinct run_ids are unaffected. Adds regression tests: a re-fired callback for one run_id persists exactly one row (red on main), and distinct run_ids each still persist a message. * fix(journal): preserve canonical response on late usage * fix(journal): preserve late usage while deduplicating responses * fix(journal): keep first callback response canonical * fix(journal): snapshot canonical response summaries --------- Co-authored-by: CorgiBoyG <CorgiBoyG@users.noreply.github.com>
This commit is contained in:
parent
ab9c1719ee
commit
a4ff4b0b3b
@ -61,6 +61,25 @@ fetch-and-decode of every message row's tool outputs on long threads.
|
||||
client input, because a welded-in seq goes stale when a fork re-seeds the feed
|
||||
(#4380).
|
||||
|
||||
**LLM response callback coalescing** (`runtime/journal.py`): a provider may fire
|
||||
`on_llm_end` twice for one LangChain run id, first without usage (or with all token
|
||||
counts zero) and immediately again with usage populated. The first callback's generation
|
||||
set is always canonical: `RunJournal` stages only its response events and immutable
|
||||
message-summary fields while retaining the first caller, and applies that callback's
|
||||
fallback state and tool-call bookkeeping immediately; those effects remain canonical.
|
||||
It must not retain provider-owned message objects because a provider may mutate and
|
||||
reuse the same response for the usage replay. Usage metadata is deep-snapshotted,
|
||||
including nested token-detail mappings, before it enters a staged or buffered event.
|
||||
An adjacent same-id positive-usage replay may enrich only each corresponding staged
|
||||
event's metadata/content usage fields. Replay
|
||||
generation-count differences never add, remove, or replace canonical messages. The next
|
||||
unrelated event, an effective buffer size (committed plus pending events) reaching the
|
||||
flush threshold, or an explicit flush commits the staged unit and updates the message
|
||||
summary. Once that ordering boundary is crossed, a late usage replay can still update the
|
||||
authoritative run token summary, but it cannot mutate the append-only message event,
|
||||
caller attribution, fallback state, or tool-call bookkeeping. Closed journals return
|
||||
from `on_llm_end` before inspecting the response or touching any run state.
|
||||
|
||||
**Run delivery receipts** (`runtime/journal.py` + `runs/worker.py`):
|
||||
`RunJournal` records each non-empty artifact update once per tool `Command` for
|
||||
the terminal `run.delivery` event. When a command contains multiple messages, a
|
||||
|
||||
@ -22,6 +22,8 @@ import logging
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
from uuid import UUID
|
||||
@ -53,6 +55,14 @@ _LEGACY_SUMMARY_MESSAGE_NAME = "summary"
|
||||
_PERSISTED_HIDDEN_HUMAN_INPUT_RESPONSE_SOURCES = frozenset({"ask_clarification", "sandbox_network"})
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PendingLlmResponse:
|
||||
llm_run_id: str
|
||||
events: list[dict]
|
||||
message_count: int
|
||||
last_ai_message: str | None
|
||||
|
||||
|
||||
def _should_persist_human_input_message(message: BaseMessage) -> bool:
|
||||
if not isinstance(message, HumanMessage):
|
||||
return False
|
||||
@ -243,6 +253,7 @@ class RunJournal(BaseCallbackHandler):
|
||||
|
||||
# Write buffer
|
||||
self._buffer: list[dict] = []
|
||||
self._pending_llm_response: _PendingLlmResponse | None = None
|
||||
self._pending_flush_tasks: set[asyncio.Task[None]] = set()
|
||||
self._pending_progress_task: asyncio.Task[None] | None = None
|
||||
self._pending_progress_delayed = False
|
||||
@ -267,6 +278,7 @@ class RunJournal(BaseCallbackHandler):
|
||||
self._counted_llm_run_ids: set[str] = set()
|
||||
self._counted_external_source_ids: set[str] = set()
|
||||
self._counted_message_llm_run_ids: set[str] = set()
|
||||
self._llm_response_callers: dict[str, str] = {}
|
||||
self._memory_context_recorded = False
|
||||
self._tool_promotion_claim_lock = threading.Lock()
|
||||
self._claimed_tool_promotions: set[str] = set()
|
||||
@ -306,6 +318,14 @@ class RunJournal(BaseCallbackHandler):
|
||||
"""Extract displayable text from a message's mixed content shape."""
|
||||
return message_to_text(message, text_attribute_fallback=True)
|
||||
|
||||
def _message_summary_text(self, message: BaseMessage, *, caller: str | None = None) -> str | None:
|
||||
"""Return the bounded user-facing AI summary text for one message."""
|
||||
is_ai_message = isinstance(message, AIMessage) or getattr(message, "type", None) == "ai"
|
||||
if not is_ai_message or (caller is not None and caller != "lead_agent"):
|
||||
return None
|
||||
text = self._message_text(message).strip()
|
||||
return text[:2000] if text else None
|
||||
|
||||
def _record_message_summary(self, message: BaseMessage, *, caller: str | None = None) -> None:
|
||||
"""Update run-level convenience fields for persisted run rows."""
|
||||
self._msg_count += 1
|
||||
@ -313,11 +333,9 @@ class RunJournal(BaseCallbackHandler):
|
||||
# ``last_ai_message`` should represent the lead agent's user-facing
|
||||
# answer. Middleware/subagent model calls and empty tool-call-only
|
||||
# AI messages must not overwrite the last useful assistant text.
|
||||
is_ai_message = isinstance(message, AIMessage) or getattr(message, "type", None) == "ai"
|
||||
if is_ai_message and (caller is None or caller == "lead_agent"):
|
||||
text = self._message_text(message).strip()
|
||||
if text:
|
||||
self._last_ai_msg = text[:2000]
|
||||
summary_text = self._message_summary_text(message, caller=caller)
|
||||
if summary_text is not None:
|
||||
self._last_ai_msg = summary_text
|
||||
|
||||
def on_chain_start(
|
||||
self,
|
||||
@ -433,7 +451,16 @@ class RunJournal(BaseCallbackHandler):
|
||||
tags: list[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
|
||||
messages: list[AnyMessage] = []
|
||||
response_events: list[dict] = []
|
||||
should_schedule_progress = False
|
||||
rid = str(run_id)
|
||||
callback_caller = self._identify_caller(tags)
|
||||
is_canonical_callback = rid not in self._counted_message_llm_run_ids
|
||||
caller = self._llm_response_callers.get(rid, callback_caller)
|
||||
logger.debug("on_llm_end %s: tags=%s", run_id, tags)
|
||||
for generation in response.generations:
|
||||
for gen in generation:
|
||||
@ -443,19 +470,20 @@ class RunJournal(BaseCallbackHandler):
|
||||
logger.warning(f"on_llm_end {run_id}: generation has no message attribute: {gen}")
|
||||
|
||||
for message in messages:
|
||||
caller = self._identify_caller(tags)
|
||||
self._remember_current_run_tool_calls(message, caller=caller)
|
||||
if is_canonical_callback:
|
||||
self._remember_current_run_tool_calls(message, caller=caller)
|
||||
|
||||
# Latency
|
||||
rid = str(run_id)
|
||||
start = self._llm_start_times.pop(rid, None)
|
||||
latency_ms = int((time.monotonic() - start) * 1000) if start else None
|
||||
|
||||
# Token usage from message
|
||||
usage = getattr(message, "usage_metadata", None)
|
||||
usage_dict = dict(usage) if usage else {}
|
||||
# Providers may mutate and reuse the same response object after the
|
||||
# callback returns, including nested token-detail mappings.
|
||||
usage_dict = deepcopy(dict(usage)) if usage else {}
|
||||
additional_kwargs = getattr(message, "additional_kwargs", None) or {}
|
||||
if isinstance(additional_kwargs, dict) and additional_kwargs.get("deerflow_error_fallback"):
|
||||
if is_canonical_callback and isinstance(additional_kwargs, dict) and additional_kwargs.get("deerflow_error_fallback"):
|
||||
self._had_llm_error_fallback = True
|
||||
detail = additional_kwargs.get("error_detail")
|
||||
reason = additional_kwargs.get("error_reason")
|
||||
@ -475,20 +503,19 @@ class RunJournal(BaseCallbackHandler):
|
||||
call_index = self._llm_call_index
|
||||
self._seen_llm_starts.add(rid)
|
||||
|
||||
# Message event: checkpoint-aligned llm.ai.response payload.
|
||||
self._put(
|
||||
event_type=LLM_AI_RESPONSE_EVENT.event_type,
|
||||
category=LLM_AI_RESPONSE_EVENT.category,
|
||||
content=message.model_dump(),
|
||||
metadata={
|
||||
"caller": caller,
|
||||
"usage": usage_dict,
|
||||
"latency_ms": latency_ms,
|
||||
"llm_call_index": call_index,
|
||||
},
|
||||
response_events.append(
|
||||
self._make_event(
|
||||
event_type=LLM_AI_RESPONSE_EVENT.event_type,
|
||||
category=LLM_AI_RESPONSE_EVENT.category,
|
||||
content=message.model_dump(),
|
||||
metadata={
|
||||
"caller": caller,
|
||||
"usage": usage_dict,
|
||||
"latency_ms": latency_ms,
|
||||
"llm_call_index": call_index,
|
||||
},
|
||||
)
|
||||
)
|
||||
if rid not in self._counted_message_llm_run_ids:
|
||||
self._record_message_summary(message, caller=caller)
|
||||
|
||||
# Token accumulation (dedup by langchain run_id to avoid double-counting
|
||||
# when the callback fires more than once for the same response)
|
||||
@ -519,10 +546,18 @@ class RunJournal(BaseCallbackHandler):
|
||||
per_call_model = response_metadata.get("model_name") or response_metadata.get("model")
|
||||
self._record_model_usage(per_call_model, input_tk, output_tk, total_tk, self._extract_cache_read(usage_dict))
|
||||
|
||||
self._schedule_progress_flush()
|
||||
should_schedule_progress = True
|
||||
|
||||
if messages:
|
||||
self._counted_message_llm_run_ids.add(str(run_id))
|
||||
self._queue_llm_response_events(
|
||||
str(run_id),
|
||||
response_events,
|
||||
messages,
|
||||
caller=caller,
|
||||
)
|
||||
|
||||
if should_schedule_progress:
|
||||
self._schedule_progress_flush()
|
||||
|
||||
def on_llm_error(self, error: BaseException, *, run_id: UUID, **kwargs: Any) -> None:
|
||||
self._llm_start_times.pop(str(run_id), None)
|
||||
@ -661,20 +696,123 @@ class RunJournal(BaseCallbackHandler):
|
||||
if self._should_reconcile_tool_message(message):
|
||||
self._persist_tool_result_message(message)
|
||||
|
||||
def _make_event(self, *, event_type: str, category: str, content: str | dict = "", metadata: dict | None = None) -> dict:
|
||||
return {
|
||||
"thread_id": self.thread_id,
|
||||
"run_id": self.run_id,
|
||||
"event_type": event_type,
|
||||
"category": category,
|
||||
"content": content,
|
||||
"metadata": metadata or {},
|
||||
"created_at": datetime.now(UTC).isoformat(),
|
||||
}
|
||||
|
||||
def _commit_pending_llm_response(self) -> None:
|
||||
pending = self._pending_llm_response
|
||||
if pending is None:
|
||||
return
|
||||
self._pending_llm_response = None
|
||||
self._buffer.extend(pending.events)
|
||||
self._msg_count += pending.message_count
|
||||
if pending.last_ai_message is not None:
|
||||
self._last_ai_msg = pending.last_ai_message
|
||||
|
||||
def _snapshot_message_summary(self, messages: Sequence[AnyMessage], *, caller: str) -> tuple[int, str | None]:
|
||||
"""Freeze summary fields before a provider can mutate replayed messages."""
|
||||
last_ai_message: str | None = None
|
||||
for message in messages:
|
||||
summary_text = self._message_summary_text(message, caller=caller)
|
||||
if summary_text is not None:
|
||||
last_ai_message = summary_text
|
||||
return len(messages), last_ai_message
|
||||
|
||||
@staticmethod
|
||||
def _has_positive_usage(events: list[dict]) -> bool:
|
||||
for event in events:
|
||||
usage = event["metadata"].get("usage")
|
||||
if not isinstance(usage, Mapping):
|
||||
continue
|
||||
for key in ("input_tokens", "output_tokens", "total_tokens"):
|
||||
try:
|
||||
if int(usage.get(key) or 0) > 0:
|
||||
return True
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _merge_response_event_usage(canonical_events: list[dict], replay_events: list[dict]) -> None:
|
||||
"""Enrich canonical generation events with only replayed usage fields."""
|
||||
for canonical, replay in zip(canonical_events, replay_events, strict=False):
|
||||
replay_metadata = replay.get("metadata")
|
||||
if isinstance(replay_metadata, Mapping):
|
||||
replay_usage = replay_metadata.get("usage")
|
||||
if isinstance(replay_usage, Mapping):
|
||||
canonical["metadata"]["usage"] = deepcopy(dict(replay_usage))
|
||||
|
||||
canonical_content = canonical.get("content")
|
||||
replay_content = replay.get("content")
|
||||
if isinstance(canonical_content, dict) and isinstance(replay_content, Mapping) and "usage_metadata" in replay_content:
|
||||
replay_content_usage = replay_content.get("usage_metadata")
|
||||
canonical_content["usage_metadata"] = deepcopy(dict(replay_content_usage)) if isinstance(replay_content_usage, Mapping) else replay_content_usage
|
||||
|
||||
def _flush_if_threshold_reached(self) -> None:
|
||||
pending_count = len(self._pending_llm_response.events) if self._pending_llm_response is not None else 0
|
||||
if len(self._buffer) + pending_count >= self._flush_threshold:
|
||||
self._flush_sync()
|
||||
|
||||
def _queue_llm_response_events(
|
||||
self,
|
||||
llm_run_id: str,
|
||||
events: list[dict],
|
||||
messages: list[AnyMessage],
|
||||
*,
|
||||
caller: str,
|
||||
) -> None:
|
||||
"""Queue one logical response and merge usage into its canonical callback."""
|
||||
if self._closed:
|
||||
return
|
||||
|
||||
has_usage = self._has_positive_usage(events)
|
||||
pending = self._pending_llm_response
|
||||
if pending is not None and pending.llm_run_id == llm_run_id:
|
||||
if has_usage:
|
||||
# The first callback's generation set, immutable summary,
|
||||
# caller, and non-usage payload are canonical. A provider's
|
||||
# immediate replay may enrich only corresponding usage fields.
|
||||
self._merge_response_event_usage(pending.events, events)
|
||||
self._commit_pending_llm_response()
|
||||
self._flush_if_threshold_reached()
|
||||
return
|
||||
if llm_run_id in self._counted_message_llm_run_ids:
|
||||
return
|
||||
|
||||
# A different event is the ordering boundary for an earlier no-usage
|
||||
# callback. Commit it before accepting this response.
|
||||
self._commit_pending_llm_response()
|
||||
self._flush_if_threshold_reached()
|
||||
|
||||
message_count, last_ai_message = self._snapshot_message_summary(messages, caller=caller)
|
||||
pending_response = _PendingLlmResponse(
|
||||
llm_run_id=llm_run_id,
|
||||
events=events,
|
||||
message_count=message_count,
|
||||
last_ai_message=last_ai_message,
|
||||
)
|
||||
self._counted_message_llm_run_ids.add(llm_run_id)
|
||||
self._llm_response_callers[llm_run_id] = caller
|
||||
self._pending_llm_response = pending_response
|
||||
if has_usage:
|
||||
self._commit_pending_llm_response()
|
||||
self._flush_if_threshold_reached()
|
||||
# Some providers immediately re-fire on_llm_end with usage filled in.
|
||||
# Defer an incomplete copy until the next event or flush.
|
||||
|
||||
def _put(self, *, event_type: str, category: str, content: str | dict = "", metadata: dict | None = None) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
self._buffer.append(
|
||||
{
|
||||
"thread_id": self.thread_id,
|
||||
"run_id": self.run_id,
|
||||
"event_type": event_type,
|
||||
"category": category,
|
||||
"content": content,
|
||||
"metadata": metadata or {},
|
||||
"created_at": datetime.now(UTC).isoformat(),
|
||||
}
|
||||
)
|
||||
self._commit_pending_llm_response()
|
||||
self._buffer.append(self._make_event(event_type=event_type, category=category, content=content, metadata=metadata))
|
||||
if len(self._buffer) >= self._flush_threshold:
|
||||
self._flush_sync()
|
||||
|
||||
@ -686,6 +824,7 @@ class RunJournal(BaseCallbackHandler):
|
||||
stay in the buffer and are flushed later by the async ``flush()``
|
||||
call in the worker's ``finally`` block.
|
||||
"""
|
||||
self._commit_pending_llm_response()
|
||||
if not self._buffer:
|
||||
return
|
||||
# Skip if a flush is already in flight — avoids concurrent writes
|
||||
@ -929,6 +1068,7 @@ class RunJournal(BaseCallbackHandler):
|
||||
"""Force flush remaining buffer. Called in worker's finally block."""
|
||||
if self._closed:
|
||||
return
|
||||
self._commit_pending_llm_response()
|
||||
if self._pending_flush_tasks:
|
||||
await asyncio.gather(*tuple(self._pending_flush_tasks), return_exceptions=True)
|
||||
while self._pending_progress_task is not None:
|
||||
@ -968,6 +1108,7 @@ class RunJournal(BaseCallbackHandler):
|
||||
self._store = None
|
||||
self._progress_reporter = None
|
||||
self._buffer.clear()
|
||||
self._pending_llm_response = None
|
||||
self._pending_flush_tasks.clear()
|
||||
self._pending_progress_task = None
|
||||
self._pending_progress_delayed = False
|
||||
@ -976,6 +1117,7 @@ class RunJournal(BaseCallbackHandler):
|
||||
self._counted_llm_run_ids.clear()
|
||||
self._counted_external_source_ids.clear()
|
||||
self._counted_message_llm_run_ids.clear()
|
||||
self._llm_response_callers.clear()
|
||||
self._llm_start_times.clear()
|
||||
self._seen_llm_starts.clear()
|
||||
self._current_run_tool_call_names.clear()
|
||||
|
||||
@ -11,7 +11,8 @@ from unittest.mock import MagicMock
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import HumanMessage
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
from langchain_core.outputs import ChatGeneration, LLMResult
|
||||
|
||||
from deerflow.runtime.events.store.memory import MemoryRunEventStore
|
||||
from deerflow.runtime.journal import RunJournal
|
||||
@ -68,6 +69,24 @@ async def test_close_flushes_and_detaches_runtime_dependencies():
|
||||
assert reporter_ref() is None
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_closed_on_llm_end_returns_before_touching_response_or_state():
|
||||
store = MemoryRunEventStore()
|
||||
journal = RunJournal("r-closed-callback", "t-closed-callback", store)
|
||||
await journal.close()
|
||||
completion_before = journal.get_completion_data()
|
||||
|
||||
# A plain object has no generations attribute, so this also pins the
|
||||
# early return ahead of response inspection.
|
||||
journal.on_llm_end(object(), run_id=uuid4(), tags=["lead_agent"])
|
||||
|
||||
assert journal.get_completion_data() == completion_before
|
||||
assert journal._pending_llm_response is None
|
||||
assert journal._buffer == []
|
||||
assert journal._counted_message_llm_run_ids == set()
|
||||
assert journal._counted_llm_run_ids == set()
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_close_preserves_buffer_and_dependencies_when_flush_fails():
|
||||
class FailOnceRunEventStore(MemoryRunEventStore):
|
||||
@ -101,6 +120,73 @@ async def test_close_preserves_buffer_and_dependencies_when_flush_fails():
|
||||
assert [event["event_type"] for event in events] == ["middleware:test"]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_close_retries_pending_no_usage_response_without_duplication():
|
||||
class FailOnceRunEventStore(MemoryRunEventStore):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.put_batch_calls = 0
|
||||
|
||||
async def put_batch(self, events):
|
||||
self.put_batch_calls += 1
|
||||
if self.put_batch_calls == 1:
|
||||
raise RuntimeError("transient store failure")
|
||||
return await super().put_batch(events)
|
||||
|
||||
async def progress_reporter(snapshot):
|
||||
del snapshot
|
||||
|
||||
store = FailOnceRunEventStore()
|
||||
journal = RunJournal(
|
||||
"r-close-pending-retry",
|
||||
"t-close-pending-retry",
|
||||
store,
|
||||
flush_threshold=100,
|
||||
progress_reporter=progress_reporter,
|
||||
)
|
||||
journal.record_middleware("before", name="test", hook="after", action="record", changes={})
|
||||
journal.on_llm_end(
|
||||
_make_llm_response("Canonical without usage"),
|
||||
run_id=uuid4(),
|
||||
parent_run_id=None,
|
||||
tags=["lead_agent"],
|
||||
)
|
||||
|
||||
assert journal._pending_llm_response is not None
|
||||
assert journal.get_completion_data()["message_count"] == 0
|
||||
|
||||
with pytest.raises(RuntimeError, match="transient store failure"):
|
||||
await journal.close()
|
||||
|
||||
assert journal._closed is False
|
||||
assert journal._store is store
|
||||
assert journal._progress_reporter is progress_reporter
|
||||
assert journal._pending_llm_response is None
|
||||
assert [event["event_type"] for event in journal._buffer] == [
|
||||
"middleware:before",
|
||||
"llm.ai.response",
|
||||
]
|
||||
assert journal.get_completion_data()["message_count"] == 1
|
||||
assert journal.get_completion_data()["last_ai_message"] == "Canonical without usage"
|
||||
|
||||
await journal.close()
|
||||
|
||||
events = await store.list_events("t-close-pending-retry", "r-close-pending-retry")
|
||||
assert [event["event_type"] for event in events] == [
|
||||
"middleware:before",
|
||||
"llm.ai.response",
|
||||
]
|
||||
responses = [event for event in events if event["event_type"] == "llm.ai.response"]
|
||||
assert len(responses) == 1
|
||||
assert responses[0]["content"]["content"] == "Canonical without usage"
|
||||
assert responses[0]["content"]["usage_metadata"] is None
|
||||
assert responses[0]["metadata"]["usage"] == {}
|
||||
assert journal.get_completion_data()["message_count"] == 1
|
||||
assert journal._closed is True
|
||||
assert journal._store is None
|
||||
assert journal._progress_reporter is None
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_close_without_flush_discards_buffer_and_detaches_runtime_dependencies():
|
||||
class TrackingRunEventStore(MemoryRunEventStore):
|
||||
@ -197,6 +283,12 @@ def _make_llm_response(content="Hello", usage=None, tool_calls=None, additional_
|
||||
return response
|
||||
|
||||
|
||||
def _combine_llm_responses(*responses):
|
||||
response = MagicMock()
|
||||
response.generations = [generation for item in responses for generation in item.generations]
|
||||
return response
|
||||
|
||||
|
||||
class TestLlmCallbacks:
|
||||
@pytest.mark.anyio
|
||||
async def test_on_chat_model_start_persists_original_user_input_without_mutating_model_message(self, journal_setup):
|
||||
@ -568,14 +660,28 @@ class TestBufferFlush:
|
||||
j, store = journal_setup
|
||||
j._flush_threshold = 2
|
||||
# Each on_llm_end emits 1 event
|
||||
j.on_llm_end(_make_llm_response("A"), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
usage = {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}
|
||||
j.on_llm_end(_make_llm_response("A", usage=usage), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
assert len(j._buffer) == 1
|
||||
j.on_llm_end(_make_llm_response("B"), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
j.on_llm_end(_make_llm_response("B", usage=usage), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
# At threshold the buffer should have been flushed asynchronously
|
||||
await asyncio.sleep(0.1)
|
||||
events = await store.list_events("t1", "r1")
|
||||
assert len(events) >= 2
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_pending_response_counts_toward_flush_threshold(self, journal_setup):
|
||||
j, store = journal_setup
|
||||
j._flush_threshold = 2
|
||||
j.record_middleware("before", name="BeforeMiddleware", hook="after_model", action="record", changes={})
|
||||
|
||||
j.on_llm_end(_make_llm_response("Pending"), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
assert j._pending_llm_response is None
|
||||
events = await store.list_events("t1", "r1")
|
||||
assert [event["event_type"] for event in events] == ["middleware:before", "llm.ai.response"]
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_events_retained_when_no_loop(self, journal_setup):
|
||||
"""Events buffered in a sync (no-loop) context should survive
|
||||
@ -611,11 +717,12 @@ class TestFeedGeneration:
|
||||
"""
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_buffering_alone_does_not_advance_it(self, journal_setup):
|
||||
async def test_pending_response_alone_does_not_advance_it(self, journal_setup):
|
||||
j, _store = journal_setup
|
||||
j.on_llm_end(_make_llm_response("A"), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
|
||||
assert len(j._buffer) == 1
|
||||
assert j._buffer == []
|
||||
assert j._pending_llm_response is not None
|
||||
assert j.feed_generation == 0
|
||||
|
||||
@pytest.mark.anyio
|
||||
@ -623,7 +730,8 @@ class TestFeedGeneration:
|
||||
j, _store = journal_setup
|
||||
j._flush_threshold = 1
|
||||
|
||||
j.on_llm_end(_make_llm_response("A"), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
usage = {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}
|
||||
j.on_llm_end(_make_llm_response("A", usage=usage), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
assert j.feed_generation == 1
|
||||
@ -734,6 +842,7 @@ class TestConvenienceFields:
|
||||
parent_run_id=None,
|
||||
tags=["lead_agent"],
|
||||
)
|
||||
await j.flush()
|
||||
|
||||
data = j.get_completion_data()
|
||||
|
||||
@ -756,6 +865,7 @@ class TestConvenienceFields:
|
||||
parent_run_id=None,
|
||||
tags=["lead_agent"],
|
||||
)
|
||||
await j.flush()
|
||||
|
||||
data = j.get_completion_data()
|
||||
|
||||
@ -766,6 +876,7 @@ class TestConvenienceFields:
|
||||
async def test_last_ai_message_extracts_mapping_content(self, journal_setup):
|
||||
j, _ = journal_setup
|
||||
j.on_llm_end(_make_llm_response({"content": "Nested answer"}), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
await j.flush()
|
||||
|
||||
data = j.get_completion_data()
|
||||
|
||||
@ -796,6 +907,7 @@ class TestConvenienceFields:
|
||||
j, _ = journal_setup
|
||||
j.on_llm_end(_make_llm_response("Lead answer"), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
j.on_llm_end(_make_llm_response("Subagent detail"), run_id=uuid4(), parent_run_id=None, tags=["subagent:research"])
|
||||
await j.flush()
|
||||
|
||||
data = j.get_completion_data()
|
||||
|
||||
@ -950,17 +1062,317 @@ class TestCallerBucketing:
|
||||
assert j._lead_agent_tokens == 15
|
||||
assert j._llm_call_count == 1
|
||||
|
||||
def test_first_no_usage_second_with_usage(self, journal_setup):
|
||||
"""First callback with no usage must not block second callback with usage for same run_id."""
|
||||
j, _ = journal_setup
|
||||
@pytest.mark.anyio
|
||||
async def test_dedup_same_run_id_persists_single_message(self, journal_setup):
|
||||
"""A re-fired on_llm_end for one run_id must persist the message once.
|
||||
|
||||
LangChain can deliver on_llm_end more than once for the same run_id.
|
||||
Token accounting already dedups on that; the durable llm.ai.response
|
||||
row must be deduped on the same premise, or count_messages and message
|
||||
pagination (which read append-only rows without dedup) inflate.
|
||||
"""
|
||||
j, store = journal_setup
|
||||
run_id = uuid4()
|
||||
response = _make_llm_response("Answer")
|
||||
j.on_llm_end(response, run_id=run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
j.on_llm_end(response, run_id=run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
await j.flush()
|
||||
messages = await store.list_messages("t1")
|
||||
assert [m["event_type"] for m in messages] == ["llm.ai.response"]
|
||||
assert await store.count_messages("t1") == 1
|
||||
# The run summary counts the message exactly once as well.
|
||||
assert j._msg_count == 1
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_adjacent_late_usage_enriches_canonical_response_only(self, journal_setup):
|
||||
j, store = journal_setup
|
||||
run_id = uuid4()
|
||||
usage = {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}
|
||||
original_tool_calls = [{"id": "call-original", "name": "search", "args": {}}]
|
||||
replay_tool_calls = [{"id": "call-replay", "name": "write_file", "args": {}}]
|
||||
|
||||
j.on_llm_end(
|
||||
_make_llm_response(
|
||||
"Canonical",
|
||||
tool_calls=original_tool_calls,
|
||||
additional_kwargs={
|
||||
"deerflow_error_fallback": True,
|
||||
"error_detail": "canonical fallback",
|
||||
},
|
||||
),
|
||||
run_id=run_id,
|
||||
parent_run_id=None,
|
||||
tags=["lead_agent"],
|
||||
)
|
||||
j.on_llm_end(
|
||||
_make_llm_response(
|
||||
"Replay",
|
||||
usage=usage,
|
||||
tool_calls=replay_tool_calls,
|
||||
additional_kwargs={
|
||||
"deerflow_error_fallback": True,
|
||||
"error_detail": "replay fallback",
|
||||
},
|
||||
),
|
||||
run_id=run_id,
|
||||
parent_run_id=None,
|
||||
tags=["subagent:research"],
|
||||
)
|
||||
await j.flush()
|
||||
|
||||
messages = await store.list_messages("t1")
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"]["content"] == "Canonical"
|
||||
assert messages[0]["content"]["tool_calls"] == original_tool_calls
|
||||
assert messages[0]["content"]["additional_kwargs"]["error_detail"] == "canonical fallback"
|
||||
assert messages[0]["content"]["usage_metadata"] == usage
|
||||
assert messages[0]["metadata"]["caller"] == "lead_agent"
|
||||
assert messages[0]["metadata"]["usage"] == usage
|
||||
assert j._current_run_tool_call_names == {"call-original": "search"}
|
||||
assert j.had_llm_error_fallback is True
|
||||
assert j.llm_error_fallback_message == "canonical fallback"
|
||||
assert j.get_completion_data()["last_ai_message"] == "Canonical"
|
||||
assert j.get_completion_data()["lead_agent_tokens"] == 15
|
||||
assert j.get_completion_data()["subagent_tokens"] == 0
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_same_message_object_replay_cannot_mutate_canonical_summary(self, journal_setup):
|
||||
j, store = journal_setup
|
||||
run_id = uuid4()
|
||||
message = AIMessage(content="Canonical answer")
|
||||
response = LLMResult(generations=[[ChatGeneration(message=message)]])
|
||||
|
||||
j.on_llm_end(response, run_id=run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
message.content = "Replay answer"
|
||||
message.usage_metadata = {
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
"input_token_details": {"cache_read": 3},
|
||||
}
|
||||
j.on_llm_end(response, run_id=run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
message.usage_metadata["input_token_details"]["cache_read"] = 999
|
||||
await j.flush()
|
||||
|
||||
messages = await store.list_messages("t1")
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"]["content"] == "Canonical answer"
|
||||
expected_usage = {
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
"input_token_details": {"cache_read": 3},
|
||||
}
|
||||
assert messages[0]["metadata"]["usage"] == expected_usage
|
||||
assert messages[0]["content"]["usage_metadata"] == expected_usage
|
||||
assert j.get_completion_data()["message_count"] == 1
|
||||
assert j.get_completion_data()["last_ai_message"] == "Canonical answer"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_positive_usage_event_does_not_retain_nested_provider_metadata(self, journal_setup):
|
||||
j, store = journal_setup
|
||||
usage = {
|
||||
"input_tokens": 8,
|
||||
"output_tokens": 3,
|
||||
"total_tokens": 11,
|
||||
"output_token_details": {"reasoning": 2},
|
||||
}
|
||||
message = AIMessage(content="Canonical with usage", usage_metadata=usage)
|
||||
response = LLMResult(generations=[[ChatGeneration(message=message)]])
|
||||
|
||||
j.on_llm_end(response, run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
message.usage_metadata["output_token_details"]["reasoning"] = 999
|
||||
await j.flush()
|
||||
|
||||
messages = await store.list_messages("t1")
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["metadata"]["usage"]["output_token_details"] == {"reasoning": 2}
|
||||
assert messages[0]["content"]["usage_metadata"]["output_token_details"] == {"reasoning": 2}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_mutating_staged_message_before_flush_cannot_mutate_canonical_summary(self, journal_setup):
|
||||
j, store = journal_setup
|
||||
message = AIMessage(content="Canonical before flush")
|
||||
response = LLMResult(generations=[[ChatGeneration(message=message)]])
|
||||
|
||||
j.on_llm_end(response, run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
message.content = "Mutation before flush"
|
||||
await j.flush()
|
||||
|
||||
messages = await store.list_messages("t1")
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"]["content"] == "Canonical before flush"
|
||||
assert j.get_completion_data()["message_count"] == 1
|
||||
assert j.get_completion_data()["last_ai_message"] == "Canonical before flush"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_nested_same_message_object_replay_cannot_mutate_canonical_summary(self, journal_setup):
|
||||
j, store = journal_setup
|
||||
run_id = uuid4()
|
||||
message = AIMessage(content=[{"type": "text", "text": "Canonical nested answer"}])
|
||||
response = LLMResult(generations=[[ChatGeneration(message=message)]])
|
||||
|
||||
j.on_llm_end(response, run_id=run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
message.content[0]["text"] = "Replay nested answer"
|
||||
message.usage_metadata = {"input_tokens": 8, "output_tokens": 3, "total_tokens": 11}
|
||||
j.on_llm_end(response, run_id=run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
await j.flush()
|
||||
|
||||
messages = await store.list_messages("t1")
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"]["content"] == [{"type": "text", "text": "Canonical nested answer"}]
|
||||
assert messages[0]["content"]["usage_metadata"] == message.usage_metadata
|
||||
assert j.get_completion_data()["message_count"] == 1
|
||||
assert j.get_completion_data()["last_ai_message"] == "Canonical nested answer"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_all_zero_usage_remains_pending_and_positive_usage_enriches_it(self, journal_setup):
|
||||
j, store = journal_setup
|
||||
run_id = uuid4()
|
||||
zero_usage = {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
|
||||
positive_usage = {"input_tokens": 4, "output_tokens": 2, "total_tokens": 6}
|
||||
|
||||
j.on_llm_end(_make_llm_response("Zero usage", usage=zero_usage), run_id=run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
assert j._buffer == []
|
||||
assert j._pending_llm_response is not None
|
||||
assert j.get_completion_data()["message_count"] == 0
|
||||
|
||||
j.on_llm_end(_make_llm_response("Replay payload", usage=positive_usage), run_id=run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
await j.flush()
|
||||
|
||||
messages = await store.list_messages("t1")
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"]["content"] == "Zero usage"
|
||||
assert messages[0]["content"]["usage_metadata"] == positive_usage
|
||||
assert messages[0]["metadata"]["usage"] == positive_usage
|
||||
assert j.get_completion_data()["message_count"] == 1
|
||||
assert j.get_completion_data()["last_ai_message"] == "Zero usage"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_replay_generation_length_cannot_change_canonical_set(self, journal_setup):
|
||||
j, store = journal_setup
|
||||
short_usage = {"input_tokens": 8, "output_tokens": 3, "total_tokens": 11}
|
||||
extra_usage = {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}
|
||||
first_run_id = uuid4()
|
||||
second_run_id = uuid4()
|
||||
|
||||
j.on_llm_end(
|
||||
_combine_llm_responses(_make_llm_response("Canonical one"), _make_llm_response("Canonical two")),
|
||||
run_id=first_run_id,
|
||||
parent_run_id=None,
|
||||
tags=["lead_agent"],
|
||||
)
|
||||
j.on_llm_end(
|
||||
_make_llm_response("Short replay", usage=short_usage),
|
||||
run_id=first_run_id,
|
||||
parent_run_id=None,
|
||||
tags=["lead_agent"],
|
||||
)
|
||||
j.on_llm_end(
|
||||
_make_llm_response("Single canonical"),
|
||||
run_id=second_run_id,
|
||||
parent_run_id=None,
|
||||
tags=["lead_agent"],
|
||||
)
|
||||
j.on_llm_end(
|
||||
_combine_llm_responses(
|
||||
_make_llm_response("Long replay one", usage=extra_usage),
|
||||
_make_llm_response("Long replay two"),
|
||||
),
|
||||
run_id=second_run_id,
|
||||
parent_run_id=None,
|
||||
tags=["lead_agent"],
|
||||
)
|
||||
await j.flush()
|
||||
|
||||
messages = await store.list_messages("t1")
|
||||
assert [message["content"]["content"] for message in messages] == [
|
||||
"Canonical one",
|
||||
"Canonical two",
|
||||
"Single canonical",
|
||||
]
|
||||
assert messages[0]["metadata"]["usage"] == short_usage
|
||||
assert messages[0]["content"]["usage_metadata"] == short_usage
|
||||
assert messages[1]["metadata"]["usage"] == {}
|
||||
assert messages[1]["content"]["usage_metadata"] is None
|
||||
assert messages[2]["metadata"]["usage"] == extra_usage
|
||||
assert messages[2]["content"]["usage_metadata"] == extra_usage
|
||||
assert j.get_completion_data()["message_count"] == 3
|
||||
assert j.get_completion_data()["last_ai_message"] == "Single canonical"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_interleaved_late_usage_updates_summary_only(self, journal_setup):
|
||||
j, store = journal_setup
|
||||
first_run_id = uuid4()
|
||||
second_run_id = uuid4()
|
||||
usage = {"input_tokens": 9, "output_tokens": 4, "total_tokens": 13}
|
||||
|
||||
j.on_llm_end(_make_llm_response("First canonical"), run_id=first_run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
j.on_llm_end(_make_llm_response("Second canonical"), run_id=second_run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
j.on_llm_end(
|
||||
_make_llm_response(
|
||||
"Late replay",
|
||||
usage=usage,
|
||||
tool_calls=[{"id": "late-call", "name": "write_file", "args": {}}],
|
||||
additional_kwargs={"deerflow_error_fallback": True, "error_detail": "late fallback"},
|
||||
),
|
||||
run_id=first_run_id,
|
||||
parent_run_id=None,
|
||||
tags=["subagent:research"],
|
||||
)
|
||||
await j.flush()
|
||||
|
||||
messages = await store.list_messages("t1")
|
||||
assert [message["content"]["content"] for message in messages] == ["First canonical", "Second canonical"]
|
||||
assert messages[0]["metadata"]["usage"] == {}
|
||||
assert messages[0]["content"]["usage_metadata"] is None
|
||||
assert j.get_completion_data()["total_tokens"] == 13
|
||||
assert j.get_completion_data()["lead_agent_tokens"] == 13
|
||||
assert j.get_completion_data()["subagent_tokens"] == 0
|
||||
assert j.get_completion_data()["message_count"] == 2
|
||||
assert j.get_completion_data()["last_ai_message"] == "Second canonical"
|
||||
assert "late-call" not in j._current_run_tool_call_names
|
||||
assert j.had_llm_error_fallback is False
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_single_no_usage_response_persists_once_at_flush(self, journal_setup):
|
||||
j, store = journal_setup
|
||||
|
||||
j.on_llm_end(_make_llm_response("No usage"), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
assert j._buffer == []
|
||||
assert j._pending_llm_response is not None
|
||||
|
||||
await j.flush()
|
||||
|
||||
messages = await store.list_messages("t1")
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"]["content"] == "No usage"
|
||||
assert messages[0]["metadata"]["usage"] == {}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_distinct_run_ids_each_persist_a_message(self, journal_setup):
|
||||
"""The dedup guard is per run_id and must not drop distinct responses."""
|
||||
j, store = journal_setup
|
||||
j.on_llm_end(_make_llm_response("First"), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
j.on_llm_end(_make_llm_response("Second"), run_id=uuid4(), parent_run_id=None, tags=["lead_agent"])
|
||||
await j.flush()
|
||||
assert await store.count_messages("t1") == 2
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_first_no_usage_second_with_usage(self, journal_setup):
|
||||
"""Late usage enriches the single canonical event and the run summary."""
|
||||
j, store = journal_setup
|
||||
run_id = uuid4()
|
||||
j.on_llm_end(_make_llm_response("A", usage=None), run_id=run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
assert str(run_id) not in j._counted_llm_run_ids
|
||||
# Second callback for the same run_id with actual usage must still count
|
||||
usage = {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}
|
||||
j.on_llm_end(_make_llm_response("A", usage=usage), run_id=run_id, parent_run_id=None, tags=["lead_agent"])
|
||||
assert j._total_tokens == 15
|
||||
assert j._lead_agent_tokens == 15
|
||||
await j.flush()
|
||||
|
||||
messages = await store.list_messages("t1")
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["metadata"]["usage"] == usage
|
||||
assert messages[0]["content"]["usage_metadata"] == usage
|
||||
assert j.get_completion_data()["total_tokens"] == 15
|
||||
|
||||
def test_track_token_usage_false_skips_buckets(self):
|
||||
"""When token tracking is disabled, caller buckets stay at 0."""
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user