diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 97e6410447b..50579124c2f 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -120,7 +120,9 @@ def _warn_on_server_name_fields( if result.is_valid: return - warning_text = "; ".join(result.warnings) if result.warnings else "Validation failed" + warning_text = ( + "; ".join(result.warnings) if result.warnings else "Validation failed" + ) verbose_logger.warning( "MCP server '%s' has invalid %s '%s': %s", server_id, @@ -171,7 +173,9 @@ def _create_sampling_callback(user_api_key_auth: Optional[Any] = None): ) auth_context = get_active_auth_context() - resolved_auth = user_api_key_auth or (auth_context.user_api_key_auth if auth_context else None) + resolved_auth = user_api_key_auth or ( + auth_context.user_api_key_auth if auth_context else None + ) return await handle_sampling_create_message( context=context, @@ -200,7 +204,11 @@ def _create_elicitation_callback(): # In Gateway mode, we relay the elicitation request to the downstream client # that triggered the current operation. downstream_session = get_active_mcp_session() - downstream_capabilities = getattr(downstream_session, "capabilities", None) if downstream_session else None + downstream_capabilities = ( + getattr(downstream_session, "capabilities", None) + if downstream_session + else None + ) return await handle_elicitation_request( context=context, @@ -242,10 +250,14 @@ class MCPServerManager: """ self._upstream_initialize_instructions_by_server_id: Dict[str, str] = {} - def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None: + def _remember_upstream_initialize_instructions( + self, server: MCPServer, client: MCPClient + ) -> None: raw = getattr(client, "_last_initialize_instructions", None) if raw and str(raw).strip(): - self._upstream_initialize_instructions_by_server_id[server.server_id] = str(raw).strip() + self._upstream_initialize_instructions_by_server_id[server.server_id] = str( + raw + ).strip() def get_registry(self) -> Dict[str, MCPServer]: """ @@ -289,10 +301,15 @@ class MCPServerManager: if mcp_aliases and alias is None: # Check if this server_name has an alias in mcp_aliases for alias_name, target_server_name in mcp_aliases.items(): - if target_server_name == server_name and alias_name not in used_aliases: + if ( + target_server_name == server_name + and alias_name not in used_aliases + ): alias = alias_name used_aliases.add(alias_name) - verbose_logger.debug(f"Mapped alias '{alias_name}' to server '{server_name}'") + verbose_logger.debug( + f"Mapped alias '{alias_name}' to server '{server_name}'" + ) break # Create a temporary server object to use with get_server_prefix utility @@ -310,10 +327,15 @@ class MCPServerManager: if mcp_aliases and alias is None: # Check if this server_name has an alias in mcp_aliases for alias_name, target_server_name in mcp_aliases.items(): - if target_server_name == server_name and alias_name not in used_aliases: + if ( + target_server_name == server_name + and alias_name not in used_aliases + ): alias = alias_name used_aliases.add(alias_name) - verbose_logger.debug(f"Mapped alias '{alias_name}' to server '{server_name}'") + verbose_logger.debug( + f"Mapped alias '{alias_name}' to server '{server_name}'" + ) break # Create a temporary server object to use with get_server_prefix utility @@ -348,7 +370,9 @@ class MCPServerManager: else: mcp_oauth_metadata = None - resolved_scopes = server_config.get("scopes") or (mcp_oauth_metadata.scopes if mcp_oauth_metadata else None) + resolved_scopes = server_config.get("scopes") or ( + mcp_oauth_metadata.scopes if mcp_oauth_metadata else None + ) resolved_authorization_url = server_config.get("authorization_url") or ( mcp_oauth_metadata.authorization_url if mcp_oauth_metadata else None ) @@ -380,7 +404,9 @@ class MCPServerManager: # TODO: utility fn the default values transport=server_config.get("transport", MCPTransport.http), auth_type=auth_type, - authentication_token=server_config.get("authentication_token", server_config.get("auth_value", None)), + authentication_token=server_config.get( + "authentication_token", server_config.get("auth_value", None) + ), mcp_info=mcp_info, extra_headers=server_config.get("extra_headers", None), allowed_tools=server_config.get("allowed_tools", None), @@ -389,7 +415,9 @@ class MCPServerManager: access_groups=server_config.get("access_groups", None), static_headers=server_config.get("static_headers", None), allow_all_keys=bool(server_config.get("allow_all_keys", False)), - available_on_public_internet=bool(server_config.get("available_on_public_internet", True)), + available_on_public_internet=bool( + server_config.get("available_on_public_internet", True) + ), # AWS SigV4 fields aws_access_key_id=server_config.get("aws_access_key_id", None), aws_secret_access_key=server_config.get("aws_secret_access_key", None), @@ -405,18 +433,24 @@ class MCPServerManager: # Check if this is an OpenAPI-based server spec_path = server_config.get("spec_path", None) if spec_path: - verbose_logger.info(f"Loading OpenAPI spec from {spec_path} for server {server_name}") + verbose_logger.info( + f"Loading OpenAPI spec from {spec_path} for server {server_name}" + ) await self._register_openapi_tools( spec_path=spec_path, server=new_server, base_url=server_config.get("url", ""), ) - verbose_logger.debug(f"Loaded MCP Servers: {json.dumps(self.config_mcp_servers, indent=4, default=str)}") + verbose_logger.debug( + f"Loaded MCP Servers: {json.dumps(self.config_mcp_servers, indent=4, default=str)}" + ) self.initialize_tool_name_to_mcp_server_name_mapping() - async def _register_openapi_tools(self, spec_path: str, server: MCPServer, base_url: str): + async def _register_openapi_tools( + self, spec_path: str, server: MCPServer, base_url: str + ): """ Register tools from an OpenAPI specification for a given server. @@ -452,7 +486,9 @@ class MCPServerManager: # Use base_url from config if provided, otherwise extract from spec if not base_url: base_url = get_openapi_base_url(spec, spec_path) - verbose_logger.info(f"Registering OpenAPI tools for server {server.name} with base URL: {base_url}") + verbose_logger.info( + f"Registering OpenAPI tools for server {server.name} with base URL: {base_url}" + ) # Get server prefix for tool naming server_prefix = get_server_prefix(server) @@ -505,14 +541,20 @@ class MCPServerManager: operation = path_item[method] # Resolve $ref params and merge path-level params into the operation. - resolved_operation = resolve_operation_params(operation, path_item, components) + resolved_operation = resolve_operation_params( + operation, path_item, components + ) # Generate tool name (without prefix initially) - operation_id = operation.get("operationId", f"{method}_{path.replace('/', '_')}") + operation_id = operation.get( + "operationId", f"{method}_{path.replace('/', '_')}" + ) base_tool_name = operation_id.replace(" ", "_").lower() # Add server prefix to tool name - prefixed_tool_name = add_server_prefix_to_name(base_tool_name, server_prefix) + prefixed_tool_name = add_server_prefix_to_name( + base_tool_name, server_prefix + ) # Get description description = operation.get( @@ -524,7 +566,9 @@ class MCPServerManager: input_schema = build_input_schema(resolved_operation) # Create tool function with headers using imported function - tool_func = create_tool_function(path, method, resolved_operation, base_url, headers=headers) + tool_func = create_tool_function( + path, method, resolved_operation, base_url, headers=headers + ) tool_func.__name__ = prefixed_tool_name tool_func.__doc__ = description @@ -537,16 +581,26 @@ class MCPServerManager: ) # Update tool name to server name mapping (for both prefixed and base names) - self.tool_name_to_mcp_server_name_mapping[base_tool_name] = server_prefix - self.tool_name_to_mcp_server_name_mapping[prefixed_tool_name] = server_prefix + self.tool_name_to_mcp_server_name_mapping[base_tool_name] = ( + server_prefix + ) + self.tool_name_to_mcp_server_name_mapping[prefixed_tool_name] = ( + server_prefix + ) registered_count += 1 - verbose_logger.debug(f"Registered OpenAPI tool: {prefixed_tool_name} for server {server.name}") + verbose_logger.debug( + f"Registered OpenAPI tool: {prefixed_tool_name} for server {server.name}" + ) - verbose_logger.info(f"Successfully registered {registered_count} OpenAPI tools for server {server.name}") + verbose_logger.info( + f"Successfully registered {registered_count} OpenAPI tools for server {server.name}" + ) except Exception as e: - verbose_logger.error(f"Failed to register OpenAPI tools for server {server.name}: {str(e)}") + verbose_logger.error( + f"Failed to register OpenAPI tools for server {server.name}: {str(e)}" + ) raise e def remove_server(self, mcp_server: LiteLLM_MCPServerTable): @@ -560,7 +614,9 @@ class MCPServerManager: del self.registry[mcp_server.server_id] verbose_logger.debug(f"Removed MCP Server: {mcp_server.server_id}") else: - verbose_logger.warning(f"Server ID {mcp_server.server_id} not found in registry") + verbose_logger.warning( + f"Server ID {mcp_server.server_id} not found in registry" + ) async def build_mcp_server_from_table( self, @@ -570,8 +626,12 @@ class MCPServerManager: ) -> MCPServer: _mcp_info: MCPInfo = mcp_server.mcp_info or {} env_dict = _deserialize_json_dict(getattr(mcp_server, "env", None)) - static_headers_dict = _deserialize_json_dict(getattr(mcp_server, "static_headers", None)) - credentials_dict = _deserialize_json_dict(getattr(mcp_server, "credentials", None)) + static_headers_dict = _deserialize_json_dict( + getattr(mcp_server, "static_headers", None) + ) + credentials_dict = _deserialize_json_dict( + getattr(mcp_server, "credentials", None) + ) encrypted_auth_value: Optional[str] = None encrypted_client_id: Optional[str] = None @@ -618,7 +678,9 @@ class MCPServerManager: client_secret_value = encrypted_client_secret # AWS SigV4 credential fields - aws_creds = self._extract_aws_credentials(credentials_dict, credentials_are_encrypted) + aws_creds = self._extract_aws_credentials( + credentials_dict, credentials_are_encrypted + ) scopes: Optional[List[str]] = None if credentials_dict: @@ -626,7 +688,9 @@ class MCPServerManager: if scopes_value is not None: scopes = self._extract_scopes(scopes_value) - name_for_prefix = mcp_server.alias or mcp_server.server_name or mcp_server.server_id + name_for_prefix = ( + mcp_server.alias or mcp_server.server_name or mcp_server.server_id + ) mcp_info: MCPInfo = _mcp_info.copy() if "server_name" not in mcp_info: @@ -636,14 +700,20 @@ class MCPServerManager: auth_type = cast(MCPAuthType, mcp_server.auth_type) server_url = mcp_server.url - needs_discovery = bool(server_url) and auth_type == MCPAuth.oauth2 and not mcp_server.authorization_url + needs_discovery = ( + bool(server_url) + and auth_type == MCPAuth.oauth2 + and not mcp_server.authorization_url + ) mcp_oauth_metadata = ( await self._descovery_metadata(server_url=server_url) # type: ignore[arg-type] if needs_discovery else None ) - resolved_scopes = scopes or (mcp_oauth_metadata.scopes if mcp_oauth_metadata else None) + resolved_scopes = scopes or ( + mcp_oauth_metadata.scopes if mcp_oauth_metadata else None + ) new_server = MCPServer( server_id=mcp_server.server_id, @@ -659,12 +729,16 @@ class MCPServerManager: extra_headers=getattr(mcp_server, "extra_headers", None), static_headers=static_headers_dict, client_id=client_id_value or getattr(mcp_server, "client_id", None), - client_secret=client_secret_value or getattr(mcp_server, "client_secret", None), + client_secret=client_secret_value + or getattr(mcp_server, "client_secret", None), oauth2_flow=getattr(mcp_server, "oauth2_flow", None), scopes=resolved_scopes, - authorization_url=mcp_server.authorization_url or getattr(mcp_oauth_metadata, "authorization_url", None), - token_url=mcp_server.token_url or getattr(mcp_oauth_metadata, "token_url", None), - registration_url=mcp_server.registration_url or getattr(mcp_oauth_metadata, "registration_url", None), + authorization_url=mcp_server.authorization_url + or getattr(mcp_oauth_metadata, "authorization_url", None), + token_url=mcp_server.token_url + or getattr(mcp_oauth_metadata, "token_url", None), + registration_url=mcp_server.registration_url + or getattr(mcp_oauth_metadata, "registration_url", None), command=getattr(mcp_server, "command", None), args=getattr(mcp_server, "args", None) or [], env=env_dict, @@ -672,11 +746,17 @@ class MCPServerManager: allowed_tools=getattr(mcp_server, "allowed_tools", None), disallowed_tools=getattr(mcp_server, "disallowed_tools", None), allow_all_keys=mcp_server.allow_all_keys, - available_on_public_internet=bool(getattr(mcp_server, "available_on_public_internet", True)), + available_on_public_internet=bool( + getattr(mcp_server, "available_on_public_internet", True) + ), created_at=getattr(mcp_server, "created_at", None), updated_at=getattr(mcp_server, "updated_at", None), - tool_name_to_display_name=_deserialize_json_dict(getattr(mcp_server, "tool_name_to_display_name", None)), - tool_name_to_description=_deserialize_json_dict(getattr(mcp_server, "tool_name_to_description", None)), + tool_name_to_display_name=_deserialize_json_dict( + getattr(mcp_server, "tool_name_to_display_name", None) + ), + tool_name_to_description=_deserialize_json_dict( + getattr(mcp_server, "tool_name_to_description", None) + ), is_byok=bool(getattr(mcp_server, "is_byok", False)), byok_description=getattr(mcp_server, "byok_description", None) or [], byok_api_key_help_url=getattr(mcp_server, "byok_api_key_help_url", None), @@ -695,7 +775,9 @@ class MCPServerManager: async def _maybe_register_openapi_tools(self, server: MCPServer): """Register OpenAPI tools if the server has a spec_path configured.""" if server.spec_path: - verbose_logger.info(f"Loading OpenAPI spec from {server.spec_path} for server {server.name}") + verbose_logger.info( + f"Loading OpenAPI spec from {server.spec_path} for server {server.name}" + ) await self._register_openapi_tools( spec_path=server.spec_path, server=server, @@ -736,9 +818,15 @@ class MCPServerManager: def get_allow_all_keys_server_ids(self) -> List[str]: """Return server IDs that bypass per-key restrictions.""" - return [server.server_id for server in self.get_registry().values() if server.allow_all_keys is True] + return [ + server.server_id + for server in self.get_registry().values() + if server.allow_all_keys is True + ] - async def get_allowed_mcp_servers(self, user_api_key_auth: Optional[UserAPIKeyAuth] = None) -> List[str]: + async def get_allowed_mcp_servers( + self, user_api_key_auth: Optional[UserAPIKeyAuth] = None + ) -> List[str]: """ Get the allowed MCP Servers for the user. @@ -763,13 +851,23 @@ class MCPServerManager: ) # If admin but NO explicit object permission, get all servers - if user_api_key_auth and _user_has_admin_view(user_api_key_auth) and not has_explicit_object_permission: - verbose_logger.debug("Admin user without explicit object_permission - returning all servers") + if ( + user_api_key_auth + and _user_has_admin_view(user_api_key_auth) + and not has_explicit_object_permission + ): + verbose_logger.debug( + "Admin user without explicit object_permission - returning all servers" + ) return list(self.get_registry().keys()) # Get allowed servers from object permissions (respects object_permission even for admins) - allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth) - verbose_logger.debug(f"Allowed MCP Servers for user api key auth: {allowed_mcp_servers}") + allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers( + user_api_key_auth + ) + verbose_logger.debug( + f"Allowed MCP Servers for user api key auth: {allowed_mcp_servers}" + ) combined_servers = set(allowed_mcp_servers) # Only skip allow_all_keys servers when the request is inside a toolset # scope. toolset_mcp_route / dynamic_mcp_route set _mcp_active_toolset_id @@ -785,7 +883,9 @@ class MCPServerManager: combined_servers.update(allow_all_server_ids) if len(combined_servers) == 0: - verbose_logger.debug("No allowed MCP Servers found for user api key auth.") + verbose_logger.debug( + "No allowed MCP Servers found for user api key auth." + ) return list(combined_servers) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}.") @@ -862,12 +962,15 @@ class MCPServerManager: keys_to_remove = [ k for k in cache_dict - if (k.startswith("toolset_perms:") and toolset_id in k) or k.startswith("toolset_name:") + if (k.startswith("toolset_perms:") and toolset_id in k) + or k.startswith("toolset_name:") ] for k in keys_to_remove: cache_dict.pop(k, None) except Exception as e: - verbose_logger.warning(f"invalidate_toolset_cache: failed to evict in-memory entries: {e}") + verbose_logger.warning( + f"invalidate_toolset_cache: failed to evict in-memory entries: {e}" + ) async def get_toolset_by_name_cached( self, @@ -905,12 +1008,18 @@ class MCPServerManager: toolset = await get_mcp_toolset_by_name(prisma_client, toolset_name) await user_api_key_cache.async_set_cache( key=cache_key, - value=(toolset.model_dump(mode="json") if toolset is not None else "__not_found__"), + value=( + toolset.model_dump(mode="json") + if toolset is not None + else "__not_found__" + ), ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) return toolset - def filter_server_ids_by_ip(self, server_ids: List[str], client_ip: Optional[str]) -> List[str]: + def filter_server_ids_by_ip( + self, server_ids: List[str], client_ip: Optional[str] + ) -> List[str]: """ Filter server IDs by client IP — external callers only see public servers. @@ -952,7 +1061,9 @@ class MCPServerManager: return [] return await self._get_tools_from_server(server) except Exception as e: - verbose_logger.warning(f"Failed to get tools from server {server_id}: {str(e)}") + verbose_logger.warning( + f"Failed to get tools from server {server_id}: {str(e)}" + ) return [] async def list_tools( @@ -1014,7 +1125,9 @@ class MCPServerManager: # Flatten results into single list list_tools_result: List[MCPTool] = [tool for tools in results for tool in tools] - verbose_logger.info(f"Successfully fetched {len(list_tools_result)} tools total from all servers") + verbose_logger.info( + f"Successfully fetched {len(list_tools_result)} tools total from all servers" + ) return list_tools_result ######################################################### @@ -1083,7 +1196,9 @@ class MCPServerManager: # Handle stdio transport if transport == MCPTransport.stdio: resolved_env = ( - stdio_env if stdio_env is not None else (dict(server.env) if server.env is not None else None) + stdio_env + if stdio_env is not None + else (dict(server.env) if server.env is not None else None) ) # Ensure npm-based STDIO MCP servers have a writable cache dir. @@ -1201,7 +1316,9 @@ class MCPServerManager: ## HANDLE OPENAPI TOOLS if server.spec_path: _tools = global_mcp_tool_registry.list_tools(tool_prefix=server.name) - tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type(_tools) + tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type( + _tools + ) # OpenAPI tools are stored in the registry with their prefix already # applied (e.g. "test_petstore-getinventory"). Do NOT pass them # through _create_prefixed_tools — that would add the prefix a second @@ -1211,7 +1328,9 @@ class MCPServerManager: sep = MCP_TOOL_PREFIX_SEPARATOR tools = [ ( - t.model_copy(update={"name": t.name[len(prefix) + len(sep) :]}) + t.model_copy( + update={"name": t.name[len(prefix) + len(sep) :]} + ) if t.name.startswith(f"{prefix}{sep}") else t ) @@ -1222,12 +1341,16 @@ class MCPServerManager: tools = await self._fetch_tools_with_timeout(client, server.name) self._remember_upstream_initialize_instructions(server, client) - prefixed_or_original_tools = self._create_prefixed_tools(tools, server, add_prefix=add_prefix) + prefixed_or_original_tools = self._create_prefixed_tools( + tools, server, add_prefix=add_prefix + ) return prefixed_or_original_tools except Exception as e: - verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}") + verbose_logger.warning( + f"Failed to get tools from server {server.name}: {str(e)}" + ) return [] async def get_prompts_from_server( @@ -1271,12 +1394,16 @@ class MCPServerManager: prompts = await client.list_prompts() - prefixed_or_original_prompts = self._create_prefixed_prompts(prompts, server, add_prefix=add_prefix) + prefixed_or_original_prompts = self._create_prefixed_prompts( + prompts, server, add_prefix=add_prefix + ) return prefixed_or_original_prompts except Exception as e: - verbose_logger.warning(f"Failed to get prompts from server {server.name}: {str(e)}") + verbose_logger.warning( + f"Failed to get prompts from server {server.name}: {str(e)}" + ) return [] async def get_resources_from_server( @@ -1311,12 +1438,16 @@ class MCPServerManager: resources = await client.list_resources() - prefixed_resources = self._create_prefixed_resources(resources, server, add_prefix=add_prefix) + prefixed_resources = self._create_prefixed_resources( + resources, server, add_prefix=add_prefix + ) return prefixed_resources except Exception as e: - verbose_logger.warning(f"Failed to get resources from server {server.name}: {str(e)}") + verbose_logger.warning( + f"Failed to get resources from server {server.name}: {str(e)}" + ) return [] async def get_resource_templates_from_server( @@ -1358,7 +1489,9 @@ class MCPServerManager: return prefixed_templates except Exception as e: - verbose_logger.warning(f"Failed to get resource templates from server {server.name}: {str(e)}") + verbose_logger.warning( + f"Failed to get resource templates from server {server.name}: {str(e)}" + ) return [] async def read_resource_from_server( @@ -1448,11 +1581,13 @@ class MCPServerManager: header_value: Optional[str] = None if exc.response is not None: - header_value = exc.response.headers.get("WWW-Authenticate") or exc.response.headers.get( - "www-authenticate" - ) + header_value = exc.response.headers.get( + "WWW-Authenticate" + ) or exc.response.headers.get("www-authenticate") - resource_metadata_url, scopes = self._parse_www_authenticate_header(header_value) + resource_metadata_url, scopes = self._parse_www_authenticate_header( + header_value + ) authorization_servers: List[str] = [] resource_scopes: Optional[List[str]] = None @@ -1460,7 +1595,9 @@ class MCPServerManager: ( authorization_servers, resource_scopes, - ) = await self._fetch_oauth_metadata_from_resource(resource_metadata_url) + ) = await self._fetch_oauth_metadata_from_resource( + resource_metadata_url + ) else: ( authorization_servers, @@ -1472,12 +1609,16 @@ class MCPServerManager: try: parsed_url = urlparse(server_url) if parsed_url.scheme and parsed_url.netloc: - authorization_servers = [f"{parsed_url.scheme}://{parsed_url.netloc}"] + authorization_servers = [ + f"{parsed_url.scheme}://{parsed_url.netloc}" + ] except Exception: authorization_servers = [] if authorization_servers: - metadata = await self._fetch_authorization_server_metadata(authorization_servers) + metadata = await self._fetch_authorization_server_metadata( + authorization_servers + ) preferred_scopes = scopes or resource_scopes if metadata is None and preferred_scopes: @@ -1487,10 +1628,14 @@ class MCPServerManager: return metadata except Exception as exc: # pragma: no cover - network/transient issues - verbose_logger.debug("MCP OAuth discovery failed for %s: %s", server_url, exc) + verbose_logger.debug( + "MCP OAuth discovery failed for %s: %s", server_url, exc + ) return None - def _parse_www_authenticate_header(self, header_value: Optional[str]) -> Tuple[Optional[str], Optional[List[str]]]: + def _parse_www_authenticate_header( + self, header_value: Optional[str] + ) -> Tuple[Optional[str], Optional[List[str]]]: if not header_value: return None, None @@ -1499,7 +1644,8 @@ class MCPServerManager: param_pattern = re.compile(r"([a-zA-Z0-9_]+)\s*=\s*\"?([^\",]+)\"?") params: Dict[str, str] = { - match.group(1).lower(): match.group(2).strip() for match in param_pattern.finditer(params_section) + match.group(1).lower(): match.group(2).strip() + for match in param_pattern.finditer(params_section) } resource_metadata_url = params.get("resource_metadata") @@ -1534,15 +1680,23 @@ class MCPServerManager: raw_servers = data.get("authorization_servers") if isinstance(raw_servers, list): - authorization_servers = [entry for entry in raw_servers if isinstance(entry, str) and entry.strip() != ""] + authorization_servers = [ + entry + for entry in raw_servers + if isinstance(entry, str) and entry.strip() != "" + ] else: authorization_servers = [] - scopes = self._extract_scopes(data.get("scopes_supported") or data.get("scopes")) + scopes = self._extract_scopes( + data.get("scopes_supported") or data.get("scopes") + ) return authorization_servers, scopes - async def _attempt_well_known_discovery(self, server_url: str) -> Tuple[List[str], Optional[List[str]]]: + async def _attempt_well_known_discovery( + self, server_url: str + ) -> Tuple[List[str], Optional[List[str]]]: try: parsed = urlparse(server_url) except Exception: @@ -1579,7 +1733,9 @@ class MCPServerManager: return metadata return None - async def _fetch_single_authorization_server_metadata(self, issuer_url: str) -> Optional[MCPOAuthMetadata]: + async def _fetch_single_authorization_server_metadata( + self, issuer_url: str + ) -> Optional[MCPOAuthMetadata]: try: parsed = urlparse(issuer_url) except Exception: @@ -1593,7 +1749,9 @@ class MCPServerManager: candidate_urls: List[str] = [] if path: - candidate_urls.append(f"{base}/.well-known/oauth-authorization-server/{path}") + candidate_urls.append( + f"{base}/.well-known/oauth-authorization-server/{path}" + ) candidate_urls.append(f"{base}/.well-known/openid-configuration/{path}") candidate_urls.append(f"{base}/.well-known/oauth-authorization-server") candidate_urls.append(f"{base}/.well-known/openid-configuration") @@ -1693,7 +1851,9 @@ class MCPServerManager: return scopes or None return None - async def _fetch_tools_with_timeout(self, client: MCPClient, server_name: str) -> List[MCPTool]: + async def _fetch_tools_with_timeout( + self, client: MCPClient, server_name: str + ) -> List[MCPTool]: """ Fetch tools from MCP client with timeout and error handling. @@ -1716,16 +1876,22 @@ class MCPServerManager: verbose_logger.warning(f"Timeout while listing tools from {server_name}") return [] except asyncio.CancelledError: - verbose_logger.warning(f"Task cancelled while listing tools from {server_name}") + verbose_logger.warning( + f"Task cancelled while listing tools from {server_name}" + ) return [] except ConnectionError as e: - verbose_logger.warning(f"Connection error while listing tools from {server_name}: {str(e)}") + verbose_logger.warning( + f"Connection error while listing tools from {server_name}: {str(e)}" + ) return [] except Exception as e: verbose_logger.warning(f"Error listing tools from {server_name}: {str(e)}") return [] - def _create_prefixed_tools(self, tools: List[MCPTool], server: MCPServer, add_prefix: bool = True) -> List[MCPTool]: + def _create_prefixed_tools( + self, tools: List[MCPTool], server: MCPServer, add_prefix: bool = True + ) -> List[MCPTool]: """ Create prefixed tools and update tool mapping. @@ -1755,7 +1921,9 @@ class MCPServerManager: self.tool_name_to_mcp_server_name_mapping[original_name] = prefix self.tool_name_to_mcp_server_name_mapping[prefixed_name] = prefix - verbose_logger.info(f"Successfully fetched {len(prefixed_tools)} tools from server {server.name}") + verbose_logger.info( + f"Successfully fetched {len(prefixed_tools)} tools from server {server.name}" + ) return prefixed_tools def _create_prefixed_prompts( @@ -1782,7 +1950,9 @@ class MCPServerManager: prompt.name = name_to_use prefixed_prompts.append(prompt) - verbose_logger.info(f"Successfully fetched {len(prefixed_prompts)} prompts from server {server.name}") + verbose_logger.info( + f"Successfully fetched {len(prefixed_prompts)} prompts from server {server.name}" + ) return prefixed_prompts def _create_prefixed_resources( @@ -1794,11 +1964,17 @@ class MCPServerManager: prefix = get_server_prefix(server) for resource in resources: - name_to_use = add_server_prefix_to_name(resource.name, prefix) if add_prefix else resource.name + name_to_use = ( + add_server_prefix_to_name(resource.name, prefix) + if add_prefix + else resource.name + ) resource.name = name_to_use prefixed_resources.append(resource) - verbose_logger.info(f"Successfully fetched {len(prefixed_resources)} resources from server {server.name}") + verbose_logger.info( + f"Successfully fetched {len(prefixed_resources)} resources from server {server.name}" + ) return prefixed_resources def _create_prefixed_resource_templates( @@ -1814,7 +1990,9 @@ class MCPServerManager: for resource_template in resource_templates: name_to_use = ( - add_server_prefix_to_name(resource_template.name, prefix) if add_prefix else resource_template.name + add_server_prefix_to_name(resource_template.name, prefix) + if add_prefix + else resource_template.name ) resource_template.name = name_to_use prefixed_templates.append(resource_template) @@ -1829,14 +2007,20 @@ class MCPServerManager: Check if the tool is allowed or banned for the given server """ if server.allowed_tools: - return tool_name in server.allowed_tools or f"{server.name}-{tool_name}" in server.allowed_tools + return ( + tool_name in server.allowed_tools + or f"{server.name}-{tool_name}" in server.allowed_tools + ) if server.disallowed_tools: return ( - tool_name not in server.disallowed_tools and f"{server.name}-{tool_name}" not in server.disallowed_tools + tool_name not in server.disallowed_tools + and f"{server.name}-{tool_name}" not in server.disallowed_tools ) return True - def validate_allowed_params(self, tool_name: str, arguments: Dict[str, Any], server: MCPServer) -> None: + def validate_allowed_params( + self, tool_name: str, arguments: Dict[str, Any], server: MCPServer + ) -> None: """ Filter arguments to only include allowed parameters for the given tool. @@ -1863,14 +2047,18 @@ class MCPServerManager: unprefixed_tool_name, _ = split_server_prefix_from_name(tool_name) # Check both prefixed and unprefixed tool names - allowed_params_list = server.allowed_params.get(tool_name) or server.allowed_params.get(unprefixed_tool_name) + allowed_params_list = server.allowed_params.get( + tool_name + ) or server.allowed_params.get(unprefixed_tool_name) # If this tool doesn't have allowed_params specified, allow all params if allowed_params_list is None: return None # Filter arguments to only include allowed parameters - disallowed_params = [param for param in arguments.keys() if param not in allowed_params_list] + disallowed_params = [ + param for param in arguments.keys() if param not in allowed_params_list + ] if disallowed_params: raise HTTPException( @@ -2034,20 +2222,38 @@ class MCPServerManager: "arguments": arguments, "server_name": server_name, "user_api_key_auth": user_api_key_auth, - "user_api_key_user_id": (getattr(user_api_key_auth, "user_id", None) if user_api_key_auth else None), - "user_api_key_team_id": (getattr(user_api_key_auth, "team_id", None) if user_api_key_auth else None), - "user_api_key_end_user_id": ( - getattr(user_api_key_auth, "end_user_id", None) if user_api_key_auth else None + "user_api_key_user_id": ( + getattr(user_api_key_auth, "user_id", None) + if user_api_key_auth + else None + ), + "user_api_key_team_id": ( + getattr(user_api_key_auth, "team_id", None) + if user_api_key_auth + else None + ), + "user_api_key_end_user_id": ( + getattr(user_api_key_auth, "end_user_id", None) + if user_api_key_auth + else None + ), + "user_api_key_hash": ( + getattr(user_api_key_auth, "api_key_hash", None) + if user_api_key_auth + else None ), - "user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None), "incoming_bearer_token": incoming_bearer_token, } # Create MCP request object for processing - mcp_request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs) + mcp_request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs( + pre_hook_kwargs + ) # Convert to LLM format for existing guardrail compatibility - synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs) + synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format( + mcp_request_obj, pre_hook_kwargs + ) hook_result: Dict[str, Any] = {} try: @@ -2059,7 +2265,11 @@ class MCPServerManager: ) if modified_data: # Convert response back to MCP format and apply modifications - modified_kwargs = proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs) + modified_kwargs = ( + proxy_logging_obj._convert_mcp_hook_response_to_kwargs( + modified_data, pre_hook_kwargs + ) + ) if modified_kwargs.get("arguments") != arguments: hook_result["arguments"] = modified_kwargs["arguments"] if modified_kwargs.get("extra_headers"): @@ -2104,7 +2314,9 @@ class MCPServerManager: "user_api_key_auth": user_api_key_auth, } - synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, during_hook_kwargs) + synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format( + request_obj, during_hook_kwargs + ) return asyncio.create_task( proxy_logging_obj.during_call_hook( @@ -2160,12 +2372,16 @@ class MCPServerManager: server_auth_header: Optional[Union[Dict[str, str], str]] = None if mcp_server_auth_headers: # Normalize keys for case-insensitive lookup - normalized_headers = {k.lower(): v for k, v in mcp_server_auth_headers.items()} + normalized_headers = { + k.lower(): v for k, v in mcp_server_auth_headers.items() + } if mcp_server.alias: server_auth_header = normalized_headers.get(mcp_server.alias.lower()) if server_auth_header is None and mcp_server.server_name: - server_auth_header = normalized_headers.get(mcp_server.server_name.lower()) + server_auth_header = normalized_headers.get( + mcp_server.server_name.lower() + ) # Fall back to deprecated mcp_auth_header if no server-specific header found if server_auth_header is None: @@ -2180,7 +2396,9 @@ class MCPServerManager: if extra_headers is None: extra_headers = {} - normalized_raw_headers = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} + normalized_raw_headers = { + str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) + } for header in mcp_server.extra_headers: if not isinstance(header, str): continue @@ -2235,9 +2453,13 @@ class MCPServerManager: ) async def _call_tool_via_client(client, params): - return await client.call_tool(params, host_progress_callback=host_progress_callback) + return await client.call_tool( + params, host_progress_callback=host_progress_callback + ) - tasks.append(asyncio.create_task(_call_tool_via_client(client, call_tool_params))) + tasks.append( + asyncio.create_task(_call_tool_via_client(client, call_tool_params)) + ) try: mcp_responses = await asyncio.gather(*tasks) @@ -2247,7 +2469,9 @@ class MCPServerManager: HTTPException, ) as e: # Re-raise guardrail exceptions to properly fail the MCP call - verbose_logger.error(f"Guardrail blocked MCP tool call during result check: {str(e)}") + verbose_logger.error( + f"Guardrail blocked MCP tool call during result check: {str(e)}" + ) raise e # If proxy_logging_obj is None, the tool call result is at index 0 @@ -2331,7 +2555,11 @@ class MCPServerManager: # oauth2_headers, look up the stored token from Redis / DB. This is the # call_tool equivalent of _get_user_oauth_extra_headers_from_db used in # list_tools. - if mcp_server.needs_user_oauth_token and not oauth2_headers and user_api_key_auth is not None: + if ( + mcp_server.needs_user_oauth_token + and not oauth2_headers + and user_api_key_auth is not None + ): user_id = getattr(user_api_key_auth, "user_id", None) if user_id: try: @@ -2355,7 +2583,9 @@ class MCPServerManager: # For OpenAPI servers, call the tool handler directly instead of via MCP client if mcp_server.spec_path: - verbose_logger.debug("Calling OpenAPI tool %s directly via HTTP handler", name) + verbose_logger.debug( + "Calling OpenAPI tool %s directly via HTTP handler", name + ) if hook_result.get("extra_headers"): verbose_logger.warning( "pre_mcp_call hook returned extra_headers for OpenAPI-backed " @@ -2364,7 +2594,11 @@ class MCPServerManager: "transport to enable hook header injection.", server_name, ) - tasks.append(asyncio.create_task(self._call_openapi_tool_handler(mcp_server, name, arguments))) + tasks.append( + asyncio.create_task( + self._call_openapi_tool_handler(mcp_server, name, arguments) + ) + ) else: return await self._call_regular_mcp_tool( mcp_server=mcp_server, @@ -2397,7 +2631,9 @@ class MCPServerManager: HTTPException, ) as e: # Re-raise guardrail exceptions to properly fail the MCP call - verbose_logger.error(f"Guardrail blocked MCP tool call during result check: {str(e)}") + verbose_logger.error( + f"Guardrail blocked MCP tool call during result check: {str(e)}" + ) raise e ######################################################### @@ -2410,7 +2646,9 @@ class MCPServerManager: """ try: if asyncio.get_running_loop(): - asyncio.create_task(self._initialize_tool_name_to_mcp_server_name_mapping()) + asyncio.create_task( + self._initialize_tool_name_to_mcp_server_name_mapping() + ) except RuntimeError as e: # no running event loop verbose_logger.exception( f"No running event loop - skipping tool name to MCP server name mapping initialization: {str(e)}" @@ -2447,12 +2685,16 @@ class MCPServerManager: if tool_name in self.tool_name_to_mcp_server_name_mapping: server_name = self.tool_name_to_mcp_server_name_mapping[tool_name] for server in self.get_registry().values(): - if normalize_server_name(server.name) == normalize_server_name(server_name): + if normalize_server_name(server.name) == normalize_server_name( + server_name + ): return server # If not found and tool name is prefixed, try extracting server name from prefix known_prefixes = { - normalize_server_name(get_server_prefix(s)) for s in self.get_registry().values() if get_server_prefix(s) + normalize_server_name(get_server_prefix(s)) + for s in self.get_registry().values() + if get_server_prefix(s) } if is_tool_name_prefixed(tool_name, known_server_prefixes=known_prefixes): ( @@ -2462,9 +2704,13 @@ class MCPServerManager: if original_tool_name in self.tool_name_to_mcp_server_name_mapping: for server in self.get_registry().values(): if server.server_name is None: - if normalize_server_name(server.name) == normalize_server_name(server_name_from_prefix): + if normalize_server_name(server.name) == normalize_server_name( + server_name_from_prefix + ): return server - elif normalize_server_name(server.server_name) == normalize_server_name(server_name_from_prefix): + elif normalize_server_name( + server.server_name + ) == normalize_server_name(server_name_from_prefix): return server return None @@ -2479,7 +2725,9 @@ class MCPServerManager: self._upstream_initialize_instructions_by_server_id.clear() # perform authz check to filter the mcp servers user has access to - prisma_client = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") + prisma_client = get_prisma_client_or_throw( + "Database not connected. Connect a database to your proxy" + ) # Load only "active", legacy "approved", and NULL (no approval workflow) rows. # Pending/rejected servers are excluded at the DB level so we never load them. from litellm.proxy._experimental.mcp_server.db import LiteLLM_MCPServerTable @@ -2517,14 +2765,18 @@ class MCPServerManager: alias=getattr(server, "alias", None), server_name=getattr(server, "server_name", None), ) - verbose_logger.debug(f"Building server from DB: {server.server_id} ({server.server_name})") + verbose_logger.debug( + f"Building server from DB: {server.server_id} ({server.server_name})" + ) new_server = await self.build_mcp_server_from_table(server) new_registry[server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.registry = new_registry - verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(new_registry)) + verbose_logger.debug( + "MCP registry refreshed (%s servers in registry)", len(new_registry) + ) def get_mcp_servers_from_ids(self, server_ids: List[str]) -> List[MCPServer]: servers = [] @@ -2546,7 +2798,9 @@ class MCPServerManager: # Fallback if proxy_server not available return {} - def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: Optional[str]) -> bool: + def _is_server_accessible_from_ip( + self, server: MCPServer, client_ip: Optional[str] + ) -> bool: """ Check if a server is accessible from the given client IP. @@ -2564,7 +2818,9 @@ class MCPServerManager: return True # Non-public server: only accessible from internal IPs general_settings = self._get_general_settings() - internal_networks = IPAddressUtils.parse_internal_networks(general_settings.get("mcp_internal_ip_ranges")) + internal_networks = IPAddressUtils.parse_internal_networks( + general_settings.get("mcp_internal_ip_ranges") + ) return IPAddressUtils.is_internal_ip(client_ip, internal_networks) def get_mcp_server_by_id(self, server_id: str) -> Optional[MCPServer]: @@ -2613,7 +2869,9 @@ class MCPServerManager: matches: List[str] = [ server_id for server_id, server in registry.items() - if server.alias == identifier or server.server_name == identifier or server.name == identifier + if server.alias == identifier + or server.server_name == identifier + or server.name == identifier ] if matches: expanded.update(matches) @@ -2653,7 +2911,9 @@ class MCPServerManager: result.setdefault(server_id, []).extend(tools or []) return result - def get_mcp_server_by_name(self, server_name: str, client_ip: Optional[str] = None) -> Optional[MCPServer]: + def get_mcp_server_by_name( + self, server_name: str, client_ip: Optional[str] = None + ) -> Optional[MCPServer]: """ Get the MCP Server from the server name. @@ -2688,7 +2948,9 @@ class MCPServerManager: return server return None - def get_filtered_registry(self, client_ip: Optional[str] = None) -> Dict[str, MCPServer]: + def get_filtered_registry( + self, client_ip: Optional[str] = None + ) -> Dict[str, MCPServer]: """ Get registry filtered by client IP access control. @@ -2699,7 +2961,11 @@ class MCPServerManager: registry = self.get_registry() if client_ip is None: return registry - return {k: v for k, v in registry.items() if self._is_server_accessible_from_ip(v, client_ip)} + return { + k: v + for k, v in registry.items() + if self._is_server_accessible_from_ip(v, client_ip) + } def _generate_stable_server_id( self, @@ -2728,7 +2994,9 @@ class MCPServerManager: A deterministic server ID string """ # Create a string from all the identifying parameters - params_string = f"{server_name}|{url}|{transport}|{auth_type or ''}|{alias or ''}" + params_string = ( + f"{server_name}|{url}|{transport}|{auth_type or ''}|{alias or ''}" + ) # Generate SHA-256 hash hash_object = hashlib.sha256(params_string.encode("utf-8")) @@ -2801,11 +3069,15 @@ class MCPServerManager: return "ok" # Add timeout wrapper to prevent hanging - await asyncio.wait_for(client.run_with_session(_noop), timeout=MCP_HEALTH_CHECK_TIMEOUT) + await asyncio.wait_for( + client.run_with_session(_noop), timeout=MCP_HEALTH_CHECK_TIMEOUT + ) self._remember_upstream_initialize_instructions(server, client) status = "healthy" except asyncio.TimeoutError: - health_check_error = f"Health check timed out after {MCP_HEALTH_CHECK_TIMEOUT} seconds" + health_check_error = ( + f"Health check timed out after {MCP_HEALTH_CHECK_TIMEOUT} seconds" + ) status = "unhealthy" except asyncio.CancelledError: health_check_error = "Health check was cancelled" @@ -2818,7 +3090,9 @@ class MCPServerManager: server_id=server.server_id, server_name=server.server_name, alias=server.alias, - description=(server.mcp_info.get("description") if server.mcp_info else None), + description=( + server.mcp_info.get("description") if server.mcp_info else None + ), url=server.url, transport=server.transport, auth_type=server.auth_type, @@ -2907,7 +3181,9 @@ class MCPServerManager: server_id=server.server_id, server_name=server.server_name, alias=server.alias, - description=(server.mcp_info.get("description") if server.mcp_info else None), + description=( + server.mcp_info.get("description") if server.mcp_info else None + ), url=server.url, spec_path=server.spec_path, transport=server.transport, @@ -2968,7 +3244,9 @@ class MCPServerManager: return await self._run_health_checks(target_server_ids) - async def _run_health_checks(self, target_server_ids: List[str]) -> List[LiteLLM_MCPServerTable]: + async def _run_health_checks( + self, target_server_ids: List[str] + ) -> List[LiteLLM_MCPServerTable]: if not target_server_ids: return [] diff --git a/litellm/proxy/_experimental/mcp_server/sampling_handler.py b/litellm/proxy/_experimental/mcp_server/sampling_handler.py index 5fe56bcc81c..8efd1d5a422 100644 --- a/litellm/proxy/_experimental/mcp_server/sampling_handler.py +++ b/litellm/proxy/_experimental/mcp_server/sampling_handler.py @@ -155,7 +155,11 @@ def _convert_single_content(content: Any) -> Dict[str, Any]: # ToolResultContent → represents tool results tool_content = getattr(content, "content", []) if isinstance(tool_content, list) and tool_content: - texts = [getattr(c, "text", str(c)) for c in tool_content if getattr(c, "type", None) == "text"] + texts = [ + getattr(c, "text", str(c)) + for c in tool_content + if getattr(c, "type", None) == "text" + ] return {"type": "text", "text": "\n".join(texts) if texts else ""} return {"type": "text", "text": str(tool_content)} # Fallback: treat as text @@ -242,7 +246,9 @@ def _extract_tool_calls(content: Any) -> List[Dict[str, Any]]: "type": "function", "function": { "name": getattr(item, "name", ""), - "arguments": json.dumps(getattr(item, "input", {}), default=str), + "arguments": json.dumps( + getattr(item, "input", {}), default=str + ), }, } ) @@ -269,7 +275,11 @@ def _extract_tool_results(content: Any) -> List[Dict[str, Any]]: # Extract text from nested content nested_content = getattr(item, "content", []) if isinstance(nested_content, list): - text_parts = [getattr(c, "text", str(c)) for c in nested_content if getattr(c, "type", None) == "text"] + text_parts = [ + getattr(c, "text", str(c)) + for c in nested_content + if getattr(c, "type", None) == "text" + ] result_text = "\n".join(text_parts) if text_parts else "" else: result_text = str(nested_content) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 7214e92481a..dd834a9194a 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -76,7 +76,9 @@ def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None: _byok_cred_cache.pop((user_id, server_id), None) -def _write_byok_cred_cache(user_id: str, server_id: str, credential: Optional[str]) -> None: +def _write_byok_cred_cache( + user_id: str, server_id: str, credential: Optional[str] +) -> None: """Write a credential value to the cache, evicting all entries if at capacity.""" if len(_byok_cred_cache) >= _BYOK_CRED_CACHE_MAX_SIZE: _byok_cred_cache.clear() @@ -101,8 +103,8 @@ try: ) from mcp.server.session import ServerSession as _McpServerSession - active_mcp_session_var: contextvars.ContextVar[Optional[_McpServerSession]] = contextvars.ContextVar( - "active_mcp_session", default=None + active_mcp_session_var: contextvars.ContextVar[Optional[_McpServerSession]] = ( + contextvars.ContextVar("active_mcp_session", default=None) ) except ImportError as e: verbose_logger.debug(f"MCP module not found: {e}") @@ -270,7 +272,9 @@ if MCP_AVAILABLE: await _session_manager_cm.__aenter__() await _sse_session_manager_cm.__aenter__() _SESSION_MANAGERS_INITIALIZED = True - verbose_logger.info("MCP Server started with StreamableHTTP session manager and SSE transport!") + verbose_logger.info( + "MCP Server started with StreamableHTTP session manager and SSE transport!" + ) async def shutdown_session_managers(): """Shutdown the session managers.""" @@ -323,8 +327,12 @@ if MCP_AVAILABLE: raw_headers, _client_ip, ) = await get_or_extract_auth_context() - verbose_logger.debug(f"MCP list_tools - User API Key Auth from context: {user_api_key_auth}") - verbose_logger.debug(f"MCP list_tools - MCP servers from context: {mcp_servers}") + verbose_logger.debug( + f"MCP list_tools - User API Key Auth from context: {user_api_key_auth}" + ) + verbose_logger.debug( + f"MCP list_tools - MCP servers from context: {mcp_servers}" + ) verbose_logger.debug( f"MCP list_tools - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}" ) @@ -340,7 +348,9 @@ if MCP_AVAILABLE: log_list_tools_to_spendlogs=True, list_tools_log_source="mcp_protocol", ) - verbose_logger.info(f"MCP list_tools - Successfully returned {len(tools)} tools") + verbose_logger.info( + f"MCP list_tools - Successfully returned {len(tools)} tools" + ) return tools except Exception as e: verbose_logger.exception(f"Error in list_tools endpoint: {str(e)}") @@ -349,7 +359,9 @@ if MCP_AVAILABLE: return [] @server.call_tool() - async def mcp_server_tool_call(name: str, arguments: Dict[str, Any] | None) -> CallToolResult: + async def mcp_server_tool_call( # noqa: PLR0915 + name: str, arguments: Dict[str, Any] | None + ) -> CallToolResult: """ Call a specific tool with the provided arguments Args: @@ -364,6 +376,14 @@ if MCP_AVAILABLE: from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request from litellm.proxy.proxy_server import proxy_config + from mcp.server.models import CallToolResult + from mcp.server.lowlevel.server import request_ctx, request_ctx_var + + req_ctx = request_ctx.get(None) + if req_ctx: + active_mcp_session_var.set(req_ctx.session) + elif request_ctx_var.get(None): + active_mcp_session_var.set(request_ctx_var.get().session) # Validate arguments ( @@ -394,12 +414,18 @@ if MCP_AVAILABLE: progress=progress, total=total, ) - verbose_logger.debug(f"Forwarded progress {progress}/{total} to Host") + verbose_logger.debug( + f"Forwarded progress {progress}/{total} to Host" + ) except Exception as e: - verbose_logger.error(f"Failed to forward progress to Host: {e}") + verbose_logger.error( + f"Failed to forward progress to Host: {e}" + ) host_progress_callback = forward_progress - verbose_logger.debug(f"Host progressToken captured: {host_token[:8]}...") + verbose_logger.debug( + f"Host progressToken captured: {host_token[:8]}..." + ) except Exception as e: verbose_logger.warning(f"Could not capture host progress context: {e}") try: @@ -451,7 +477,11 @@ if MCP_AVAILABLE: except GuardrailRaisedException as e: verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}") return CallToolResult( - content=[TextContent(text=f"Error: Guardrail violation - {str(e)}", type="text")], + content=[ + TextContent( + text=f"Error: Guardrail violation - {str(e)}", type="text" + ) + ], isError=True, ) except HTTPException as e: @@ -471,8 +501,14 @@ if MCP_AVAILABLE: @server.list_prompts() async def list_prompts() -> List[Prompt]: """ - List all available prompts + List all available prompts. + Also captures the active session for propagation to callbacks. """ + from mcp.server.lowlevel.server import request_ctx + + req_ctx = request_ctx.get(None) + if req_ctx: + active_mcp_session_var.set(req_ctx.session) try: # Get user authentication from context variable ( @@ -484,8 +520,12 @@ if MCP_AVAILABLE: raw_headers, _client_ip, ) = await get_or_extract_auth_context() - verbose_logger.debug(f"MCP list_prompts - User API Key Auth from context: {user_api_key_auth}") - verbose_logger.debug(f"MCP list_prompts - MCP servers from context: {mcp_servers}") + verbose_logger.debug( + f"MCP list_prompts - User API Key Auth from context: {user_api_key_auth}" + ) + verbose_logger.debug( + f"MCP list_prompts - MCP servers from context: {mcp_servers}" + ) verbose_logger.debug( f"MCP list_prompts - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}" ) @@ -499,7 +539,9 @@ if MCP_AVAILABLE: oauth2_headers=oauth2_headers, raw_headers=raw_headers, ) - verbose_logger.info(f"MCP list_prompts - Successfully returned {len(prompts)} prompts") + verbose_logger.info( + f"MCP list_prompts - Successfully returned {len(prompts)} prompts" + ) return prompts except Exception as e: verbose_logger.exception(f"Error in list_prompts endpoint: {str(e)}") @@ -508,15 +550,23 @@ if MCP_AVAILABLE: return [] @server.get_prompt() - async def get_prompt(name: str, arguments: dict[str, str] | None) -> GetPromptResult: + async def get_prompt( + name: str, arguments: dict[str, str] | None + ) -> GetPromptResult: """ - Get a specific prompt with the provided arguments + Get a specific prompt with the provided arguments. + Also captures the active session for propagation to callbacks. Args: name (str): Name of the prompt to get arguments (Dict[str, Any] | None): Arguments to pass to the prompt Returns: GetPromptResult: Getting prompt execution results """ + from mcp.server.lowlevel.server import request_ctx + + req_ctx = request_ctx.get(None) + if req_ctx: + active_mcp_session_var.set(req_ctx.session) # Validate arguments ( user_api_key_auth, @@ -527,7 +577,9 @@ if MCP_AVAILABLE: raw_headers, _client_ip, ) = await get_or_extract_auth_context() - verbose_logger.debug(f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}") + verbose_logger.debug( + f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}" + ) return await mcp_get_prompt( name=name, arguments=arguments, @@ -541,7 +593,15 @@ if MCP_AVAILABLE: @server.list_resources() async def list_resources() -> List[Resource]: - """List all available resources.""" + """ + List all available resources. + Also captures the active session for propagation to callbacks. + """ + from mcp.server.lowlevel.server import request_ctx + + req_ctx = request_ctx.get(None) + if req_ctx: + active_mcp_session_var.set(req_ctx.session) try: ( user_api_key_auth, @@ -552,8 +612,12 @@ if MCP_AVAILABLE: raw_headers, _client_ip, ) = await get_or_extract_auth_context() - verbose_logger.debug(f"MCP list_resources - User API Key Auth from context: {user_api_key_auth}") - verbose_logger.debug(f"MCP list_resources - MCP servers from context: {mcp_servers}") + verbose_logger.debug( + f"MCP list_resources - User API Key Auth from context: {user_api_key_auth}" + ) + verbose_logger.debug( + f"MCP list_resources - MCP servers from context: {mcp_servers}" + ) verbose_logger.debug( f"MCP list_resources - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}" ) @@ -565,7 +629,9 @@ if MCP_AVAILABLE: oauth2_headers=oauth2_headers, raw_headers=raw_headers, ) - verbose_logger.info(f"MCP list_resources - Successfully returned {len(resources)} resources") + verbose_logger.info( + f"MCP list_resources - Successfully returned {len(resources)} resources" + ) return resources except Exception as e: verbose_logger.exception(f"Error in list_resources endpoint: {str(e)}") @@ -573,7 +639,15 @@ if MCP_AVAILABLE: @server.list_resource_templates() async def list_resource_templates() -> List[ResourceTemplate]: - """List all available resource templates.""" + """ + List all available resource templates. + Also captures the active session for propagation to callbacks. + """ + from mcp.server.lowlevel.server import request_ctx + + req_ctx = request_ctx.get(None) + if req_ctx: + active_mcp_session_var.set(req_ctx.session) try: ( user_api_key_auth, @@ -584,8 +658,12 @@ if MCP_AVAILABLE: raw_headers, _client_ip, ) = await get_or_extract_auth_context() - verbose_logger.debug(f"MCP list_resource_templates - User API Key Auth from context: {user_api_key_auth}") - verbose_logger.debug(f"MCP list_resource_templates - MCP servers from context: {mcp_servers}") + verbose_logger.debug( + f"MCP list_resource_templates - User API Key Auth from context: {user_api_key_auth}" + ) + verbose_logger.debug( + f"MCP list_resource_templates - MCP servers from context: {mcp_servers}" + ) verbose_logger.debug( f"MCP list_resource_templates - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}" ) @@ -602,11 +680,22 @@ if MCP_AVAILABLE: ) return resource_templates except Exception as e: - verbose_logger.exception(f"Error in list_resource_templates endpoint: {str(e)}") + verbose_logger.exception( + f"Error in list_resource_templates endpoint: {str(e)}" + ) return [] @server.read_resource() async def read_resource(url: AnyUrl) -> list[ReadResourceContents]: + """ + Read resource contents from upstream MCP servers. + Also captures the active session for propagation to callbacks. + """ + from mcp.server.lowlevel.server import request_ctx + + req_ctx = request_ctx.get(None) + if req_ctx: + active_mcp_session_var.set(req_ctx.session) ( user_api_key_auth, mcp_auth_header, @@ -662,8 +751,10 @@ if MCP_AVAILABLE: break if not server_name_matched: try: - access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups( - [server_or_group] + access_group_server_ids = ( + await MCPRequestHandler._get_mcp_servers_from_access_groups( + [server_or_group] + ) ) # Only include servers that the user has access to for server_id in access_group_server_ids: @@ -671,7 +762,9 @@ if MCP_AVAILABLE: if server_id == server.server_id: filtered_server[server.server_id] = server except Exception as e: - verbose_logger.debug(f"Could not resolve '{server_or_group}' as access group: {e}") + verbose_logger.debug( + f"Could not resolve '{server_or_group}' as access group: {e}" + ) if filtered_server: return list(filtered_server.values()) return allowed_mcp_servers @@ -718,11 +811,17 @@ if MCP_AVAILABLE: tools_to_return = tools # Filter by allowed_tools (whitelist) if mcp_server.allowed_tools: - tools_to_return = [tool for tool in tools if _tool_name_matches(tool.name, mcp_server.allowed_tools)] + tools_to_return = [ + tool + for tool in tools + if _tool_name_matches(tool.name, mcp_server.allowed_tools) + ] # Filter by disallowed_tools (blacklist) if mcp_server.disallowed_tools: tools_to_return = [ - tool for tool in tools_to_return if not _tool_name_matches(tool.name, mcp_server.disallowed_tools) + tool + for tool in tools_to_return + if not _tool_name_matches(tool.name, mcp_server.disallowed_tools) ] return tools_to_return @@ -783,11 +882,15 @@ if MCP_AVAILABLE: "MCP _get_allowed_mcp_servers called without client_ip and no auth context. " "IP filtering will be skipped. This is expected for internal calls." ) - allowed_mcp_server_ids = await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth) + allowed_mcp_server_ids = ( + await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth) + ) ( allowed_mcp_server_ids, _ip_blocked, - ) = global_mcp_server_manager.filter_server_ids_by_ip_with_info(allowed_mcp_server_ids, client_ip) + ) = global_mcp_server_manager.filter_server_ids_by_ip_with_info( + allowed_mcp_server_ids, client_ip + ) verbose_logger.debug( "MCP IP filter: client_ip=%s, allowed_server_ids=%s", client_ip, @@ -805,7 +908,9 @@ if MCP_AVAILABLE: ) allowed_mcp_servers: List[MCPServer] = [] for allowed_mcp_server_id in allowed_mcp_server_ids: - mcp_server = global_mcp_server_manager.get_mcp_server_by_id(allowed_mcp_server_id) + mcp_server = global_mcp_server_manager.get_mcp_server_by_id( + allowed_mcp_server_id + ) if mcp_server is not None: allowed_mcp_servers.append(mcp_server) if mcp_servers is not None: @@ -871,7 +976,9 @@ if MCP_AVAILABLE: prisma_client = get_prisma_client_or_throw( "Database not connected. Connect a database to use OAuth2 MCP tools." ) - cred = await get_user_oauth_credential(prisma_client, user_id, server_id) + cred = await get_user_oauth_credential( + prisma_client, user_id, server_id + ) if not cred or not cred.get("access_token"): return None if is_oauth_credential_expired(cred): @@ -925,12 +1032,16 @@ if MCP_AVAILABLE: exp_dt = datetime.fromisoformat(expires_at) if exp_dt.tzinfo is None: exp_dt = exp_dt.replace(tzinfo=timezone.utc) - remaining = int((exp_dt - datetime.now(timezone.utc)).total_seconds()) + remaining = int( + (exp_dt - datetime.now(timezone.utc)).total_seconds() + ) raw_expires = max(remaining, 0) if remaining > 0 else None except (ValueError, TypeError): pass ttl = _compute_per_user_token_ttl(server, raw_expires) - await mcp_per_user_token_cache.set(user_id, server_id, access_token, ttl) + await mcp_per_user_token_cache.set( + user_id, server_id, access_token, ttl + ) return {"Authorization": f"Bearer {access_token}"} except Exception as e: verbose_logger.warning( @@ -947,7 +1058,9 @@ if MCP_AVAILABLE: """Fetch all OAuth2 credentials for the user in one DB query. Returns a dict keyed by server_id to avoid N+1 queries in asyncio.gather loops. """ - user_id = getattr(user_api_key_auth, "user_id", None) if user_api_key_auth else None + user_id = ( + getattr(user_api_key_auth, "user_id", None) if user_api_key_auth else None + ) if not user_id: return {} try: @@ -962,7 +1075,9 @@ if MCP_AVAILABLE: creds = await list_user_oauth_credentials(prisma_client, user_id) return {c["server_id"]: c for c in creds if "server_id" in c} except Exception as e: - verbose_logger.warning(f"_prefetch_oauth_creds_for_user: failed to prefetch for user={user_id}: {e}") + verbose_logger.warning( + f"_prefetch_oauth_creds_for_user: failed to prefetch for user={user_id}: {e}" + ) return {} def _prepare_mcp_server_headers( @@ -985,7 +1100,9 @@ if MCP_AVAILABLE: if server.extra_headers and raw_headers: if extra_headers is None: extra_headers = {} - normalized_raw_headers = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} + normalized_raw_headers = { + str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) + } for header in server.extra_headers: if not isinstance(header, str): continue @@ -1005,13 +1122,21 @@ if MCP_AVAILABLE: return None texts: List[Tuple[str, str]] = [] for server in allowed_mcp_servers: - label = server.alias or server.server_name or server.name or server.server_id or "mcp" + label = ( + server.alias + or server.server_name + or server.name + or server.server_id + or "mcp" + ) if server.instructions and server.instructions.strip(): texts.append((label, server.instructions.strip())) continue if server.spec_path: continue - cached = global_mcp_server_manager._upstream_initialize_instructions_by_server_id.get(server.server_id) + cached = global_mcp_server_manager._upstream_initialize_instructions_by_server_id.get( + server.server_id + ) if cached and cached.strip(): texts.append((label, cached.strip())) if not texts: @@ -1070,7 +1195,9 @@ if MCP_AVAILABLE: rules_obj = Rules() list_tools_call_id = str(uuid.uuid4()) # Derive trace_id from raw_headers when not explicitly passed (same as A2A / MCP call_tool) - effective_litellm_trace_id = litellm_trace_id or get_chain_id_from_headers(raw_headers) + effective_litellm_trace_id = litellm_trace_id or get_chain_id_from_headers( + raw_headers + ) spend_logs_metadata: Dict[str, Any] = { "mcp_operation": "list_tools", } @@ -1104,9 +1231,9 @@ if MCP_AVAILABLE: user_api_key_dict=user_api_key_auth, _metadata_variable_name="metadata", ) - 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 try: @@ -1120,7 +1247,9 @@ if MCP_AVAILABLE: litellm_logging_obj.call_type = CallTypes.list_mcp_tools.value litellm_logging_obj.model = "MCP: list_tools" except Exception as logging_error: - verbose_logger.debug("Failed to initialize logging for MCP list_tools: %s", logging_error) + verbose_logger.debug( + "Failed to initialize logging for MCP list_tools: %s", logging_error + ) litellm_logging_obj = None try: allowed_mcp_servers = await _get_allowed_mcp_servers( @@ -1129,9 +1258,14 @@ if MCP_AVAILABLE: ) # Pre-fetch OAuth credentials only when at least one server uses OAuth2, # to avoid an unnecessary DB round-trip on requests with no OAuth2 MCP servers. - _has_oauth2_server = any(getattr(s, "auth_type", None) == MCPAuth.oauth2 for s in allowed_mcp_servers) + _has_oauth2_server = any( + getattr(s, "auth_type", None) == MCPAuth.oauth2 + for s in allowed_mcp_servers + ) _prefetched_oauth_creds = ( - await _prefetch_oauth_creds_for_user(user_api_key_auth) if _has_oauth2_server else {} + await _prefetch_oauth_creds_for_user(user_api_key_auth) + if _has_oauth2_server + else {} ) async def _fetch_and_filter_server_tools( @@ -1176,11 +1310,15 @@ if MCP_AVAILABLE: ) return filtered_tools except Exception as e: - verbose_logger.exception(f"Error getting tools from server {server.name}: {str(e)}") + verbose_logger.exception( + f"Error getting tools from server {server.name}: {str(e)}" + ) return [] # Fetch tools from all servers in parallel - tasks = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers] + tasks = [ + _fetch_and_filter_server_tools(server) for server in allowed_mcp_servers + ] results = await asyncio.gather(*tasks) # Flatten results into single list all_tools: List[MCPTool] = [tool for tools in results for tool in tools] @@ -1212,7 +1350,9 @@ if MCP_AVAILABLE: start_time=list_tools_start_time, end_time=end_time, ) - verbose_logger.info(f"Successfully fetched {len(all_tools)} tools total from all MCP servers") + verbose_logger.info( + f"Successfully fetched {len(all_tools)} tools total from all MCP servers" + ) return all_tools except Exception as e: # Only fire failure hook if logging was requested for this list-tools execution @@ -1221,7 +1361,9 @@ if MCP_AVAILABLE: from litellm.proxy.proxy_server import proxy_logging_obj if proxy_logging_obj: - 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=list_tools_request_data or {}, original_exception=e, @@ -1230,7 +1372,9 @@ if MCP_AVAILABLE: traceback_str=traceback_str, ) except Exception: - verbose_logger.debug("Failed to log MCP list_tools failure via post_call_failure_hook") + verbose_logger.debug( + "Failed to log MCP list_tools failure via post_call_failure_hook" + ) raise async def _get_prompts_from_mcp_servers( @@ -1279,11 +1423,17 @@ if MCP_AVAILABLE: raw_headers=raw_headers, ) all_prompts.extend(prompts) - verbose_logger.debug(f"Successfully fetched {len(prompts)} prompts from server {server.name}") + verbose_logger.debug( + f"Successfully fetched {len(prompts)} prompts from server {server.name}" + ) except Exception as e: - verbose_logger.exception(f"Error getting prompts from server {server.name}: {str(e)}") + verbose_logger.exception( + f"Error getting prompts from server {server.name}: {str(e)}" + ) # Continue with other servers instead of failing completely - verbose_logger.info(f"Successfully fetched {len(all_prompts)} prompts total from all MCP servers") + verbose_logger.info( + f"Successfully fetched {len(all_prompts)} prompts total from all MCP servers" + ) return all_prompts async def _get_resources_from_mcp_servers( @@ -1321,10 +1471,16 @@ if MCP_AVAILABLE: raw_headers=raw_headers, ) all_resources.extend(resources) - verbose_logger.debug(f"Successfully fetched {len(resources)} resources from server {server.name}") + verbose_logger.debug( + f"Successfully fetched {len(resources)} resources from server {server.name}" + ) except Exception as e: - verbose_logger.exception(f"Error getting resources from server {server.name}: {str(e)}") - verbose_logger.info(f"Successfully fetched {len(all_resources)} resources total from all MCP servers") + verbose_logger.exception( + f"Error getting resources from server {server.name}: {str(e)}" + ) + verbose_logger.info( + f"Successfully fetched {len(all_resources)} resources total from all MCP servers" + ) return all_resources async def _get_resource_templates_from_mcp_servers( @@ -1354,12 +1510,14 @@ if MCP_AVAILABLE: raw_headers=raw_headers, ) try: - resource_templates = await global_mcp_server_manager.get_resource_templates_from_server( - server=server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, + resource_templates = ( + await global_mcp_server_manager.get_resource_templates_from_server( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, # Always add server prefix + raw_headers=raw_headers, + ) ) all_resource_templates.extend(resource_templates) verbose_logger.debug( @@ -1425,7 +1583,11 @@ if MCP_AVAILABLE: toolset_ids = getattr(op, "mcp_toolsets", None) or [] if not toolset_ids: return user_api_key_auth - toolset_perms = await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=toolset_ids) + toolset_perms = ( + await global_mcp_server_manager.resolve_toolset_tool_permissions( + toolset_ids=toolset_ids + ) + ) if not toolset_perms: return user_api_key_auth # Merge toolset_perms into existing mcp_tool_permissions (union) @@ -1439,7 +1601,9 @@ if MCP_AVAILABLE: # filtering doesn't silently drop servers that the toolset references but that # aren't already in the key's explicit mcp_servers list. merged_servers = list(set(op.mcp_servers or []) | set(existing.keys())) - updated_op = op.model_copy(update={"mcp_servers": merged_servers, "mcp_tool_permissions": existing}) + updated_op = op.model_copy( + update={"mcp_servers": merged_servers, "mcp_tool_permissions": existing} + ) return user_api_key_auth.model_copy(update={"object_permission": updated_op}) async def _list_mcp_tools( @@ -1480,9 +1644,13 @@ if MCP_AVAILABLE: log_list_tools_to_spendlogs=log_list_tools_to_spendlogs, list_tools_log_source=list_tools_log_source, ) - verbose_logger.debug(f"Successfully fetched {len(managed_tools)} tools from managed MCP servers") + verbose_logger.debug( + f"Successfully fetched {len(managed_tools)} tools from managed MCP servers" + ) except Exception as e: - verbose_logger.exception(f"Error getting tools from managed MCP servers: {str(e)}") + verbose_logger.exception( + f"Error getting tools from managed MCP servers: {str(e)}" + ) # Continue with empty managed tools list instead of failing completely return managed_tools @@ -1517,9 +1685,13 @@ if MCP_AVAILABLE: oauth2_headers=oauth2_headers, raw_headers=raw_headers, ) - verbose_logger.debug(f"Successfully fetched {len(managed_prompts)} prompts from managed MCP servers") + verbose_logger.debug( + f"Successfully fetched {len(managed_prompts)} prompts from managed MCP servers" + ) except Exception as e: - verbose_logger.exception(f"Error getting tools from managed MCP servers: {str(e)}") + verbose_logger.exception( + f"Error getting tools from managed MCP servers: {str(e)}" + ) # Continue with empty managed tools list instead of failing completely return managed_prompts @@ -1544,9 +1716,13 @@ if MCP_AVAILABLE: oauth2_headers=oauth2_headers, raw_headers=raw_headers, ) - verbose_logger.debug(f"Successfully fetched {len(managed_resources)} resources from managed MCP servers") + verbose_logger.debug( + f"Successfully fetched {len(managed_resources)} resources from managed MCP servers" + ) except Exception as e: - verbose_logger.exception(f"Error getting resources from managed MCP servers: {str(e)}") + verbose_logger.exception( + f"Error getting resources from managed MCP servers: {str(e)}" + ) return managed_resources async def _list_mcp_resource_templates( @@ -1595,7 +1771,9 @@ if MCP_AVAILABLE: display_map = server.tool_name_to_display_name or {} for unprefixed_name, display_name in display_map.items(): if display_name == name: - return add_server_prefix_to_name(unprefixed_name, get_server_prefix(server)) + return add_server_prefix_to_name( + unprefixed_name, get_server_prefix(server) + ) return name async def _get_byok_credential( @@ -1652,7 +1830,9 @@ if MCP_AVAILABLE: "server_name": mcp_server.server_name or mcp_server.name, "message": "User identity is required for BYOK servers", }, - headers={"WWW-Authenticate": 'Bearer resource_metadata="/.well-known/oauth-protected-resource"'}, + headers={ + "WWW-Authenticate": 'Bearer resource_metadata="/.well-known/oauth-protected-resource"' + }, ) # Check shared credential cache before hitting the DB. cache_key = (user_id, mcp_server.server_id) @@ -1711,7 +1891,9 @@ if MCP_AVAILABLE: "Complete the OAuth authorization flow to provide your API key." ), }, - headers={"WWW-Authenticate": 'Bearer resource_metadata="/.well-known/oauth-protected-resource"'}, + headers={ + "WWW-Authenticate": 'Bearer resource_metadata="/.well-known/oauth-protected-resource"' + }, ) async def execute_mcp_tool( # noqa: PLR0915 @@ -1766,25 +1948,33 @@ if MCP_AVAILABLE: status_code=403, detail=f"User not allowed to call this tool. Allowed MCP servers: {allowed_mcp_servers}", ) - standard_logging_mcp_tool_call: StandardLoggingMCPToolCall = _get_standard_logging_mcp_tool_call( - name=original_tool_name, # Use original name for logging - arguments=arguments, - server_name=server_name, + standard_logging_mcp_tool_call: StandardLoggingMCPToolCall = ( + _get_standard_logging_mcp_tool_call( + name=original_tool_name, # Use original name for logging + arguments=arguments, + server_name=server_name, + ) + ) + litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get( + "litellm_logging_obj", None ) - litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj", None) if litellm_logging_obj: - litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = standard_logging_mcp_tool_call + litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = ( + standard_logging_mcp_tool_call + ) litellm_logging_obj.model = f"MCP: {name}" # Resolve the MCP server early so BYOK checks and credential injection # apply to ALL dispatch paths (local tool registry AND managed MCP server). if mcp_server is None: mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) if mcp_server: - standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get( - "mcp_server_cost_info" - ) + standard_logging_mcp_tool_call["mcp_server_cost_info"] = ( + mcp_server.mcp_info or {} + ).get("mcp_server_cost_info") if litellm_logging_obj: - litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = standard_logging_mcp_tool_call + litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = ( + standard_logging_mcp_tool_call + ) # BYOK: retrieve the stored per-user credential. A single DB call # both checks existence and fetches the value, avoiding a double query. if mcp_server.is_byok and not mcp_auth_header: @@ -1821,7 +2011,9 @@ if MCP_AVAILABLE: # configured auth_type so the generator doesn't need to know the prefix. auth_header_value: Optional[str] = None if mcp_auth_header: - server_auth_type = getattr(mcp_server, "auth_type", None) if mcp_server else None + server_auth_type = ( + getattr(mcp_server, "auth_type", None) if mcp_server else None + ) if server_auth_type == MCPAuth.api_key: auth_header_value = f"ApiKey {mcp_auth_header}" elif server_auth_type == MCPAuth.basic: @@ -1875,17 +2067,25 @@ if MCP_AVAILABLE: Call a specific tool with the provided arguments (handles prefixed tool names). """ start_time = datetime.now() - litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj", None) + litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get( + "litellm_logging_obj", None + ) try: if arguments is None: - raise HTTPException(status_code=400, detail="Request arguments are required") + raise HTTPException( + status_code=400, detail="Request arguments are required" + ) ## CHECK IF USER IS ALLOWED TO CALL THIS TOOL - allowed_mcp_server_ids = await global_mcp_server_manager.get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, + allowed_mcp_server_ids = ( + await global_mcp_server_manager.get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + ) ) allowed_mcp_servers: List[MCPServer] = [] for allowed_mcp_server_id in allowed_mcp_server_ids: - allowed_server = global_mcp_server_manager.get_mcp_server_by_id(allowed_mcp_server_id) + allowed_server = global_mcp_server_manager.get_mcp_server_by_id( + allowed_mcp_server_id + ) if allowed_server is not None: allowed_mcp_servers.append(allowed_server) allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( @@ -1933,7 +2133,9 @@ if MCP_AVAILABLE: end_time=end_time, ) litellm_logging_obj.call_type = CallTypes.call_mcp_tool.value - await litellm_logging_obj.async_success_handler(result=response, start_time=start_time, end_time=end_time) + await litellm_logging_obj.async_success_handler( + result=response, start_time=start_time, end_time=end_time + ) return response async def mcp_get_prompt( @@ -2126,15 +2328,21 @@ if MCP_AVAILABLE: # Path found at the end, remove it from servers path_part = "/" + path_match.group(1) servers_part = servers_and_path[: -len(path_part)] - mcp_servers_from_path = [s.strip() for s in servers_part.split(",") if s.strip()] + mcp_servers_from_path = [ + s.strip() for s in servers_part.split(",") if s.strip() + ] else: # No path, just comma-separated servers - mcp_servers_from_path = [s.strip() for s in servers_and_path.split(",") if s.strip()] + mcp_servers_from_path = [ + s.strip() for s in servers_and_path.split(",") if s.strip() + ] else: # Single server case - use regex approach for server/path separation # This handles cases like "custom_solutions/user_123/chat/completions" # where we want to extract "custom_solutions/user_123" as the server name - single_server_match = re.match(r"^([^/]+(?:/[^/]+)?)(?:/.*)?$", servers_and_path) + single_server_match = re.match( + r"^([^/]+(?:/[^/]+)?)(?:/.*)?$", servers_and_path + ) if single_server_match: server_name = single_server_match.group(1) mcp_servers_from_path = [server_name] @@ -2247,7 +2455,11 @@ if MCP_AVAILABLE: "Stripping stale header to force new session creation.", _session_id, ) - scope["headers"] = [(k, v) for k, v in _headers if _normalize_header_name(k) != _mcp_session_header] + scope["headers"] = [ + (k, v) + for k, v in _headers + if _normalize_header_name(k) != _mcp_session_header + ] return False async def _apply_toolset_scope( @@ -2279,7 +2491,11 @@ if MCP_AVAILABLE: status_code=403, detail=f"API key does not have access to toolset '{toolset_id}'.", ) - tool_permissions = await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=[toolset_id]) + tool_permissions = ( + await global_mcp_server_manager.resolve_toolset_tool_permissions( + toolset_ids=[toolset_id] + ) + ) server_ids = list(tool_permissions.keys()) existing_op = user_api_key_auth.object_permission if existing_op is not None: @@ -2300,7 +2516,9 @@ if MCP_AVAILABLE: ) return user_api_key_auth.model_copy(update={"object_permission": updated_op}) - async def handle_streamable_http_mcp(scope: Scope, receive: Receive, send: Send) -> None: + async def handle_streamable_http_mcp( + scope: Scope, receive: Receive, send: Send + ) -> None: """Handle MCP requests through StreamableHTTP.""" try: path = scope.get("path", "") @@ -2314,42 +2532,52 @@ if MCP_AVAILABLE: ) = await extract_mcp_auth_context(scope, path) # Extract client IP for MCP access control _client_ip = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope)) - verbose_logger.debug(f"MCP request mcp_servers (header/path): {mcp_servers}") + verbose_logger.debug( + f"MCP request mcp_servers (header/path): {mcp_servers}" + ) verbose_logger.debug( f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}" ) # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response for server_name in mcp_servers or []: - server = global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=_client_ip) + server = global_mcp_server_manager.get_mcp_server_by_name( + server_name, client_ip=_client_ip + ) if server and server.auth_type == MCPAuth.oauth2 and not oauth2_headers: # For per-user OAuth servers, only skip the pre-emptive 401 when # a stored token actually exists for this user+server pair. # If no stored token exists, fail fast with 401 so clients can # kick off PKCE/interactive OAuth flow immediately. if server.needs_user_oauth_token: - stored_oauth_headers = await _get_user_oauth_extra_headers_from_db( - server=server, - user_api_key_auth=user_api_key_auth, + stored_oauth_headers = ( + await _get_user_oauth_extra_headers_from_db( + server=server, + user_api_key_auth=user_api_key_auth, + ) ) if stored_oauth_headers: continue request = StarletteRequest(scope) base_url = get_request_base_url(request) - authorization_uri = ( - f"Bearer authorization_uri={base_url}/.well-known/oauth-authorization-server/{server_name}" - ) + authorization_uri = f"Bearer authorization_uri={base_url}/.well-known/oauth-authorization-server/{server_name}" raise HTTPException( status_code=401, detail="Unauthorized", headers={"www-authenticate": authorization_uri}, ) # Strip any client-supplied x-mcp-toolset-id to prevent forgery. - scope["headers"] = [(k, v) for k, v in scope.get("headers", []) if k.lower() != b"x-mcp-toolset-id"] + scope["headers"] = [ + (k, v) + for k, v in scope.get("headers", []) + if k.lower() != b"x-mcp-toolset-id" + ] # Apply toolset scope if set server-side via ContextVar (set by # /toolset/{name}/mcp and /{name}/mcp route handlers in proxy_server.py). active_toolset_id = _mcp_active_toolset_id.get() if active_toolset_id and user_api_key_auth is not None: - user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id) + user_api_key_auth = await _apply_toolset_scope( + user_api_key_auth, active_toolset_id + ) # Inject masked debug headers when client sends x-litellm-mcp-debug: true _debug_headers = MCPDebug.maybe_build_debug_headers( raw_headers=raw_headers, @@ -2379,7 +2607,9 @@ if MCP_AVAILABLE: await asyncio.sleep(0.1) # Handle stale session IDs - either strip them for reconnection # or return success for idempotent DELETE operations - handled = await _handle_stale_mcp_session(scope, receive, send, session_manager) + handled = await _handle_stale_mcp_session( + scope, receive, send, session_manager + ) if handled: # Request was fully handled (e.g., DELETE on non-existent session) return @@ -2405,21 +2635,22 @@ if MCP_AVAILABLE: ) await error_response(scope, receive, send) except Exception as response_error: - verbose_logger.exception(f"Failed to send error response: {response_error}") + verbose_logger.exception( + f"Failed to send error response: {response_error}" + ) # If we can't send a proper response, re-raise the original error raise e - async def handle_sse_mcp_endpoint(request: StarletteRequest): + async def handle_sse_mcp_endpoint( + scope: Scope, receive: Receive, send: Send + ) -> None: """ - Handle MCP SSE GET requests. - This is a Starlette Route endpoint handler (takes Request, returns Response). + Handle MCP SSE GET requests as a raw ASGI app. Follows the pattern documented in the official MCP SDK source at mcp/server/sse.py lines 6-31. - CRITICAL: Must return Response() after the SSE connection ends to prevent - "TypeError: 'NoneType' object is not callable" when the client disconnects. """ try: - scope = request.scope + request = StarletteRequest(scope, receive) path = scope.get("path", "") ( user_api_key_auth, @@ -2431,7 +2662,9 @@ if MCP_AVAILABLE: ) = await extract_mcp_auth_context(scope, path) # Extract client IP for MCP access control _sse_client_ip = IPAddressUtils.get_mcp_client_ip(request) - verbose_logger.debug(f"MCP SSE request mcp_servers (header/path): {mcp_servers}") + verbose_logger.debug( + f"MCP SSE request mcp_servers (header/path): {mcp_servers}" + ) verbose_logger.debug( f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}" ) @@ -2448,7 +2681,7 @@ if MCP_AVAILABLE: # ContextVars are lost when the MCP SDK spawns internal tasks # (e.g. _receive_loop), so tool handlers can't read auth_context_var. # Storing it on the server object makes it available everywhere. - server._litellm_auth_context = MCPAuthenticatedUser( + server._litellm_auth_context = MCPAuthenticatedUser( # type: ignore user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -2467,8 +2700,10 @@ if MCP_AVAILABLE: ): verbose_logger.info("Initializing SSE session...") options = server.create_initialization_options() - async with sse.connect_sse(request.scope, request.receive, request._send) as streams: - verbose_logger.info("SSE connection established, running server loop...") + async with sse.connect_sse(scope, receive, send) as streams: + verbose_logger.info( + "SSE connection established, running server loop..." + ) try: # Capture the session for propagation to sampling/elicitation callbacks # Since server.run doesn't return the session, we use a middleware-like @@ -2489,16 +2724,62 @@ if MCP_AVAILABLE: status_code=HTTP_500_INTERNAL_SERVER_ERROR, content={"error": "MCP request failed", "details": str(e)}, ) - await error_response(request.scope, request.receive, request._send) + await error_response(scope, receive, send) except Exception as response_error: - verbose_logger.exception(f"Failed to send error response: {response_error}") + verbose_logger.exception( + f"Failed to send error response: {response_error}" + ) # If we can't send a proper response, re-raise the original error raise e - # CRITICAL: Return empty Response to prevent NoneType crash. - # See MCP SDK docstring at mcp/server/sse.py lines 25-26. - from starlette.responses import Response as StarletteResponse + # No need to return Response for raw ASGI app. - return StarletteResponse() + async def handle_sse_post_messages( + scope: Scope, receive: Receive, send: Send + ) -> None: + """Handle SSE POST messages by delegating to the SDK's handle_post_message.""" + try: + request = StarletteRequest(scope, receive) + path = scope.get("path", "") + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + ) = await extract_mcp_auth_context(scope, path) + _sse_client_ip = IPAddressUtils.get_mcp_client_ip(request) + set_auth_context( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=_sse_client_ip, + ) + except Exception as e: + verbose_logger.warning( + f"Failed to extract auth context in POST /messages: {e}" + ) + # The SDK's handler calls `send` directly. + await sse.handle_post_message(scope, receive, send) + # No need to return NoOpResponse for raw ASGI app. + + def get_active_mcp_session() -> Optional[_McpServerSession]: + """Get the active downstream MCP session from the current context.""" + return active_mcp_session_var.get() + + def get_active_auth_context() -> Optional[MCPAuthenticatedUser]: + """Get the active auth context from the server object or context var.""" + # Check context var first + auth = auth_context_var.get() + if auth and isinstance(auth, MCPAuthenticatedUser): + return auth + elif auth: + return cast(MCPAuthenticatedUser, auth) + # Fallback to server object + return getattr(server, "_litellm_auth_context", None) app = FastAPI( title=LITELLM_MCP_SERVER_NAME, @@ -2521,62 +2802,19 @@ if MCP_AVAILABLE: # Include the MCP router app.include_router(router) # Mount SSE handlers using the SDK's documented pattern. - # We use Starlette Route for the SSE GET endpoint (must return Response), - # and a FastAPI POST route for the POST messages endpoint. + # We use raw ASGI apps for both the SSE GET endpoint and the POST messages endpoint + # to avoid accessing private Starlette request attributes. from starlette.routing import Route as StarletteRoute - app.routes.insert(0, StarletteRoute("/sse", endpoint=handle_sse_mcp_endpoint, methods=["GET"])) - from starlette.responses import Response as StarletteResponse - - class NoOpResponse(StarletteResponse): - """A response that does nothing. Used when the underlying ASGI app already sent the response.""" - - async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: - pass - - @app.post("/messages", include_in_schema=False) - async def handle_sse_post_messages(request: StarletteRequest): - """Handle SSE POST messages by delegating to the SDK's handle_post_message.""" - try: - scope = request.scope - path = scope.get("path", "") - ( - user_api_key_auth, - mcp_auth_header, - mcp_servers, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - ) = await extract_mcp_auth_context(scope, path) - _sse_client_ip = IPAddressUtils.get_mcp_client_ip(request) - set_auth_context( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - client_ip=_sse_client_ip, - ) - except Exception as e: - verbose_logger.warning(f"Failed to extract auth context in POST /messages: {e}") - # The SDK's handler calls `send` directly. - await sse.handle_post_message(request.scope, request.receive, request._send) - # Return NoOpResponse to prevent Starlette from sending a second response. - return NoOpResponse() - - def get_active_mcp_session() -> Optional[_McpServerSession]: - """Get the active downstream MCP session from the current context.""" - return active_mcp_session_var.get() - - def get_active_auth_context() -> Optional[MCPAuthenticatedUser]: - """Get the active auth context from the server object or context var.""" - # Check context var first - auth = auth_context_var.get() - if auth: - return auth - # Fallback to server object - return getattr(server, "_litellm_auth_context", None) + app.routes.insert( + 0, StarletteRoute("/sse", endpoint=handle_sse_mcp_endpoint, methods=["GET"]) + ) + app.routes.insert( + 0, + StarletteRoute( + "/messages", endpoint=handle_sse_post_messages, methods=["POST"] + ), + ) # StreamableHTTP catch-all mounts (must come after specific routes) app.mount("/mcp", handle_streamable_http_mcp) @@ -2670,9 +2908,13 @@ if MCP_AVAILABLE: # Fallback: read from server object if ContextVar was lost if user_api_key_auth is None: stored = getattr(server, "_litellm_auth_context", None) - verbose_logger.debug(f"get_or_extract_auth_context FALLBACK: stored={stored}, type={type(stored)}") + verbose_logger.debug( + f"get_or_extract_auth_context FALLBACK: stored={stored}, type={type(stored)}" + ) if stored and isinstance(stored, MCPAuthenticatedUser): - verbose_logger.debug("get_or_extract_auth_context: Recovered auth from server object") + verbose_logger.debug( + "get_or_extract_auth_context: Recovered auth from server object" + ) user_api_key_auth = stored.user_api_key_auth mcp_auth_header = stored.mcp_auth_header mcp_servers = stored.mcp_servers diff --git a/tests/mcp_sampling_elicitation/test_live_mcp.py b/tests/mcp_sampling_elicitation/test_live_mcp.py index 4b15cac7342..71e8e358972 100644 --- a/tests/mcp_sampling_elicitation/test_live_mcp.py +++ b/tests/mcp_sampling_elicitation/test_live_mcp.py @@ -28,19 +28,25 @@ async def main(): return ElicitResult(action="accept", content=user_response) - async with sse_client("http://localhost:4000/mcp/sse", headers={"Authorization": "Bearer sk-1234"}) as ( + async with sse_client( + "http://localhost:4000/mcp/sse", headers={"Authorization": "Bearer sk-1234"} + ) as ( read_stream, write_stream, ): logger.info("SSE connection established.") - async with ClientSession(read_stream, write_stream, elicitation_callback=my_elicitation_callback) as session: + async with ClientSession( + read_stream, write_stream, elicitation_callback=my_elicitation_callback + ) as session: await session.initialize() logger.info("Initialized!") logger.info("\n--- Testing Complex Pipeline (Elicitation + Sampling) ---") logger.info("Calling 'test_server-test_complex_pipeline'...") try: - result = await session.call_tool("test_server-test_complex_pipeline", arguments={}) + result = await session.call_tool( + "test_server-test_complex_pipeline", arguments={} + ) logger.info("\nFINAL TOOL RESULT:") logger.info("==================") logger.info(result.content[0].text) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_auth.py index d1884e89d8c..6d6f550b4a8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_auth.py @@ -31,7 +31,11 @@ async def test_auth_context_persistence(): @pytest.mark.asyncio async def test_get_or_extract_auth_context_fallback(): """Test get_or_extract_auth_context fallback to server object.""" - from litellm.proxy._experimental.mcp_server.server import server, MCPAuthenticatedUser, auth_context_var + from litellm.proxy._experimental.mcp_server.server import ( + server, + MCPAuthenticatedUser, + auth_context_var, + ) auth_data = UserAPIKeyAuth(api_key="fallback-key") auth_user = MCPAuthenticatedUser(user_api_key_auth=auth_data) @@ -61,7 +65,9 @@ async def test_extract_mcp_auth_context_with_key(): mock_user_auth = UserAPIKeyAuth(api_key="sk-123") - with patch("litellm.proxy.auth.auth_checks.common_checks", new_callable=AsyncMock) as mock_auth: + with patch( + "litellm.proxy.auth.auth_checks.common_checks", new_callable=AsyncMock + ) as mock_auth: mock_auth.return_value = mock_user_auth result = await extract_mcp_auth_context(mock_scope, "/mcp/sse") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_conversions.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_conversions.py index 4b0d38ad95d..5c598b90853 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_conversions.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_conversions.py @@ -22,7 +22,10 @@ def test_convert_image_content(): mock_image.mimeType = "image/jpeg" result = _convert_single_content(mock_image) - assert result == {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,base64data"}} + assert result == { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,base64data"}, + } def test_convert_audio_content(): @@ -32,7 +35,10 @@ def test_convert_audio_content(): mock_audio.mimeType = "audio/mp3" result = _convert_single_content(mock_audio) - assert result == {"type": "input_audio", "input_audio": {"data": "audiobase64", "format": "mp3"}} + assert result == { + "type": "input_audio", + "input_audio": {"data": "audiobase64", "format": "mp3"}, + } def test_convert_list_content(): @@ -64,7 +70,10 @@ def test_resolve_model_from_hints(): original_router = proxy_server.llm_router try: proxy_server.llm_router = MagicMock() - proxy_server.llm_router.get_model_names.return_value = ["gpt-4", "claude-3-5-sonnet"] + proxy_server.llm_router.get_model_names.return_value = [ + "gpt-4", + "claude-3-5-sonnet", + ] result = _resolve_model_from_preferences(mock_prefs) assert result == "claude-3-5-sonnet" finally: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_handler.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_handler.py index cd5e871d480..06d81928635 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_handler.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_handler.py @@ -78,7 +78,9 @@ def _make_params( return params -def _make_completion_response(content="Hello!", model="gpt-4o-mini", finish_reason="stop", tool_calls=None): +def _make_completion_response( + content="Hello!", model="gpt-4o-mini", finish_reason="stop", tool_calls=None +): """Create a mock litellm completion response.""" response = MagicMock() choice = MagicMock() @@ -155,7 +157,9 @@ class TestConvertMCPMessagesToOpenAI: ) user_msg = _make_sampling_message("user", _make_text_content("Hi")) - assistant_msg = _make_sampling_message("assistant", _make_text_content("Hello!")) + assistant_msg = _make_sampling_message( + "assistant", _make_text_content("Hello!") + ) result = _convert_mcp_messages_to_openai([user_msg, assistant_msg]) assert len(result) == 2 assert result[0]["role"] == "user" @@ -173,7 +177,9 @@ class TestResolveModel: _resolve_model_from_preferences, ) - result = _resolve_model_from_preferences(None, default_model="claude-3.5-sonnet") + result = _resolve_model_from_preferences( + None, default_model="claude-3.5-sonnet" + ) assert result == "claude-3.5-sonnet" @patch("litellm.proxy.proxy_server.llm_router", None) @@ -309,7 +315,9 @@ class TestConvertOpenAIResponseToMCPResult: _convert_openai_response_to_mcp_result, ) - response = _make_completion_response(content="Partial...", finish_reason="length") + response = _make_completion_response( + content="Partial...", finish_reason="length" + ) result = _convert_openai_response_to_mcp_result(response, "gpt-4o") assert result.stopReason == "maxTokens" @@ -322,7 +330,9 @@ class TestConvertOpenAIResponseToMCPResult: tc.id = "call_123" tc.function.name = "get_weather" tc.function.arguments = '{"city": "NYC"}' - response = _make_completion_response(content=None, finish_reason="tool_calls", tool_calls=[tc]) + response = _make_completion_response( + content=None, finish_reason="tool_calls", tool_calls=[tc] + ) result = _convert_openai_response_to_mcp_result(response, "gpt-4o") assert result.stopReason == "toolUse" assert isinstance(result.content, list) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 43178c0da13..ffd5c0e41dc 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -35,14 +35,18 @@ from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPSer def _reload_mcp_manager_module(): utils_module = sys.modules["litellm.proxy._experimental.mcp_server.utils"] - manager_module = sys.modules["litellm.proxy._experimental.mcp_server.mcp_server_manager"] + manager_module = sys.modules[ + "litellm.proxy._experimental.mcp_server.mcp_server_manager" + ] importlib.reload(utils_module) reloaded = importlib.reload(manager_module) # After reload, server.py still holds a stale reference to the old # global_mcp_server_manager. Update it so tests that exercise server.py # functions (e.g. _get_tools_from_mcp_servers) use the fresh instance. server_module = sys.modules.get("litellm.proxy._experimental.mcp_server.server") - if server_module is not None and hasattr(server_module, "global_mcp_server_manager"): + if server_module is not None and hasattr( + server_module, "global_mcp_server_manager" + ): server_module.global_mcp_server_manager = reloaded.global_mcp_server_manager return reloaded @@ -303,7 +307,9 @@ class TestMCPServerManager: # Mock get_allowed_mcp_servers to return our test servers manager.get_allowed_mcp_servers = AsyncMock(return_value=["github", "zapier"]) - manager.get_mcp_server_by_id = MagicMock(side_effect=lambda x: server1 if x == "github" else server2) + manager.get_mcp_server_by_id = MagicMock( + side_effect=lambda x: server1 if x == "github" else server2 + ) # Mock _get_tools_from_server to return different results async def mock_get_tools_from_server( @@ -331,7 +337,9 @@ class TestMCPServerManager: "zapier": "zapier-api-key", } - result = await manager.list_tools(mcp_server_auth_headers=mcp_server_auth_headers) + result = await manager.list_tools( + mcp_server_auth_headers=mcp_server_auth_headers + ) # Verify that both servers were called with their specific auth headers assert len(result) == 3 # 2 from github + 1 from zapier @@ -402,7 +410,9 @@ class TestMCPServerManager: mcp_protocol_version=None, raw_headers=None, ): - assert mcp_auth_header == "server-specific-token" # Should use server-specific header + assert ( + mcp_auth_header == "server-specific-token" + ) # Should use server-specific header tool = MagicMock() tool.name = "github_tool_1" return [tool] @@ -433,7 +443,9 @@ class TestMCPServerManager: ) mock_client = AsyncMock() - mock_client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False)) + mock_client.call_tool = AsyncMock( + return_value=CallToolResult(content=[], isError=False) + ) captured_extra_headers = None async def capture_create_mcp_client( @@ -547,7 +559,9 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_resources = [Resource(name="file", uri="https://example.com/file")] mock_client.list_resources = AsyncMock(return_value=mock_resources) - prefixed_resources = [Resource(name="alias-server-file", uri="https://example.com/file")] + prefixed_resources = [ + Resource(name="alias-server-file", uri="https://example.com/file") + ] with ( patch.object( @@ -678,7 +692,9 @@ class TestMCPServerManager: mock_create_client.assert_called_once() called_kwargs = mock_create_client.call_args.kwargs assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "1"} - mock_client.read_resource.assert_awaited_once_with("https://example.com/resource") + mock_client.read_resource.assert_awaited_once_with( + "https://example.com/resource" + ) assert result is read_result @pytest.mark.asyncio @@ -725,7 +741,9 @@ class TestMCPServerManager: ) def raise_http_error(): - raise httpx.HTTPStatusError("unauthorized", request=request, response=response_obj) + raise httpx.HTTPStatusError( + "unauthorized", request=request, response=response_obj + ) response_obj.raise_for_status = MagicMock(side_effect=raise_http_error) @@ -843,7 +861,9 @@ class TestMCPServerManager: mcp_protocol_version=None, raw_headers=None, ): - assert mcp_auth_header == "server-specific-token" # Should use server-specific header via server_name + assert ( + mcp_auth_header == "server-specific-token" + ) # Should use server-specific header via server_name tool = MagicMock() tool.name = "github_tool_1" return [tool] @@ -910,7 +930,9 @@ class TestMCPServerManager: # Mock failed client.run_with_session mock_client = AsyncMock() - mock_client.run_with_session = AsyncMock(side_effect=Exception("Connection timeout")) + mock_client.run_with_session = AsyncMock( + side_effect=Exception("Connection timeout") + ) manager._create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check @@ -1032,7 +1054,9 @@ class TestMCPServerManager: # Capture the extra_headers passed to _create_mcp_client captured_extra_headers = None - async def capture_create_mcp_client(server, mcp_auth_header, extra_headers, stdio_env): + async def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env + ): nonlocal captured_extra_headers captured_extra_headers = extra_headers return mock_client @@ -1374,7 +1398,9 @@ class TestMCPServerManager: proxy_logging_obj = MagicMock() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock( + return_value={} + ) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) @@ -1418,8 +1444,13 @@ class TestMCPServerManager: ) assert exc_info.value.status_code == 403 - assert "Tool blocked_tool is not allowed for server test-server" in exc_info.value.detail["error"] - assert "Contact proxy admin to allow this tool" in exc_info.value.detail["error"] + assert ( + "Tool blocked_tool is not allowed for server test-server" + in exc_info.value.detail["error"] + ) + assert ( + "Contact proxy admin to allow this tool" in exc_info.value.detail["error"] + ) @pytest.mark.asyncio async def test_pre_call_tool_check_disallowed_tools_list_allows_tool(self): @@ -1443,7 +1474,9 @@ class TestMCPServerManager: proxy_logging_obj = MagicMock() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock( + return_value={} + ) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) @@ -1487,8 +1520,13 @@ class TestMCPServerManager: ) assert exc_info.value.status_code == 403 - assert "Tool banned_tool is not allowed for server test-server" in exc_info.value.detail["error"] - assert "Contact proxy admin to allow this tool" in exc_info.value.detail["error"] + assert ( + "Tool banned_tool is not allowed for server test-server" + in exc_info.value.detail["error"] + ) + assert ( + "Contact proxy admin to allow this tool" in exc_info.value.detail["error"] + ) @pytest.mark.asyncio async def test_pre_call_tool_check_no_restrictions_allows_any_tool(self): @@ -1512,7 +1550,9 @@ class TestMCPServerManager: proxy_logging_obj = MagicMock() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock( + return_value={} + ) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) @@ -1549,7 +1589,9 @@ class TestMCPServerManager: proxy_logging_obj = MagicMock() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock( + return_value={} + ) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) @@ -1575,7 +1617,10 @@ class TestMCPServerManager: ) assert exc_info.value.status_code == 403 - assert "Tool tool3 is not allowed for server test-server" in exc_info.value.detail["error"] + assert ( + "Tool tool3 is not allowed for server test-server" + in exc_info.value.detail["error"] + ) async def test_get_tools_from_server_add_prefix(self): """Verify _get_tools_from_server respects add_prefix True/False.""" @@ -1606,7 +1651,9 @@ class TestMCPServerManager: assert tools_prefixed[0].name == "zapier-send_email" # Case 2: add_prefix=False (single-server) -> expect unprefixed - tools_unprefixed = await manager._get_tools_from_server(server, add_prefix=False) + tools_unprefixed = await manager._get_tools_from_server( + server, add_prefix=False + ) assert len(tools_unprefixed) == 1 assert tools_unprefixed[0].name == "send_email" @@ -1641,9 +1688,13 @@ class TestMCPServerManager: # Mapping should include both original and prefixed names -> resolves calls either way assert manager.tool_name_to_mcp_server_name_mapping["create_issue"] == "jira" - assert manager.tool_name_to_mcp_server_name_mapping["jira-create_issue"] == "jira" + assert ( + manager.tool_name_to_mcp_server_name_mapping["jira-create_issue"] == "jira" + ) assert manager.tool_name_to_mcp_server_name_mapping["close_issue"] == "jira" - assert manager.tool_name_to_mcp_server_name_mapping["jira-close_issue"] == "jira" + assert ( + manager.tool_name_to_mcp_server_name_mapping["jira-close_issue"] == "jira" + ) def test_get_mcp_server_from_tool_name_with_prefixed_and_unprefixed(self): """After mapping is populated, manager resolves both prefixed and unprefixed tool names to the same server.""" @@ -1673,7 +1724,9 @@ class TestMCPServerManager: assert resolved_server_unpref.server_id == server.server_id # Prefixed resolution - resolved_server_pref = manager._get_mcp_server_from_tool_name("zapier-create_zap") + resolved_server_pref = manager._get_mcp_server_from_tool_name( + "zapier-create_zap" + ) assert resolved_server_pref is not None assert resolved_server_pref.server_id == server.server_id @@ -1718,7 +1771,9 @@ class TestMCPServerManager: new=AsyncMock(return_value=[tool1, tool2, tool3]), ): # Call the REST endpoint helper - filtered_response = await _get_tools_for_single_server(server, server_auth_header=None) + filtered_response = await _get_tools_for_single_server( + server, server_auth_header=None + ) # Verify only allowed tools are in the response assert len(filtered_response) == 2 @@ -1768,7 +1823,9 @@ class TestMCPServerManager: new=AsyncMock(return_value=[tool1, tool2, tool3]), ): # Call the REST endpoint helper - all_tools_response = await _get_tools_for_single_server(server, server_auth_header=None) + all_tools_response = await _get_tools_for_single_server( + server, server_auth_header=None + ) # Verify all tools are returned (no filtering) assert len(all_tools_response) == 3 @@ -1813,7 +1870,9 @@ class TestMCPServerManager: new=AsyncMock(return_value=[tool1, tool2]), ): # Call the REST endpoint helper - all_tools_response = await _get_tools_for_single_server(server, server_auth_header=None) + all_tools_response = await _get_tools_for_single_server( + server, server_auth_header=None + ) # Verify all tools are returned (no filtering) assert len(all_tools_response) == 2 @@ -1883,7 +1942,9 @@ class TestMCPServerManager: ) proxy_logging = MagicMock() - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock( + return_value={} + ) proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) @@ -1926,7 +1987,9 @@ class TestMCPServerManager: ) proxy_logging = MagicMock() - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock( + return_value={} + ) proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) @@ -2071,7 +2134,9 @@ class TestMCPServerManager: proxy_logging_obj = MagicMock() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock( + return_value={} + ) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) @@ -2107,8 +2172,13 @@ class TestMCPServerManager: ) assert exc_info.value.status_code == 403 - assert "Tool deletepet is not allowed for server my_api_mcp" in exc_info.value.detail["error"] - assert "Contact proxy admin to allow this tool" in exc_info.value.detail["error"] + assert ( + "Tool deletepet is not allowed for server my_api_mcp" + in exc_info.value.detail["error"] + ) + assert ( + "Contact proxy admin to allow this tool" in exc_info.value.detail["error"] + ) @pytest.mark.asyncio async def test_call_tool_without_broken_pipe_error(self): @@ -2133,7 +2203,9 @@ class TestMCPServerManager: # Register the server and map a tool to it manager.registry = {"test-server": server} manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server" - manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server" + manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = ( + "test-server" + ) # Create mock client that tracks call_tool usage mock_client = AsyncMock() @@ -2157,7 +2229,9 @@ class TestMCPServerManager: # Mock proxy logging proxy_logging_obj = MagicMock() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock( + return_value={} + ) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) @@ -2222,7 +2296,9 @@ class TestMCPServerManager: # Verify MCPRequestHandler.get_allowed_mcp_servers was called with user_api_key_auth mock_get_allowed.assert_called_once() call_args = mock_get_allowed.call_args - assert call_args[0][0] is user_api_key_auth # First positional arg should be user_api_key_auth + assert ( + call_args[0][0] is user_api_key_auth + ) # First positional arg should be user_api_key_auth assert call_args[0][0].user_id == "user-123" assert call_args[0][0].object_permission_id == "perm_123" assert call_args[0][0].object_permission is not None @@ -2467,7 +2543,10 @@ class TestMCPServerManagerUpstreamInstructionsCache: def test_get_returns_none_when_empty(self): """Empty cache returns None for any key.""" manager = MCPServerManager() - assert manager._upstream_initialize_instructions_by_server_id.get("nonexistent") is None + assert ( + manager._upstream_initialize_instructions_by_server_id.get("nonexistent") + is None + ) def test_remember_stores_stripped_value(self): """_remember_upstream_initialize_instructions stores a stripped string.""" @@ -2475,7 +2554,9 @@ class TestMCPServerManagerUpstreamInstructionsCache: fake_server = MagicMock(server_id="srv") fake_client = MagicMock(_last_initialize_instructions=" hello \n") manager._remember_upstream_initialize_instructions(fake_server, fake_client) - assert manager._upstream_initialize_instructions_by_server_id.get("srv") == "hello" + assert ( + manager._upstream_initialize_instructions_by_server_id.get("srv") == "hello" + ) def test_remember_ignores_empty_string(self): """Whitespace-only instructions are not stored.""" @@ -2553,7 +2634,9 @@ class TestMCPServerManagerExpandPermissionList: def test_expands_server_name(self): manager = MCPServerManager() - manager.config_mcp_servers["id-usw1"] = self._make_server("id-usw1", server_name="a") + manager.config_mcp_servers["id-usw1"] = self._make_server( + "id-usw1", server_name="a" + ) assert manager.expand_permission_list(["a"]) == ["id-usw1"] @@ -2577,7 +2660,9 @@ class TestMCPServerManagerExpandPermissionList: def test_name_collision_expands_to_all_matches(self): """Two servers sharing a server_name both resolve — the documented behavior.""" manager = MCPServerManager() - manager.config_mcp_servers["id-config"] = self._make_server("id-config", server_name="shared") + manager.config_mcp_servers["id-config"] = self._make_server( + "id-config", server_name="shared" + ) manager.registry["id-db"] = self._make_server("id-db", server_name="shared") assert sorted(manager.expand_permission_list(["shared"])) == [ @@ -2587,7 +2672,9 @@ class TestMCPServerManagerExpandPermissionList: def test_searches_config_and_registry_union(self): manager = MCPServerManager() - manager.config_mcp_servers["cfg-id"] = self._make_server("cfg-id", server_name="a") + manager.config_mcp_servers["cfg-id"] = self._make_server( + "cfg-id", server_name="a" + ) manager.registry["reg-id"] = self._make_server("reg-id", server_name="b") assert manager.expand_permission_list(["a"]) == ["cfg-id"] @@ -2599,15 +2686,23 @@ class TestMCPServerManagerExpandPermissionList: servers whose server_name happens to equal that id. """ manager = MCPServerManager() - manager.config_mcp_servers["id-1"] = self._make_server("id-1", server_name="other_name") - manager.config_mcp_servers["id-2"] = self._make_server("id-2", server_name="id-1") + manager.config_mcp_servers["id-1"] = self._make_server( + "id-1", server_name="other_name" + ) + manager.config_mcp_servers["id-2"] = self._make_server( + "id-2", server_name="id-1" + ) assert manager.expand_permission_list(["id-1"]) == ["id-1"] def test_mixed_ids_and_names_in_same_list(self): manager = MCPServerManager() - manager.config_mcp_servers["uuid-1"] = self._make_server("uuid-1", server_name="a") - manager.config_mcp_servers["uuid-2"] = self._make_server("uuid-2", server_name="b") + manager.config_mcp_servers["uuid-1"] = self._make_server( + "uuid-1", server_name="a" + ) + manager.config_mcp_servers["uuid-2"] = self._make_server( + "uuid-2", server_name="b" + ) # ["uuid-1", "b"] -> uuid-1 passes through, "b" resolves to uuid-2 assert sorted(manager.expand_permission_list(["uuid-1", "b"])) == [ @@ -2618,7 +2713,9 @@ class TestMCPServerManagerExpandPermissionList: def test_deduplicates_overlapping_id_and_name_entries(self): """If a list references the same server by both id and name, return it once.""" manager = MCPServerManager() - manager.config_mcp_servers["uuid-1"] = self._make_server("uuid-1", server_name="a") + manager.config_mcp_servers["uuid-1"] = self._make_server( + "uuid-1", server_name="a" + ) assert manager.expand_permission_list(["uuid-1", "a"]) == ["uuid-1"] @@ -2628,10 +2725,14 @@ class TestMCPServerManagerExpandPermissionList: the cross-region portability the customer is asking for. """ usw1 = MCPServerManager() - usw1.config_mcp_servers["hash-usw1"] = self._make_server("hash-usw1", server_name="a") + usw1.config_mcp_servers["hash-usw1"] = self._make_server( + "hash-usw1", server_name="a" + ) usc1 = MCPServerManager() - usc1.config_mcp_servers["hash-usc1"] = self._make_server("hash-usc1", server_name="a") + usc1.config_mcp_servers["hash-usc1"] = self._make_server( + "hash-usc1", server_name="a" + ) assert usw1.expand_permission_list(["a"]) == ["hash-usw1"] assert usc1.expand_permission_list(["a"]) == ["hash-usc1"] @@ -2660,14 +2761,18 @@ class TestMCPServerManagerExpandToolPermissions: concrete server_id, otherwise `.get(server_id)` misses and the tool restriction is silently dropped (caller treats None as allow-all).""" manager = MCPServerManager() - manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="my-alias") + manager.config_mcp_servers["uuid-a"] = self._make_server( + "uuid-a", server_name="my-alias" + ) result = manager.expand_tool_permissions({"my-alias": ["read_file"]}) assert result == {"uuid-a": ["read_file"]} def test_passes_through_existing_server_id_key(self): manager = MCPServerManager() - manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha") + manager.config_mcp_servers["uuid-a"] = self._make_server( + "uuid-a", server_name="alpha" + ) result = manager.expand_tool_permissions({"uuid-a": ["read_file"]}) assert result == {"uuid-a": ["read_file"]} @@ -2686,7 +2791,9 @@ class TestMCPServerManagerExpandToolPermissions: """Two servers sharing a server_name both match; their tool lists get the restriction (matches the list-expansion collision semantics).""" manager = MCPServerManager() - manager.config_mcp_servers["uuid-1"] = self._make_server("uuid-1", server_name="shared") + manager.config_mcp_servers["uuid-1"] = self._make_server( + "uuid-1", server_name="shared" + ) manager.registry["uuid-2"] = self._make_server("uuid-2", server_name="shared") result = manager.expand_tool_permissions({"shared": ["read_file"]}) @@ -2699,9 +2806,13 @@ class TestMCPServerManagerExpandToolPermissions: both refer to the same server, the tool lists are unioned rather than one overwriting the other.""" manager = MCPServerManager() - manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alias-a") + manager.config_mcp_servers["uuid-a"] = self._make_server( + "uuid-a", server_name="alias-a" + ) - result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["write_file"]}) + result = manager.expand_tool_permissions( + {"uuid-a": ["read_file"], "alias-a": ["write_file"]} + ) assert sorted(result["uuid-a"]) == ["read_file", "write_file"] @pytest.mark.asyncio @@ -2740,7 +2851,9 @@ class TestMCPServerManagerExpandToolPermissions: "litellm.proxy._experimental.mcp_server.elicitation_handler.handle_elicitation_request", new_callable=AsyncMock, ) as mock_handle, - patch("litellm.proxy._experimental.mcp_server.server.get_active_mcp_session") as mock_get_session, + patch( + "litellm.proxy._experimental.mcp_server.server.get_active_mcp_session" + ) as mock_get_session, ): mock_session = MagicMock() mock_session.capabilities = "test_capabilities"