mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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
This commit is contained in:
parent
8e8e96d6a8
commit
119926f112
2 changed files with 86 additions and 10 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue