diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 52117d86066..4e5c73be806 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -807,9 +807,9 @@ if MCP_AVAILABLE: "user_api_key" ] = user_api_key - user_identifier = getattr(user_api_key_auth, "end_user_id", None) or getattr( - user_api_key_auth, "user_id", None - ) + user_identifier = getattr( + user_api_key_auth, "end_user_id", None + ) or getattr(user_api_key_auth, "user_id", None) if user_identifier: list_tools_request_data["user"] = user_identifier @@ -838,7 +838,9 @@ if MCP_AVAILABLE: # Decide whether to add prefix based on number of allowed servers add_prefix = not (len(allowed_mcp_servers) == 1) - async def _fetch_and_filter_server_tools(server: MCPServer) -> List[MCPTool]: + async def _fetch_and_filter_server_tools( + server: MCPServer, + ) -> List[MCPTool]: """Fetch and filter tools from a single server with error handling.""" if server is None: return [] @@ -1517,11 +1519,9 @@ if MCP_AVAILABLE: **kwargs, ) except Exception as e: - traceback_str = traceback.format_exc( - limit=MAXIMUM_TRACEBACK_LINES_TO_LOG - ) + traceback_str = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) from litellm.proxy.proxy_server import proxy_logging_obj - + if proxy_logging_obj and user_api_key_auth: await proxy_logging_obj.post_call_failure_hook( request_data=kwargs, diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 78d358a1e39..83c23a58500 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -169,7 +169,9 @@ async def aresponses_api_with_mcp( # Process MCP tools through the complete pipeline (fetch + filter + deduplicate + transform) # Extract user_api_key_auth from litellm_metadata (where it's added by add_user_api_key_auth_to_request_metadata) - user_api_key_auth = kwargs.get("user_api_key_auth") or kwargs.get("litellm_metadata", {}).get("user_api_key_auth") + user_api_key_auth = kwargs.get("user_api_key_auth") or kwargs.get( + "litellm_metadata", {} + ).get("user_api_key_auth") # Get original MCP tools (for events) and OpenAI tools (for LLM) by reusing existing methods ( @@ -280,7 +282,7 @@ async def aresponses_api_with_mcp( user_api_key_auth = kwargs.get("litellm_metadata", {}).get( "user_api_key_auth" ) - + # Extract MCP auth headers from the request to pass to MCP server secret_fields: Optional[Dict[str, Any]] = kwargs.get("secret_fields") ( @@ -292,7 +294,7 @@ async def aresponses_api_with_mcp( secret_fields=secret_fields, tools=tools, ) - + tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=tool_server_map, tool_calls=tool_calls, @@ -590,9 +592,12 @@ def responses( ######################################################### # Update input with provider-specific file IDs if managed files are used ######################################################### - input = cast(Union[str, ResponseInputParam], update_responses_input_with_model_file_ids(input=input)) + input = cast( + Union[str, ResponseInputParam], + update_responses_input_with_model_file_ids(input=input), + ) local_vars["input"] = input - + ######################################################### # Native MCP Responses API ######################################################### @@ -627,11 +632,11 @@ def responses( ) # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=model, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=model, + provider=litellm.LlmProviders(custom_llm_provider), ) local_vars.update(kwargs) @@ -826,11 +831,11 @@ def delete_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=None, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=None, + provider=litellm.LlmProviders(custom_llm_provider), ) if responses_api_provider_config is None: @@ -1006,11 +1011,11 @@ def get_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=None, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=None, + provider=litellm.LlmProviders(custom_llm_provider), ) if responses_api_provider_config is None: @@ -1163,11 +1168,11 @@ def list_input_items( if custom_llm_provider is None: raise ValueError("custom_llm_provider is required but passed as None") - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=None, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=None, + provider=litellm.LlmProviders(custom_llm_provider), ) if responses_api_provider_config is None: @@ -1321,11 +1326,11 @@ def cancel_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=None, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=None, + provider=litellm.LlmProviders(custom_llm_provider), ) if responses_api_provider_config is None: @@ -1503,11 +1508,11 @@ def compact_responses( raise ValueError("custom_llm_provider is required but passed as None") # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=model, - provider=litellm.LlmProviders(custom_llm_provider), - ) + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=model, + provider=litellm.LlmProviders(custom_llm_provider), ) if responses_api_provider_config is None: diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index e22898ae0ff..7cc821f5e02 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -19,7 +19,12 @@ from litellm.proxy._experimental.mcp_server.utils import split_server_prefix_fro from litellm.responses.main import aresponses from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator from litellm.types.llms.openai import ResponsesAPIResponse -from litellm.types.utils import CallTypes, Choices, ModelResponse, StandardLoggingMCPToolCall +from litellm.types.utils import ( + CallTypes, + Choices, + ModelResponse, + StandardLoggingMCPToolCall, +) from litellm.utils import Rules, function_setup if TYPE_CHECKING: @@ -564,9 +569,9 @@ class LiteLLM_Proxy_MCP_Handler: if user_api_key: logging_request_data["metadata"]["user_api_key"] = user_api_key - user_identifier = getattr(user_api_key_auth, "end_user_id", None) or getattr( - user_api_key_auth, "user_id", None - ) + user_identifier = getattr( + user_api_key_auth, "end_user_id", None + ) or getattr(user_api_key_auth, "user_id", None) if user_identifier: logging_request_data["user"] = user_identifier @@ -620,7 +625,9 @@ class LiteLLM_Proxy_MCP_Handler: standard_logging_mcp_tool_call["mcp_server_logo_url"] = logo_url cost_info = mcp_info.get("mcp_server_cost_info") if cost_info: - standard_logging_mcp_tool_call["mcp_server_cost_info"] = cost_info + standard_logging_mcp_tool_call[ + "mcp_server_cost_info" + ] = cost_info if litellm_logging_obj: litellm_logging_obj.model_call_details[ @@ -895,9 +902,7 @@ class LiteLLM_Proxy_MCP_Handler: return try: - traceback_str = traceback.format_exc( - limit=MAXIMUM_TRACEBACK_LINES_TO_LOG - ) + traceback_str = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) await proxy_logging_obj.post_call_failure_hook( request_data=request_data, original_exception=error, @@ -948,7 +953,8 @@ class LiteLLM_Proxy_MCP_Handler: mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events tool_server_map=tool_server_map, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, - user_api_key_auth=kwargs.get("user_api_key_auth") or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"), + user_api_key_auth=kwargs.get("user_api_key_auth") + or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"), original_request_params=request_params, ) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index 53f39164e91..731aa5c692b 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -273,9 +273,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.finished = False # Event queues and generation flags - self.mcp_discovery_events: List[ResponsesAPIStreamingResponse] = ( - mcp_events # Pre-generated MCP discovery events - ) + self.mcp_discovery_events: List[ + ResponsesAPIStreamingResponse + ] = mcp_events # Pre-generated MCP discovery events self.tool_execution_events: List[ResponsesAPIStreamingResponse] = [] self.mcp_discovery_generated = True # Events are already generated self.mcp_events = ( @@ -284,9 +284,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.tool_server_map = tool_server_map # Iterator references - self.base_iterator: Optional[Union[Any, ResponsesAPIResponse]] = ( - base_iterator # Will be created when needed - ) + self.base_iterator: Optional[ + Union[Any, ResponsesAPIResponse] + ] = base_iterator # Will be created when needed self.follow_up_iterator: Optional[Any] = None # Response collection for tool execution @@ -305,7 +305,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): # Mark as async iterator self.is_async = True - + def _extract_mcp_headers_from_params(self) -> None: """Extract MCP headers from original request params to pass to tool calls""" from typing import Dict, Optional @@ -313,25 +313,31 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) - + # Extract headers from secret_fields in original_request_params raw_headers_from_request: Optional[Dict[str, str]] = None secret_fields = self.original_request_params.get("secret_fields") if secret_fields and isinstance(secret_fields, dict): raw_headers_from_request = secret_fields.get("raw_headers") - + # Extract MCP-specific headers self.mcp_auth_header: Optional[str] = None self.mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None self.oauth2_headers: Optional[Dict[str, str]] = None self.raw_headers: Optional[Dict[str, str]] = raw_headers_from_request - + if raw_headers_from_request: headers_obj = Headers(raw_headers_from_request) - self.mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers_obj) - self.mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers_obj) - self.oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers_obj) - + self.mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers( + headers_obj + ) + self.mcp_server_auth_headers = ( + MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers_obj) + ) + self.oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers( + headers_obj + ) + # Also check if headers are provided in tools array (from request body) tools = self.original_request_params.get("tools") if tools: @@ -341,17 +347,26 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): if tool_headers and isinstance(tool_headers, dict): # Merge tool headers into mcp_server_auth_headers headers_obj_from_tool = Headers(tool_headers) - tool_mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers_obj_from_tool) - + tool_mcp_server_auth_headers = ( + MCPRequestHandler._get_mcp_server_auth_headers_from_headers( + headers_obj_from_tool + ) + ) + if tool_mcp_server_auth_headers: if self.mcp_server_auth_headers is None: self.mcp_server_auth_headers = {} # Merge the headers from tool into existing headers - for server_alias, headers_dict in tool_mcp_server_auth_headers.items(): + for ( + server_alias, + headers_dict, + ) in tool_mcp_server_auth_headers.items(): if server_alias not in self.mcp_server_auth_headers: self.mcp_server_auth_headers[server_alias] = {} - self.mcp_server_auth_headers[server_alias].update(headers_dict) - + self.mcp_server_auth_headers[server_alias].update( + headers_dict + ) + # Also merge raw headers if self.raw_headers is None: self.raw_headers = {} @@ -489,9 +504,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): # Use the pre-fetched all_tools from original_request_params (no re-processing needed) params_for_llm = {} for key, value in params.items(): - params_for_llm[key] = ( - value # Copy all params as-is since tools are already processed - ) + params_for_llm[ + key + ] = value # Copy all params as-is since tools are already processed tools_count = ( len(params_for_llm.get("tools", [])) @@ -545,9 +560,11 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): return for tool_call in tool_calls: - tool_name, tool_arguments, tool_call_id = ( - LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) - ) + ( + tool_name, + tool_arguments, + tool_call_id, + ) = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) if tool_name and tool_call_id: # Create MCP call events for this tool execution call_events = create_mcp_call_events( diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 324380db2b4..cac2fe85541 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -63,11 +63,11 @@ def _generate_id(): # private helper function return "chatcmpl-" + str(uuid.uuid4()) - class SafeAttributeModel: """ A base model that provides safe attribute access. """ + def __delattr__(self, name): try: super().__delattr__(name) @@ -125,13 +125,14 @@ class SearchContextCostPerQuery(TypedDict, total=False): class AgenticLoopParams(TypedDict, total=False): """ Parameters passed to agentic loop hooks (e.g., WebSearch interception). - + Stored in logging_obj.model_call_details["agentic_loop_params"] to provide agentic hooks with the original request context needed for follow-up calls. """ + model: str """The model string with provider prefix (e.g., 'bedrock/invoke/...')""" - + custom_llm_provider: str """The LLM provider name (e.g., 'bedrock', 'anthropic')""" @@ -1345,8 +1346,7 @@ class CacheCreationTokenDetails(BaseModel): class PromptTokensDetailsWrapper( - SafeAttributeModel, - PromptTokensDetails + SafeAttributeModel, PromptTokensDetails ): # extends with image generation fields (text_tokens, image_tokens) text_tokens: Optional[int] = None """Text tokens sent to the model."""