From 27b2b6768032010efb65d77e15d058e5d0882a25 Mon Sep 17 00:00:00 2001 From: Coder-xiaosuo <3476584763@qq.com> Date: Sat, 5 Sep 2026 14:10:18 +0800 Subject: [PATCH] fix(models): restore usage_metadata in MindIE tool-mode simulated streaming (#5195) In tool-enabled requests MindIEChatModel._astream falls back to awaiting the full _agenerate response and re-emitting it as simulated AIMessageChunks. The full response carries usage_metadata, but none of the simulated chunks copied it, so chunk aggregation (add_ai_message_chunks) produced a final message with usage_metadata=None. Token usage therefore vanished from token accounting, run stats, persistence and the UI for every tool-enabled streamed turn. Mirror OpenAI's terminal-usage-frame convention: attach msg.usage_metadata to exactly the last simulated chunk (the trailing tool-call chunk when present, else the last text chunk / the single tool-only chunk) so the aggregated message carries it exactly once. add_usage() is per-chunk additive, so attaching usage to every chunk would multiply the totals. Scope: MindIEChatModel only; other providers keep native streaming and ainvoke/non-tool astream were already correct. Tests: regression guard asserting exactly one carrier chunk equals the last one and that merged usage equals the original across all three simulated-stream branches, plus chain-level tests driving the public astream() wrapper and asserting the persisted model_dump() shape. Closes #5192 --- .../deerflow/models/mindie_provider.py | 17 ++- backend/tests/test_mindie_provider.py | 133 +++++++++++++++++- 2 files changed, 146 insertions(+), 4 deletions(-) diff --git a/backend/packages/harness/deerflow/models/mindie_provider.py b/backend/packages/harness/deerflow/models/mindie_provider.py index 6d88e7e13..581389d40 100644 --- a/backend/packages/harness/deerflow/models/mindie_provider.py +++ b/backend/packages/harness/deerflow/models/mindie_provider.py @@ -238,17 +238,28 @@ class MindIEChatModel(ChatOpenAI): msg = gen.message content = msg.content standard_tool_calls = getattr(msg, "tool_calls", []) + # Attach the full response's terminal usage to the *last* simulated + # chunk (OpenAI terminal-frame style) so add_usage() counts it once. + usage_metadata = getattr(msg, "usage_metadata", None) # Yield text in chunks to allow downstream UI/Markdown parsers to render smoothly if isinstance(content, str) and content: chunk_size = 15 for i in range(0, len(content), chunk_size): chunk_text = content[i : i + chunk_size] - chunk_msg = AIMessageChunk(content=chunk_text, id=msg.id, response_metadata=msg.response_metadata if i == 0 else {}) + # Without tool calls the last text chunk terminates the stream. + is_final_chunk = i + chunk_size >= len(content) + chunk_msg = AIMessageChunk( + content=chunk_text, + id=msg.id, + response_metadata=msg.response_metadata if i == 0 else {}, + usage_metadata=usage_metadata if (not standard_tool_calls and is_final_chunk) else None, + ) yield ChatGenerationChunk(message=chunk_msg, generation_info=gen.generation_info if i == 0 else None) if standard_tool_calls: - yield ChatGenerationChunk(message=AIMessageChunk(content="", id=msg.id, tool_calls=standard_tool_calls, invalid_tool_calls=getattr(msg, "invalid_tool_calls", []))) + # Tool-call chunk terminates the stream: carry the usage here. + yield ChatGenerationChunk(message=AIMessageChunk(content="", id=msg.id, tool_calls=standard_tool_calls, invalid_tool_calls=getattr(msg, "invalid_tool_calls", []), usage_metadata=usage_metadata)) else: - chunk_msg = AIMessageChunk(content=content, id=msg.id, tool_calls=standard_tool_calls, invalid_tool_calls=getattr(msg, "invalid_tool_calls", [])) + chunk_msg = AIMessageChunk(content=content, id=msg.id, tool_calls=standard_tool_calls, invalid_tool_calls=getattr(msg, "invalid_tool_calls", []), usage_metadata=usage_metadata) yield ChatGenerationChunk(message=chunk_msg, generation_info=gen.generation_info) diff --git a/backend/tests/test_mindie_provider.py b/backend/tests/test_mindie_provider.py index 51fb21165..66747dc29 100644 --- a/backend/tests/test_mindie_provider.py +++ b/backend/tests/test_mindie_provider.py @@ -20,10 +20,12 @@ from deerflow.models.mindie_provider import ( # ═════════════════════════════════════════════════════════════════════════════ -def _make_chat_result(content: str, tool_calls=None) -> ChatResult: +def _make_chat_result(content: str, tool_calls=None, usage_metadata=None) -> ChatResult: msg = AIMessage(content=content) if tool_calls: msg.tool_calls = tool_calls + if usage_metadata is not None: + msg.usage_metadata = usage_metadata gen = ChatGeneration(message=msg) return ChatResult(generations=[gen]) @@ -492,3 +494,132 @@ class TestAStream: chunks = await self._collect(model._astream([HumanMessage(content="q")], tools=[{"type": "function", "function": {"name": "x"}}])) assert any(getattr(c.message, "tool_calls", []) for c in chunks) + + # ── Issue #5192: usage_metadata dropped in tool-mode simulated stream ──── + + _USAGE = {"input_tokens": 12, "output_tokens": 8, "total_tokens": 20} + + @staticmethod + def _tool(name: str) -> dict: + return {"type": "function", "function": {"name": name}} + + async def _collect_stream_with_usage(self, content, tool_calls): + """Collect the tool-mode simulated stream whose underlying `_agenerate` + result carries usage_metadata; returns (chunks, source_usage).""" + with patch.object(MindIEChatModel, "_agenerate", new_callable=AsyncMock) as mock_ag, patch.object(MindIEChatModel, "__init__", return_value=None): + mock_ag.return_value = _make_chat_result(content, tool_calls=tool_calls, usage_metadata=self._USAGE) + model = MindIEChatModel.__new__(MindIEChatModel) + chunks = await self._collect(model._astream([HumanMessage(content="q")], tools=[self._tool("fn")])) + source_usage = mock_ag.return_value.generations[0].message.usage_metadata + + return chunks, source_usage + + @staticmethod + def _merge_messages(chunks): + merged = chunks[0].message + for chunk in chunks[1:]: + merged = merged + chunk.message + return merged + + @pytest.mark.parametrize( + ("content", "tool_calls"), + [ + ("A" * 40, None), # text-only simulated stream + ("A" * 40, [{"name": "fn", "args": {"x": 1}, "id": "c1"}]), # text + trailing tool-call chunk + ("", [{"name": "fn", "args": {"x": 1}, "id": "c1"}]), # tool-call only + ], + ) + @pytest.mark.asyncio + async def test_with_tools_usage_metadata_survives_simulated_stream(self, content, tool_calls): + """Issue #5192 regression guard: usage must survive the simulated stream. + + Chunk level: exactly the *last* emitted chunk carries the usage snapshot + (mirroring OpenAI's terminal usage frame, so chunk summation cannot + double count). Aggregate level: merging the simulated chunks must + reproduce the original usage_metadata. + """ + chunks, source_usage = await self._collect_stream_with_usage(content, tool_calls) + + # Sanity: the underlying full response really did carry usage. + assert source_usage == self._USAGE + + carriers = [c for c in chunks if c.message.usage_metadata is not None] + assert len(carriers) == 1 + assert carriers[0] is chunks[-1] + assert carriers[0].message.usage_metadata == self._USAGE + + merged = self._merge_messages(chunks) + assert merged.usage_metadata == self._USAGE + + +# ═════════════════════════════════════════════════════════════════════════════ +# 7. Chain-level regression (Issue #5192): public astream() → persisted shape +# ═════════════════════════════════════════════════════════════════════════════ + + +class TestAStreamUsageChain: + """End-to-end guard for the tool-mode usage path. + + Drives the *public* ``astream()`` wrapper (the real BaseChatModel path that + LangGraph state accumulation, journal ``on_llm_end`` and the front-end + ``usage_metadata`` field all consume), then checks the aggregated message + shape that gets persisted/streamed (``model_dump()``). This covers the + interaction with the wrapper's trailing ``chunk_position="last"`` empty + chunk, which the unit-level merge above does not exercise. + """ + + _USAGE = {"input_tokens": 12, "output_tokens": 8, "total_tokens": 20} + _TOOLS = [{"type": "function", "function": {"name": "fn"}}] + + @pytest.mark.asyncio + async def test_public_astream_keeps_usage_for_text_and_tool_call(self): + usage = self._USAGE + tool_calls = [{"name": "fn", "args": {"x": 1}, "id": "c1"}] + long_text = "A" * 40 + + with patch.object(MindIEChatModel, "_agenerate", new_callable=AsyncMock) as mock_ag: + mock_ag.return_value = _make_chat_result(long_text, tool_calls=tool_calls, usage_metadata=usage) + model = MindIEChatModel(model="mindie-test", api_key="test-key") + + # Collect from the public wrapper, exactly as a graph node would. + chunks = [] + async for chunk in model.astream([HumanMessage(content="q")], tools=self._TOOLS): + chunks.append(chunk) + + assert chunks, "public astream() yielded nothing" + merged = chunks[0] + for chunk in chunks[1:]: + merged = merged + chunk + + # Text and tool calls survive the simulated stream untouched. Note the + # aggregated AIMessage normalises tool calls to include ``type``. + assert merged.content == long_text + assert merged.tool_calls == [{**tool_calls[0], "type": "tool_call"}] + # Usage survives the wrapper aggregation exactly once (no chunk-count + # multiplication even with the synthetic trailing empty chunk), and is + # present in the exact shape journal/persistence/front-end consume. + assert merged.usage_metadata == usage + dumped = merged.model_dump() + assert dumped.get("usage_metadata") == usage + + @pytest.mark.asyncio + async def test_public_astream_keeps_usage_for_tool_call_only(self): + usage = self._USAGE + tool_calls = [{"name": "fn", "args": {"x": 1}, "id": "c2"}] + + with patch.object(MindIEChatModel, "_agenerate", new_callable=AsyncMock) as mock_ag: + mock_ag.return_value = _make_chat_result("", tool_calls=tool_calls, usage_metadata=usage) + model = MindIEChatModel(model="mindie-test", api_key="test-key") + + chunks = [] + async for chunk in model.astream([HumanMessage(content="q")], tools=self._TOOLS): + chunks.append(chunk) + + assert chunks + merged = chunks[0] + for chunk in chunks[1:]: + merged = merged + chunk + + assert merged.tool_calls == [{**tool_calls[0], "type": "tool_call"}] + assert merged.usage_metadata == usage + assert merged.model_dump().get("usage_metadata") == usage