diff --git a/backend/packages/harness/deerflow/runtime/AGENTS.md b/backend/packages/harness/deerflow/runtime/AGENTS.md index 4da7ab93e..d77bc9f23 100644 --- a/backend/packages/harness/deerflow/runtime/AGENTS.md +++ b/backend/packages/harness/deerflow/runtime/AGENTS.md @@ -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 diff --git a/backend/packages/harness/deerflow/runtime/journal.py b/backend/packages/harness/deerflow/runtime/journal.py index 1a5945889..3aa4ff446 100644 --- a/backend/packages/harness/deerflow/runtime/journal.py +++ b/backend/packages/harness/deerflow/runtime/journal.py @@ -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() diff --git a/backend/tests/test_run_journal.py b/backend/tests/test_run_journal.py index 061eba45e..c7115e2a0 100644 --- a/backend/tests/test_run_journal.py +++ b/backend/tests/test_run_journal.py @@ -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."""