From 9597c6c102155bcd2ef023dae258ad898335d943 Mon Sep 17 00:00:00 2001 From: tarunag10 Date: Thu, 20 Aug 2026 12:51:55 +0530 Subject: [PATCH] style: apply repository lint rules --- .../responses/mcp/chat_completions_handler.py | 143 ++++++------------ .../mcp/litellm_proxy_mcp_handler.py | 12 +- .../responses/mcp/mcp_streaming_iterator.py | 6 +- .../mcp/test_litellm_proxy_mcp_handler.py | 57 +++---- 4 files changed, 69 insertions(+), 149 deletions(-) diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 7df0c34aa47..2b3c0d79192 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -3,13 +3,14 @@ import logging from typing import TYPE_CHECKING, Any, Final, cast -from litellm.responses.mcp.litellm_proxy_mcp_handler import ( - LiteLLM_Proxy_MCP_Handler, -) from litellm.responses.mcp.request_context import MCPRequestContext from litellm.types.utils import Message, ModelResponse from litellm.utils import CustomStreamWrapper +from litellm.responses.mcp.litellm_proxy_mcp_handler import ( + LiteLLM_Proxy_MCP_Handler, +) + if TYPE_CHECKING: from litellm.proxy._types import UserAPIKeyAuth @@ -138,9 +139,7 @@ async def acompletion_with_mcp( ) # Combine with other tools - all_tools: Final = ( - openai_tools + other_tools if (openai_tools or other_tools) else None - ) + all_tools: Final = openai_tools + other_tools if (openai_tools or other_tools) else None # Determine if we should auto-execute tools should_auto_execute: Final = LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools( @@ -241,9 +240,7 @@ async def acompletion_with_mcp( def __aiter__(self): return self - def _add_mcp_list_tools_to_chunk( - self, chunk: ModelResponseStream - ) -> ModelResponseStream: + def _add_mcp_list_tools_to_chunk(self, chunk: ModelResponseStream) -> ModelResponseStream: """Add mcp_list_tools to the first chunk.""" from litellm.types.utils import ( StreamingChoices, @@ -255,19 +252,10 @@ async def acompletion_with_mcp( if hasattr(chunk, "choices") and chunk.choices: for choice in chunk.choices: - if ( - isinstance(choice, StreamingChoices) - and hasattr(choice, "delta") - and choice.delta - ): + if isinstance(choice, StreamingChoices) and hasattr(choice, "delta") and choice.delta: # Get existing provider_specific_fields or create new dict - existing_fields = ( - getattr(choice.delta, "provider_specific_fields", None) - or {} - ) - provider_fields = dict( - existing_fields - ) # Create a copy to avoid mutating the original + 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 @@ -278,9 +266,7 @@ async def acompletion_with_mcp( return chunk - def _add_mcp_tool_metadata_to_final_chunk( - self, chunk: ModelResponseStream - ) -> ModelResponseStream: + def _add_mcp_tool_metadata_to_final_chunk(self, chunk: ModelResponseStream) -> ModelResponseStream: """Add mcp_tool_calls and mcp_call_results to the final chunk.""" from litellm.types.utils import ( StreamingChoices, @@ -289,25 +275,15 @@ async def acompletion_with_mcp( if hasattr(chunk, "choices") and chunk.choices: for choice in chunk.choices: - if ( - isinstance(choice, StreamingChoices) - and hasattr(choice, "delta") - and choice.delta - ): + if isinstance(choice, StreamingChoices) and hasattr(choice, "delta") and choice.delta: # Get existing provider_specific_fields or create new dict # Access the attribute directly to handle Pydantic model attributes correctly existing_fields = {} if hasattr(choice.delta, "provider_specific_fields"): - attr_value = getattr( - choice.delta, "provider_specific_fields", None - ) + attr_value = getattr(choice.delta, "provider_specific_fields", None) if attr_value is not None: # Create a copy to avoid mutating the original - existing_fields = ( - dict(attr_value) - if isinstance(attr_value, dict) - else {} - ) + existing_fields = dict(attr_value) if isinstance(attr_value, dict) else {} provider_fields = existing_fields @@ -382,9 +358,7 @@ async def acompletion_with_mcp( # If we have chunks, yield the final one with metadata if self.collected_chunks: final_chunk = self.collected_chunks[-1] - final_chunk = self._add_mcp_tool_metadata_to_final_chunk( - final_chunk - ) + final_chunk = self._add_mcp_tool_metadata_to_final_chunk(final_chunk) # If we have tool results, prepare follow-up call if self.tool_results and self.complete_response: await self._prepare_follow_up_call() @@ -422,9 +396,7 @@ async def acompletion_with_mcp( ): from litellm._logging import verbose_logger - verbose_logger.warning( - "Follow-up stream was not created despite having tool results" - ) + verbose_logger.warning("Follow-up stream was not created despite having tool results") raise StopAsyncIteration @@ -453,19 +425,17 @@ async def acompletion_with_mcp( if self.tool_calls: # Execute tool calls - self.tool_results = ( - await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( - tool_server_map=self.tool_server_map, - tool_calls=self.tool_calls, - user_api_key_auth=self.user_api_key_auth, - mcp_auth_header=self.mcp_auth_header, - mcp_server_auth_headers=self.mcp_server_auth_headers, - oauth2_headers=self.oauth2_headers, - raw_headers=self.raw_headers, - litellm_call_id=self.litellm_call_id, - litellm_trace_id=self.litellm_trace_id, - request_tags=self.request_tags, - ) + self.tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map=self.tool_server_map, + tool_calls=self.tool_calls, + user_api_key_auth=self.user_api_key_auth, + mcp_auth_header=self.mcp_auth_header, + mcp_server_auth_headers=self.mcp_server_auth_headers, + oauth2_headers=self.oauth2_headers, + raw_headers=self.raw_headers, + litellm_call_id=self.litellm_call_id, + litellm_trace_id=self.litellm_trace_id, + request_tags=self.request_tags, ) async def _prepare_follow_up_call(self): @@ -477,19 +447,15 @@ async def acompletion_with_mcp( return # Create follow-up messages with tool results - follow_up_messages: Final = ( - LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat( - original_messages=self.messages, - response=self.complete_response, - tool_results=self.tool_results, - ) + follow_up_messages: Final = LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat( + original_messages=self.messages, + response=self.complete_response, + tool_results=self.tool_results, ) # Make follow-up call with streaming - follow_up_call_args: Final = ( - LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params( - self.base_call_args - ) + follow_up_call_args: Final = LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params( + self.base_call_args ) follow_up_call_args["messages"] = follow_up_messages follow_up_call_args["stream"] = True @@ -500,9 +466,7 @@ async def acompletion_with_mcp( # This ensures the patch works correctly in tests import litellm - follow_up_response: Final = await litellm.acompletion( - **follow_up_call_args - ) + follow_up_response: Final = await litellm.acompletion(**follow_up_call_args) # Ensure follow-up response is a CustomStreamWrapper if isinstance(follow_up_response, CustomStreamWrapper): @@ -515,8 +479,7 @@ async def acompletion_with_mcp( from litellm._logging import verbose_logger verbose_logger.warning( - "Follow-up response is not a CustomStreamWrapper: %s", - type(follow_up_response), + "Follow-up response is not a CustomStreamWrapper: %s", type(follow_up_response) ) self.follow_up_stream = None @@ -540,24 +503,16 @@ async def acompletion_with_mcp( # Create a wrapper class that delegates to our custom iterator # We'll use a simple approach: just replace the __aiter__ method class MCPStreamWrapper(CustomStreamWrapper): - def __init__( - self, - original_wrapper: CustomStreamWrapper, - custom_iterator: MCPStreamingIterator, - ): + def __init__(self, original_wrapper: CustomStreamWrapper, custom_iterator: MCPStreamingIterator): # Initialize with the same parameters as original wrapper super().__init__( completion_stream=None, model=getattr(original_wrapper, "model", "unknown"), logging_obj=original_wrapper.logging_obj, - custom_llm_provider=getattr( - original_wrapper, "custom_llm_provider", None - ), + custom_llm_provider=getattr(original_wrapper, "custom_llm_provider", None), stream_options=getattr(original_wrapper, "stream_options", None), make_call=getattr(original_wrapper, "make_call", None), - _response_headers=getattr( - original_wrapper, "_response_headers", None - ), + _response_headers=getattr(original_wrapper, "_response_headers", None), ) self._original_wrapper = original_wrapper self._custom_iterator = custom_iterator @@ -581,9 +536,7 @@ async def acompletion_with_mcp( 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 - ) + self._sync_iterator = _SyncIteratorWrapper(self._custom_iterator, self._sync_loop) return self._sync_iterator def __next__(self): @@ -636,11 +589,7 @@ async def acompletion_with_mcp( return initial_response # Extract tool calls from response - tool_calls: Final = ( - LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response( - response=initial_response - ) - ) + tool_calls: Final = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(response=initial_response) if not tool_calls: _add_mcp_metadata_to_response( @@ -672,18 +621,14 @@ async def acompletion_with_mcp( return initial_response # Create follow-up messages with tool results - follow_up_messages: Final = ( - LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat( - original_messages=messages, - response=initial_response, - tool_results=tool_results, - ) + follow_up_messages: Final = LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat( + original_messages=messages, + response=initial_response, + tool_results=tool_results, ) # Make follow-up call with original stream setting - follow_up_call_args: Final = ( - LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params(base_call_args) - ) + follow_up_call_args: Final = LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params(base_call_args) follow_up_call_args["messages"] = follow_up_messages follow_up_call_args["stream"] = stream diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 46bb067eeeb..d2e9c8f37ba 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -4,9 +4,6 @@ from collections.abc import Iterable, Mapping, Sequence from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypedDict, overload -from openai.types.chat import ChatCompletionToolParam -from openai.types.responses.function_tool_param import FunctionToolParam - from litellm._logging import verbose_logger from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -33,13 +30,14 @@ from litellm.types.utils import ( StandardLoggingMCPToolCall, ) from litellm.utils import Rules, function_setup +from openai.types.chat import ChatCompletionToolParam +from openai.types.responses.function_tool_param import FunctionToolParam if TYPE_CHECKING: - from mcp.types import CallToolResult - from mcp.types import Tool as MCPTool - from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging + from mcp.types import CallToolResult + from mcp.types import Tool as MCPTool else: MCPTool = Any @@ -669,7 +667,6 @@ class LiteLLM_Proxy_MCP_Handler: ) -> list[MCPToolResult]: """Execute tool calls and return results.""" from fastapi import HTTPException - from litellm._uuid import uuid from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -1205,6 +1202,7 @@ class LiteLLM_Proxy_MCP_Handler: List of MCP tool execution events for streaming """ from litellm._uuid import uuid + from litellm.responses.mcp.mcp_streaming_iterator import create_mcp_call_events tool_execution_events: Final[list[ResponsesAPIStreamingResponse]] = [] diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index 69e698f7f65..89774568d57 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -22,9 +22,9 @@ from litellm.types.llms.openai import ( ) if TYPE_CHECKING: + from litellm.proxy._types import UserAPIKeyAuth from mcp.types import Tool as MCPTool - from litellm.proxy._types import UserAPIKeyAuth from litellm.responses.mcp.litellm_proxy_mcp_handler import MCPToolResult else: MCPTool = Any @@ -331,11 +331,10 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): def _extract_mcp_headers_from_params(self) -> None: """Extract MCP headers from original request params to pass to tool calls""" - from starlette.datastructures import Headers - from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) + from starlette.datastructures import Headers # Extract headers from secret_fields in original_request_params raw_headers_from_request: dict[str, str] | None = None @@ -768,6 +767,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): return from litellm.responses.main import aresponses + from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, ) diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index a1039a20489..429adb2962d 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -1,20 +1,20 @@ +import importlib import subprocess import sys import textwrap import types +from typing import Any, cast from unittest.mock import AsyncMock, MagicMock import pytest from fastapi import HTTPException -import importlib - from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing +from litellm.types.responses.main import OutputFunctionToolCall +from litellm.types.utils import ModelResponse + from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, ) -from typing import Any, cast -from litellm.types.utils import ModelResponse -from litellm.types.responses.main import OutputFunctionToolCall class _DummyMCPResult: @@ -105,9 +105,7 @@ def test_extract_tool_calls_from_chat_response_handles_tool_calls(): object="chat.completion", ) - tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response( - response - ) + tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(response) assert len(tool_calls) == 1 assert tool_calls[0]["function"]["name"] == "foo" @@ -177,9 +175,7 @@ def test_transform_mcp_tools_to_openai_uses_chat_format(monkeypatch): fake_transform_responses, ) - chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( - ["tool"], target_format="chat" - ) + chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"], target_format="chat") resp_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"]) assert chat_tools == [{"chat": True}] @@ -299,9 +295,7 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n ) from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm - _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock( - return_value=fake_server - ) + _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=fake_server) tool_name = "my_deepwiki-read_wiki_structure" tool_calls = [ @@ -375,7 +369,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey fake_manager = types.SimpleNamespace( get_registry=MagicMock(return_value={}), - call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")) + call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")), ) monkeypatch.setattr( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -383,9 +377,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey ) tool_name = "deepwiki-read_wiki_structure" - tool_calls = [ - {"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}} - ] + tool_calls = [{"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}}] user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user") @@ -403,10 +395,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey post_call_failure_hook.assert_awaited_once() assert post_call_failure_hook.await_args is not None - assert ( - post_call_failure_hook.await_args.kwargs.get("route") - == "/responses/mcp/call_tool" - ) + assert post_call_failure_hook.await_args.kwargs.get("route") == "/responses/mcp/call_tool" @pytest.mark.asyncio @@ -429,9 +418,7 @@ async def test_execute_tool_calls_uses_unique_call_ids_and_preserves_parent_cont # NOTE: Don't patch via dotted string path here because `litellm.responses` # is a function attribute on the `litellm` package (shadowing the submodule), # which breaks monkeypatch's importpath resolution. - handler_module = importlib.import_module( - "litellm.responses.mcp.litellm_proxy_mcp_handler" - ) + handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure" @@ -544,9 +531,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user") tools, _server_names = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( user_api_key_auth=user_auth, - mcp_tools_with_litellm_proxy=[ - {"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"} - ], + mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}], ) assert tools == [] @@ -557,9 +542,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch def test_get_parent_request_tags_from_metadata(): - tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags( - {"metadata": {"tags": ["team-a", "prod"]}} - ) + tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags({"metadata": {"tags": ["team-a", "prod"]}}) assert tags == ["team-a", "prod"] @@ -596,9 +579,7 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch): await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( user_api_key_auth=types.SimpleNamespace(api_key="k", user_id="u"), - mcp_tools_with_litellm_proxy=[ - {"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"} - ], + mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}], request_tags=["team-a"], ) @@ -618,9 +599,7 @@ async def test_execute_tool_calls_exposes_sanitized_client_headers_to_logging(mo captured.update(kwargs) return None, None - handler_module = importlib.import_module( - "litellm.responses.mcp.litellm_proxy_mcp_handler" - ) + handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure" @@ -646,9 +625,7 @@ async def test_execute_tool_calls_propagates_request_tags_to_function_setup(monk captured.update(kwargs) return None, None - handler_module = importlib.import_module( - "litellm.responses.mcp.litellm_proxy_mcp_handler" - ) + handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure"