From f8ceab94b4839e118fc2041841b010b85ace3929 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Fri, 23 Jan 2026 13:34:07 +0900 Subject: [PATCH] fix: for test --- .../litellm_core_utils/streaming_handler.py | 55 +++++++++++++++ .../responses/mcp/chat_completions_handler.py | 69 ++++++++++++++++--- .../mcp/litellm_proxy_mcp_handler.py | 6 ++ 3 files changed, 119 insertions(+), 11 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 3304759f749..c6f0f67976f 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1571,6 +1571,50 @@ class CustomStreamWrapper: ) return chunk + def _add_mcp_list_tools_to_first_chunk(self, chunk: ModelResponseStream) -> ModelResponseStream: + """ + Add mcp_list_tools from _hidden_params to the first chunk's delta.provider_specific_fields. + + This method checks if MCP metadata with mcp_list_tools is stored in _hidden_params + and adds it to the first chunk's delta.provider_specific_fields. + """ + try: + # Check if MCP metadata should be added to first chunk + if not hasattr(self, "_hidden_params") or not self._hidden_params: + return chunk + + mcp_metadata = self._hidden_params.get("mcp_metadata") + if not mcp_metadata or not isinstance(mcp_metadata, dict): + return chunk + + # Only add mcp_list_tools to first chunk (not tool_calls or tool_results) + mcp_list_tools = mcp_metadata.get("mcp_list_tools") + if not mcp_list_tools: + return chunk + + # Add mcp_list_tools to delta.provider_specific_fields + if hasattr(chunk, "choices") and chunk.choices: + for choice in chunk.choices: + if isinstance(choice, StreamingChoices) and hasattr(choice, "delta") and choice.delta: + # Get existing provider_specific_fields or create new dict + provider_fields = ( + getattr(choice.delta, "provider_specific_fields", None) or {} + ) + + # Add only mcp_list_tools to first chunk + provider_fields["mcp_list_tools"] = mcp_list_tools + + # Set the provider_specific_fields + setattr(choice.delta, "provider_specific_fields", provider_fields) + + except Exception as e: + from litellm._logging import verbose_logger + verbose_logger.exception( + f"Error adding MCP list tools to first chunk: {str(e)}" + ) + + return chunk + def _add_mcp_metadata_to_final_chunk(self, chunk: ModelResponseStream) -> ModelResponseStream: """ Add MCP metadata from _hidden_params to the final chunk's delta.provider_specific_fields. @@ -1727,6 +1771,12 @@ class CustomStreamWrapper: ) # HANDLE STREAM OPTIONS self.chunks.append(response) + + # Add mcp_list_tools to first chunk if present + if not self.sent_first_chunk: + response = self._add_mcp_list_tools_to_first_chunk(response) + self.sent_first_chunk = True + if hasattr( response, "usage" ): # remove usage from chunk, only send on final chunk @@ -1894,6 +1944,11 @@ class CustomStreamWrapper: input=self.response_uptil_now, model=self.model ) self.chunks.append(processed_chunk) + + # Add mcp_list_tools to first chunk if present + if not self.sent_first_chunk: + processed_chunk = self._add_mcp_list_tools_to_first_chunk(processed_chunk) + self.sent_first_chunk = True if hasattr( processed_chunk, "usage" ): # remove usage from chunk, only send on final chunk diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index d05e6c8c684..2b4f91a35cb 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -244,14 +244,14 @@ async def acompletion_with_mcp( for choice in chunk.choices: if isinstance(choice, StreamingChoices) and hasattr(choice, "delta") and choice.delta: # Get existing provider_specific_fields or create new dict - provider_fields = ( - getattr(choice.delta, "provider_specific_fields", None) or {} - ) + existing_fields = getattr(choice.delta, "provider_specific_fields", None) or {} + provider_fields = dict(existing_fields) # Create a copy to avoid mutating the original # Add only mcp_list_tools to first chunk provider_fields["mcp_list_tools"] = self.openai_tools - # Set the provider_specific_fields + # Set provider_specific_fields directly using setattr + # This ensures the modification is preserved setattr(choice.delta, "provider_specific_fields", provider_fields) return chunk @@ -264,9 +264,8 @@ async def acompletion_with_mcp( for choice in chunk.choices: if isinstance(choice, StreamingChoices) and hasattr(choice, "delta") and choice.delta: # Get existing provider_specific_fields or create new dict - provider_fields = ( - getattr(choice.delta, "provider_specific_fields", None) or {} - ) + existing_fields = getattr(choice.delta, "provider_specific_fields", None) or {} + provider_fields = dict(existing_fields) # Create a copy to avoid mutating the original # Add tool_calls and tool_results if available if self.tool_calls: @@ -274,7 +273,8 @@ async def acompletion_with_mcp( if self.tool_results: provider_fields["mcp_call_results"] = self.tool_results - # Set the provider_specific_fields + # Set provider_specific_fields directly using setattr + # This ensures the modification is preserved setattr(choice.delta, "provider_specific_fields", provider_fields) return chunk @@ -405,8 +405,6 @@ async def acompletion_with_mcp( if not self.tool_results or not self.complete_response: return - from litellm import acompletion as litellm_acompletion - # Create follow-up messages with tool results follow_up_messages = LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat( original_messages=self.messages, @@ -421,7 +419,10 @@ async def acompletion_with_mcp( # Ensure follow-up call doesn't trigger MCP handler again follow_up_call_args["_skip_mcp_handler"] = True - follow_up_response = await litellm_acompletion(**follow_up_call_args) + # Import litellm here to ensure we get the patched version + # This ensures the patch works correctly in tests + import litellm + follow_up_response = await litellm.acompletion(**follow_up_call_args) # Ensure follow-up response is a CustomStreamWrapper if isinstance(follow_up_response, CustomStreamWrapper): @@ -471,14 +472,60 @@ async def acompletion_with_mcp( # Copy important attributes from original wrapper if hasattr(original_wrapper, "_hidden_params"): self._hidden_params = original_wrapper._hidden_params + # For synchronous iteration, we need to run the async iterator + self._sync_iterator = None + self._sync_loop = None def __aiter__(self): return self._custom_iterator + def __iter__(self): + # For synchronous iteration, create a sync wrapper + if self._sync_iterator is None: + import asyncio + try: + self._sync_loop = asyncio.get_event_loop() + except RuntimeError: + self._sync_loop = asyncio.new_event_loop() + asyncio.set_event_loop(self._sync_loop) + self._sync_iterator = _SyncIteratorWrapper(self._custom_iterator, self._sync_loop) + return self._sync_iterator + + def __next__(self): + # Delegate to sync iterator + if self._sync_iterator is None: + self.__iter__() + return next(self._sync_iterator) + def __getattr__(self, name): # Delegate all other attributes to original wrapper return getattr(self._original_wrapper, name) + # Helper class to wrap async iterator for sync iteration + class _SyncIteratorWrapper: + def __init__(self, async_iterator, loop): + self._async_iterator = async_iterator + self._loop = loop + self._iterator = None + + def __iter__(self): + return self + + def __next__(self): + if self._iterator is None: + # __aiter__ might be async, so we need to await it + aiter_result = self._async_iterator.__aiter__() + if hasattr(aiter_result, '__await__'): + # It's a coroutine, await it + self._iterator = self._loop.run_until_complete(aiter_result) + else: + # It's already an iterator + self._iterator = aiter_result + try: + return self._loop.run_until_complete(self._iterator.__anext__()) + except StopAsyncIteration: + raise StopIteration + return cast(CustomStreamWrapper, MCPStreamWrapper(initial_stream, iterator)) # Non-streaming mode: use existing logic diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 4376e076a95..930f7261cd7 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -775,6 +775,12 @@ class LiteLLM_Proxy_MCP_Handler: first_choice, "message", None ): message_to_append = first_choice.message.model_dump(exclude_none=True) + # Ensure tool_calls have arguments field (required by OpenAI API) + if message_to_append.get("tool_calls"): + for tool_call in message_to_append["tool_calls"]: + if isinstance(tool_call, dict) and "function" in tool_call: + if "arguments" not in tool_call["function"]: + tool_call["function"]["arguments"] = "{}" except Exception: verbose_logger.exception("Failed to convert assistant message for MCP flow")