From 119926f1123496336066f77000c2aeebe3e67f0d Mon Sep 17 00:00:00 2001 From: fangkangmi Date: Thu, 2 Jul 2026 22:45:28 +0100 Subject: [PATCH] fix(mcp): dedupe held tool-call chunk in abrupt-termination fallback Review follow-up: when the initial stream ends via StopAsyncIteration without a finish_reason chunk and the last collected chunk is a held tool-call delta, the chunk existed in both held_tool_call_chunks and collected_chunks[-1], so the fallback flush yielded it twice. The flush now filters the final chunk out of the held list by identity, via a shared _flush_held_and_final helper used by both fallback paths Also switch pending_chunks from list.pop(0) to collections.deque per review suggestion --- .../responses/mcp/chat_completions_handler.py | 23 +++--- .../mcp/test_chat_completions_handler.py | 73 +++++++++++++++++++ 2 files changed, 86 insertions(+), 10 deletions(-) diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index d440d9a93ac..14d1cb641f9 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -1,6 +1,7 @@ """Helpers for handling MCP-aware `/chat/completions` requests.""" import logging +from collections import deque from typing import ( Any, List, @@ -246,7 +247,7 @@ async def acompletion_with_mcp( self.follow_up_iterator = None self.follow_up_exhausted = False self.held_tool_call_chunks: list[ModelResponseStream] = [] - self.pending_chunks: list[ModelResponseStream] = [] + self.pending_chunks: deque[ModelResponseStream] = deque() self.any_chunk_yielded = False async def __aiter__(self): @@ -352,9 +353,17 @@ async def acompletion_with_mcp( if self.follow_up_stream is not None: self.held_tool_call_chunks = [] + def _flush_held_and_final(self, final_chunk: ModelResponseStream) -> ModelResponseStream: + flushed_final = self._add_mcp_tool_metadata_to_final_chunk(final_chunk) + self.pending_chunks = deque( + [chunk for chunk in self.held_tool_call_chunks if chunk is not final_chunk] + [flushed_final] + ) + self.held_tool_call_chunks = [] + return self._yield_chunk(self.pending_chunks.popleft()) + async def __anext__(self): if self.pending_chunks: - return self._yield_chunk(self.pending_chunks.pop(0)) + return self._yield_chunk(self.pending_chunks.popleft()) # Phase 1: Collect and yield initial stream chunks while not self.stream_exhausted: @@ -376,10 +385,7 @@ async def acompletion_with_mcp( await self._finish_initial_turn() if self.follow_up_stream is not None: break - final_chunk = self._add_mcp_tool_metadata_to_final_chunk(self.collected_chunks[-1]) - self.pending_chunks = self.held_tool_call_chunks + [final_chunk] - self.held_tool_call_chunks = [] - return self._yield_chunk(self.pending_chunks.pop(0)) + return self._flush_held_and_final(self.collected_chunks[-1]) self.collected_chunks.append(chunk) @@ -388,10 +394,7 @@ async def acompletion_with_mcp( await self._finish_initial_turn() if self.follow_up_stream is not None: break - chunk = self._add_mcp_tool_metadata_to_final_chunk(chunk) - self.pending_chunks = self.held_tool_call_chunks + [chunk] - self.held_tool_call_chunks = [] - return self._yield_chunk(self.pending_chunks.pop(0)) + return self._flush_held_and_final(chunk) if self._chunk_has_tool_call_delta(chunk): self.held_tool_call_chunks.append(chunk) diff --git a/tests/test_litellm/responses/mcp/test_chat_completions_handler.py b/tests/test_litellm/responses/mcp/test_chat_completions_handler.py index 64808c529ef..780a3f7a656 100644 --- a/tests/test_litellm/responses/mcp/test_chat_completions_handler.py +++ b/tests/test_litellm/responses/mcp/test_chat_completions_handler.py @@ -1605,3 +1605,76 @@ async def test_acompletion_with_mcp_streaming_flushes_tool_call_turn_when_no_fol getattr(chunk.choices[0].delta, "tool_calls", None) for chunk in all_chunks if chunk.choices ), "Held tool-call chunks must be flushed when no follow-up stream is created" assert all_chunks[-1].choices[0].finish_reason == "tool_calls" + + +@pytest.mark.asyncio +async def test_acompletion_with_mcp_streaming_no_duplicate_chunk_on_abrupt_termination(monkeypatch): + """ + Regression test for a duplicate-chunk bug in the abrupt-termination fallback: + when the initial stream ends via StopAsyncIteration without a finish_reason + chunk and the last collected chunk is a held tool-call delta, that chunk is + both in held_tool_call_chunks and collected_chunks[-1]. It must be yielded + to the client exactly once. + """ + from litellm.types.utils import ChatCompletionDeltaToolCall, Function + + tools = [{"type": "mcp", "server_url": "litellm_proxy/mcp/local"}] + openai_tools = [{"type": "function", "function": {"name": "local_search"}}] + tool_calls = [ + { + "id": "call-1", + "type": "function", + "function": {"name": "local_search", "arguments": "{}"}, + } + ] + + initial_chunks = [ + _create_mcp_stream_chunk("partial answer "), + _create_mcp_stream_chunk( + None, + tool_calls=[ + ChatCompletionDeltaToolCall( + id="call-1", + type="function", + function=Function(name="local_search", arguments="{}"), + index=0, + ) + ], + ), + ] + + InitialStream = _make_mock_stream_class(initial_chunks) + + async def mock_acompletion(**kwargs): + return InitialStream() + + mock_acompletion_func = AsyncMock(side_effect=mock_acompletion) + _patch_mcp_auto_exec_scaffolding(monkeypatch, tools, openai_tools, tool_calls, tool_results=[]) + + with ( + patch("litellm.acompletion", mock_acompletion_func), + patch.object( + chat_completions_handler, + "litellm_acompletion", + mock_acompletion_func, + create=True, + ), + ): + result = await acompletion_with_mcp( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "hello"}], + tools=tools, + stream=True, + ) + + all_chunks = [] + async for chunk in result: + all_chunks.append(chunk) + + tool_call_chunk_count = sum( + 1 for chunk in all_chunks if chunk.choices and getattr(chunk.choices[0].delta, "tool_calls", None) + ) + assert tool_call_chunk_count == 1, ( + f"The held tool-call chunk must be yielded exactly once on abrupt termination. Got chunks: {all_chunks}" + ) + assert len(all_chunks) == len(set(id(chunk) for chunk in all_chunks)), "No chunk object may be yielded twice"