diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 618ac4dba49..97e6410447b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -120,9 +120,7 @@ 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, @@ -156,7 +154,7 @@ def _deserialize_json_dict(data: Any) -> Optional[Dict[str, str]]: return data -def _create_sampling_callback(): +def _create_sampling_callback(user_api_key_auth: Optional[Any] = None): """ Create a sampling callback for MCP ClientSession. Returns a callable that handles sampling/createMessage requests from @@ -172,11 +170,14 @@ def _create_sampling_callback(): get_active_auth_context, ) + auth_context = get_active_auth_context() + 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, params=params, default_model=getattr(litellm, "default_mcp_sampling_model", None), - user_api_key_auth=get_active_auth_context(), + user_api_key_auth=resolved_auth, ) return _sampling_callback @@ -199,11 +200,7 @@ 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, @@ -245,14 +242,10 @@ 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]: """ @@ -296,15 +289,10 @@ 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 @@ -322,15 +310,10 @@ 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 @@ -365,9 +348,7 @@ 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 ) @@ -399,9 +380,7 @@ 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), @@ -410,9 +389,7 @@ 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), @@ -428,24 +405,18 @@ 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. @@ -481,9 +452,7 @@ 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) @@ -518,8 +487,7 @@ class MCPServerManager: ) verbose_logger.debug( - f"Using headers for OpenAPI tools (excluding sensitive values): " - f"{list(headers.keys())}" + f"Using headers for OpenAPI tools (excluding sensitive values): {list(headers.keys())}" ) # Extract and register tools from OpenAPI paths @@ -537,20 +505,14 @@ 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( @@ -562,9 +524,7 @@ 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 @@ -577,26 +537,16 @@ 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): @@ -610,9 +560,7 @@ 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, @@ -622,12 +570,8 @@ 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 @@ -674,9 +618,7 @@ 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: @@ -684,9 +626,7 @@ 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: @@ -696,20 +636,14 @@ 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, @@ -725,16 +659,12 @@ 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, @@ -742,17 +672,11 @@ 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), @@ -771,9 +695,7 @@ 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, @@ -814,15 +736,9 @@ 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. @@ -847,23 +763,13 @@ 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 @@ -879,9 +785,7 @@ 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)}.") @@ -958,15 +862,12 @@ 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, @@ -1004,18 +905,12 @@ 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. @@ -1057,9 +952,7 @@ 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( @@ -1121,9 +1014,7 @@ 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 ######################################################### @@ -1162,6 +1053,7 @@ class MCPServerManager: mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, extra_headers: Optional[Dict[str, str]] = None, stdio_env: Optional[Dict[str, str]] = None, + user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> MCPClient: """ Create an MCPClient instance for the given server. @@ -1185,15 +1077,13 @@ class MCPServerManager: transport = server.transport or MCPTransport.sse # Create sampling and elicitation callbacks for this client - sampling_cb = _create_sampling_callback() + sampling_cb = _create_sampling_callback(user_api_key_auth=user_api_key_auth) elicitation_cb = _create_elicitation_callback() # Handle stdio transport if transport == MCPTransport.stdio: resolved_env = ( - stdio_env - if stdio_env is not None - else (dict(server.env) if server.env 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. @@ -1311,9 +1201,7 @@ 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 @@ -1323,9 +1211,7 @@ 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 ) @@ -1336,16 +1222,12 @@ 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( @@ -1389,16 +1271,12 @@ 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( @@ -1433,16 +1311,12 @@ 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( @@ -1484,9 +1358,7 @@ 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( @@ -1576,13 +1448,11 @@ 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 @@ -1590,9 +1460,7 @@ 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, @@ -1604,16 +1472,12 @@ 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: @@ -1623,14 +1487,10 @@ 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 @@ -1639,8 +1499,7 @@ 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") @@ -1675,23 +1534,15 @@ 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: @@ -1728,9 +1579,7 @@ 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: @@ -1744,9 +1593,7 @@ 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") @@ -1846,9 +1693,7 @@ 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. @@ -1871,22 +1716,16 @@ 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. @@ -1916,9 +1755,7 @@ 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( @@ -1945,9 +1782,7 @@ 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( @@ -1959,17 +1794,11 @@ 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( @@ -1985,9 +1814,7 @@ 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) @@ -2002,20 +1829,14 @@ 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. @@ -2042,18 +1863,14 @@ 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( @@ -2217,38 +2034,20 @@ 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_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 + 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), "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: @@ -2260,11 +2059,7 @@ 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"): @@ -2309,9 +2104,7 @@ 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( @@ -2334,6 +2127,7 @@ class MCPServerManager: proxy_logging_obj: Optional[ProxyLogging], host_progress_callback: Optional[Callable] = None, hook_extra_headers: Optional[Dict[str, str]] = None, + user_api_key_auth: Optional[Any] = None, ) -> CallToolResult: """ Call a regular MCP tool using the MCP client. @@ -2366,16 +2160,12 @@ 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: @@ -2390,9 +2180,7 @@ 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 @@ -2438,6 +2226,7 @@ class MCPServerManager: mcp_auth_header=server_auth_header, extra_headers=extra_headers, stdio_env=stdio_env, + user_api_key_auth=user_api_key_auth, ) call_tool_params = MCPCallToolRequestParams( @@ -2446,13 +2235,9 @@ 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) @@ -2462,9 +2247,7 @@ 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 @@ -2548,11 +2331,7 @@ 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: @@ -2568,8 +2347,7 @@ class MCPServerManager: oauth2_headers = stored_headers except Exception as _lookup_exc: verbose_logger.debug( - "call_tool: per-user token lookup failed for " - "user=%s server=%s: %s", + "call_tool: per-user token lookup failed for user=%s server=%s: %s", user_id, mcp_server.server_id, _lookup_exc, @@ -2577,9 +2355,7 @@ 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 " @@ -2588,13 +2364,8 @@ 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: - # For regular MCP servers, use the MCP client return await self._call_regular_mcp_tool( mcp_server=mcp_server, original_tool_name=name, @@ -2607,6 +2378,7 @@ class MCPServerManager: proxy_logging_obj=proxy_logging_obj, host_progress_callback=host_progress_callback, hook_extra_headers=hook_result.get("extra_headers"), + user_api_key_auth=user_api_key_auth, ) # For OpenAPI tools, await outside the client context @@ -2625,9 +2397,7 @@ 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 ######################################################### @@ -2640,9 +2410,7 @@ 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)}" @@ -2679,16 +2447,12 @@ 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): ( @@ -2698,13 +2462,9 @@ 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 @@ -2719,9 +2479,7 @@ 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 @@ -2759,18 +2517,14 @@ 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 = [] @@ -2792,9 +2546,7 @@ 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. @@ -2812,9 +2564,7 @@ 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]: @@ -2863,9 +2613,7 @@ 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) @@ -2905,9 +2653,7 @@ 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. @@ -2942,9 +2688,7 @@ 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. @@ -2955,11 +2699,7 @@ 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, @@ -2988,9 +2728,7 @@ 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")) @@ -3063,15 +2801,11 @@ 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" @@ -3084,9 +2818,7 @@ 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, @@ -3175,9 +2907,7 @@ 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, @@ -3238,9 +2968,7 @@ 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 3150ebf6e29..5fe56bcc81c 100644 --- a/litellm/proxy/_experimental/mcp_server/sampling_handler.py +++ b/litellm/proxy/_experimental/mcp_server/sampling_handler.py @@ -117,7 +117,7 @@ def _convert_mcp_content_to_openai( return _convert_single_content(content) -def _convert_single_content(content: Any) -> Union[str, Dict[str, Any]]: +def _convert_single_content(content: Any) -> Dict[str, Any]: """Convert a single MCP content item to OpenAI format.""" content_type = getattr(content, "type", None) if content_type == "text": @@ -155,11 +155,7 @@ def _convert_single_content(content: Any) -> Union[str, 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 @@ -246,9 +242,7 @@ 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), }, } ) @@ -275,11 +269,7 @@ 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) @@ -368,7 +358,7 @@ def _convert_openai_response_to_mcp_result( tool_calls = getattr(message, "tool_calls", None) if tool_calls: # Build ToolUseContent items - content_parts = [] + content_parts: "List[Any]" = [] # Include text content if present if message.content: content_parts.append(TextContent(type="text", text=message.content)) @@ -488,8 +478,7 @@ async def handle_sampling_create_message( completion_kwargs["metadata"]["user_api_key_team_id"] = team_id verbose_logger.debug( - "MCP sampling: calling litellm.acompletion with model=%s, " - "num_messages=%d, has_tools=%s", + "MCP sampling: calling litellm.acompletion with model=%s, num_messages=%d, has_tools=%s", model, len(openai_messages), bool(openai_tools), diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index b31103446a5..7214e92481a 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -76,9 +76,7 @@ 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() @@ -103,8 +101,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}") @@ -272,9 +270,7 @@ 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.""" @@ -327,12 +323,8 @@ 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}" ) @@ -348,9 +340,7 @@ 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)}") @@ -359,9 +349,7 @@ 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(name: str, arguments: Dict[str, Any] | None) -> CallToolResult: """ Call a specific tool with the provided arguments Args: @@ -406,18 +394,12 @@ 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: @@ -469,11 +451,7 @@ 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: @@ -506,12 +484,8 @@ 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}" ) @@ -525,9 +499,7 @@ 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)}") @@ -536,9 +508,7 @@ 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 Args: @@ -557,9 +527,7 @@ 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, @@ -584,12 +552,8 @@ 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}" ) @@ -601,9 +565,7 @@ 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)}") @@ -622,12 +584,8 @@ 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}" ) @@ -640,14 +598,11 @@ if MCP_AVAILABLE: raw_headers=raw_headers, ) verbose_logger.info( - "MCP list_resource_templates - Successfully returned " - f"{len(resource_templates)} resource templates" + f"MCP list_resource_templates - Successfully returned {len(resource_templates)} resource templates" ) 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() @@ -707,10 +662,8 @@ 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: @@ -718,9 +671,7 @@ 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 @@ -767,17 +718,11 @@ 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 @@ -838,15 +783,11 @@ 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, @@ -864,9 +805,7 @@ 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: @@ -916,8 +855,7 @@ if MCP_AVAILABLE: cached_token = await mcp_per_user_token_cache.get(user_id, server_id) if cached_token is not None: verbose_logger.debug( - "_get_user_oauth_extra_headers_from_db: Redis hit for " - "user=%s server=%s", + "_get_user_oauth_extra_headers_from_db: Redis hit for user=%s server=%s", user_id, server_id, ) @@ -933,15 +871,12 @@ 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): verbose_logger.debug( - "_get_user_oauth_extra_headers_from_db: token expired for " - "user=%s server=%s — attempting refresh", + "_get_user_oauth_extra_headers_from_db: token expired for user=%s server=%s — attempting refresh", user_id, server_id, ) @@ -963,8 +898,7 @@ if MCP_AVAILABLE: ) except Exception as refresh_exc: verbose_logger.warning( - "_get_user_oauth_extra_headers_from_db: refresh failed " - "for user=%s server=%s: %s", + "_get_user_oauth_extra_headers_from_db: refresh failed for user=%s server=%s: %s", user_id, server_id, refresh_exc, @@ -991,21 +925,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( - "_get_user_oauth_extra_headers_from_db: failed to retrieve credential for " - "user=%s server=%s: %s", + "_get_user_oauth_extra_headers_from_db: failed to retrieve credential for user=%s server=%s: %s", user_id, server_id, e, @@ -1018,9 +947,7 @@ 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: @@ -1035,9 +962,7 @@ 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( @@ -1060,9 +985,7 @@ 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 @@ -1082,21 +1005,13 @@ 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: @@ -1155,9 +1070,7 @@ 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", } @@ -1191,9 +1104,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: @@ -1207,9 +1120,7 @@ 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( @@ -1218,14 +1129,9 @@ 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( @@ -1270,15 +1176,11 @@ 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] @@ -1310,9 +1212,7 @@ 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 @@ -1321,9 +1221,7 @@ 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, @@ -1332,9 +1230,7 @@ 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( @@ -1383,17 +1279,11 @@ 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( @@ -1431,16 +1321,10 @@ 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( @@ -1470,14 +1354,12 @@ 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( @@ -1543,11 +1425,7 @@ 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) @@ -1561,9 +1439,7 @@ 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( @@ -1604,13 +1480,9 @@ 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 @@ -1645,13 +1517,9 @@ 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 @@ -1676,13 +1544,9 @@ 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( @@ -1731,9 +1595,7 @@ 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( @@ -1790,9 +1652,7 @@ 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) @@ -1851,9 +1711,7 @@ 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 @@ -1908,33 +1766,25 @@ 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, - ) - ) - litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get( - "litellm_logging_obj", None + 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) 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: @@ -1971,9 +1821,7 @@ 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: @@ -2027,25 +1875,17 @@ 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( @@ -2093,9 +1933,7 @@ 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( @@ -2167,8 +2005,7 @@ if MCP_AVAILABLE: raise HTTPException( status_code=400, detail=( - "Multiple MCP servers configured; read_resource currently " - "supports exactly one allowed server." + "Multiple MCP servers configured; read_resource currently supports exactly one allowed server." ), ) server = allowed_mcp_servers[0] @@ -2289,21 +2126,15 @@ 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] @@ -2393,8 +2224,7 @@ if MCP_AVAILABLE: return False except Exception: verbose_logger.debug( - "Unable to inspect active MCP sessions for '%s'. " - "Deferring to session manager.", + "Unable to inspect active MCP sessions for '%s'. Deferring to session manager.", _session_id, ) return False @@ -2402,8 +2232,7 @@ if MCP_AVAILABLE: method = scope.get("method", "").upper() if method == "DELETE": verbose_logger.info( - "DELETE request for non-existent MCP session '%s'. " - "Returning success (idempotent DELETE).", + "DELETE request for non-existent MCP session '%s'. Returning success (idempotent DELETE).", _session_id, ) success_response = JSONResponse( @@ -2418,11 +2247,7 @@ 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( @@ -2454,11 +2279,7 @@ 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: @@ -2479,9 +2300,7 @@ 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", "") @@ -2495,36 +2314,29 @@ 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=" - f"{base_url}/.well-known/oauth-authorization-server/{server_name}" + f"Bearer authorization_uri={base_url}/.well-known/oauth-authorization-server/{server_name}" ) raise HTTPException( status_code=401, @@ -2532,18 +2344,12 @@ if MCP_AVAILABLE: 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, @@ -2573,9 +2379,7 @@ 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 @@ -2601,15 +2405,10 @@ 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(scope: Scope, receive: Receive, send: Send) -> None: - """Handle MCP requests through SSE.""" - async def handle_sse_mcp_endpoint(request: StarletteRequest): """ Handle MCP SSE GET requests. @@ -2632,9 +2431,7 @@ 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}" ) @@ -2670,12 +2467,8 @@ 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(request.scope, request.receive, request._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 @@ -2698,9 +2491,7 @@ if MCP_AVAILABLE: ) await error_response(request.scope, request.receive, request._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. @@ -2734,9 +2525,7 @@ if MCP_AVAILABLE: # and a FastAPI POST route for the POST messages endpoint. from starlette.routing import Route as StarletteRoute - app.routes.insert( - 0, StarletteRoute("/sse", endpoint=handle_sse_mcp_endpoint, methods=["GET"]) - ) + app.routes.insert(0, StarletteRoute("/sse", endpoint=handle_sse_mcp_endpoint, methods=["GET"])) from starlette.responses import Response as StarletteResponse class NoOpResponse(StarletteResponse): @@ -2770,9 +2559,7 @@ if MCP_AVAILABLE: client_ip=_sse_client_ip, ) except Exception as e: - verbose_logger.warning( - f"Failed to extract auth context in POST /messages: {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. @@ -2794,7 +2581,6 @@ if MCP_AVAILABLE: # StreamableHTTP catch-all mounts (must come after specific routes) app.mount("/mcp", handle_streamable_http_mcp) app.mount("/{mcp_server_name}/mcp", handle_streamable_http_mcp) - app.mount("/sse", handle_sse_mcp) app.mount("/", handle_streamable_http_mcp) app.add_middleware(AuthContextMiddleware) @@ -2884,13 +2670,9 @@ 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/mcp_test_config.yaml b/tests/mcp_sampling_elicitation/mcp_test_config.yaml index e50ad9bf751..d94a472ccdd 100644 --- a/tests/mcp_sampling_elicitation/mcp_test_config.yaml +++ b/tests/mcp_sampling_elicitation/mcp_test_config.yaml @@ -8,6 +8,6 @@ litellm_settings: mcp_servers: test_server: transport: stdio - command: "c:\\Users\\DELL\\Desktop\\litellm\\.venv\\Scripts\\python.exe" - args: ["c:\\Users\\DELL\\Desktop\\litellm\\tests\\mcp_sampling_elicitation\\custom_mcp_server.py"] + command: "python" + args: ["tests/mcp_sampling_elicitation/custom_mcp_server.py"] allow_all_keys: true diff --git a/tests/mcp_sampling_elicitation/test_live_mcp.py b/tests/mcp_sampling_elicitation/test_live_mcp.py index 22d3d6cad8e..4b15cac7342 100644 --- a/tests/mcp_sampling_elicitation/test_live_mcp.py +++ b/tests/mcp_sampling_elicitation/test_live_mcp.py @@ -1,14 +1,22 @@ import asyncio +import os +import logging from mcp.client.sse import sse_client from mcp.client.session import ClientSession from mcp.types import ElicitResult +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + async def main(): - print("Connecting to LiteLLM Proxy via SSE...") + if os.environ.get("RUN_LIVE_MCP_TEST") != "1": + logger.info("Skipping live integration test. Set RUN_LIVE_MCP_TEST=1 to run.") + return + logger.info("Connecting to LiteLLM Proxy via SSE...") async def my_elicitation_callback(context, params): - print(f"\n[CLIENT] Received elicitation request from upstream!") + logger.info("\n[CLIENT] Received elicitation request from upstream!") # We will simulate the user filling out the form user_response = { @@ -16,31 +24,28 @@ async def main(): "adjective": "suspenseful", } - print(f"[CLIENT] User is filling the form with: {user_response}") + logger.info(f"[CLIENT] User is filling the form with: {user_response}") return ElicitResult(action="accept", content=user_response) - async with sse_client( - "http://localhost:4000/mcp/sse", headers={"Authorization": "Bearer sk-1234"} - ) as (read_stream, write_stream): - print("SSE connection established.") - async with ClientSession( - read_stream, write_stream, elicitation_callback=my_elicitation_callback - ) as session: + 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: await session.initialize() - print("Initialized!") + logger.info("Initialized!") - print("\n--- Testing Complex Pipeline (Elicitation + Sampling) ---") - print("Calling 'test_server-test_complex_pipeline'...") + 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={} - ) - print("\nFINAL TOOL RESULT:") - print("==================") - print(result.content[0].text) + result = await session.call_tool("test_server-test_complex_pipeline", arguments={}) + logger.info("\nFINAL TOOL RESULT:") + logger.info("==================") + logger.info(result.content[0].text) except Exception as e: - print(f"Error calling test_complex_pipeline: {e}") + logger.info(f"Error calling test_complex_pipeline: {e}") if __name__ == "__main__": diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 6af07585796..db1830a3341 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -445,64 +445,6 @@ async def test_streamable_http_mcp_handler_mock(): mock_session_manager.handle_request.assert_called_once() -@pytest.mark.asyncio -async def test_sse_mcp_handler_mock(): - """Test the SSE MCP handler functionality""" - from litellm.proxy._types import UserAPIKeyAuth - - # Mock the SSE session manager and its methods - mock_sse_session_manager = AsyncMock() - mock_sse_session_manager.handle_request = AsyncMock() - - # Mock scope, receive, send with proper ASGI scope format - mock_scope = { - "type": "http", - "method": "GET", - "path": "/mcp/sse", - "headers": [(b"accept", b"text/event-stream")], - "query_string": b"", - "server": ("localhost", 8000), - "scheme": "http", - } - mock_receive = AsyncMock() - mock_send = AsyncMock() - - mock_auth_result = ( - UserAPIKeyAuth(), - None, - None, - {}, - {}, - [], - ) - - with ( - patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), - patch( - "litellm.proxy._experimental.mcp_server.server.sse_session_manager", - mock_sse_session_manager, - ), - patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - new=AsyncMock(return_value=mock_auth_result), - ), - patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", - ), - ): - from litellm.proxy._experimental.mcp_server.server import handle_sse_mcp - - # Call the handler - await handle_sse_mcp(mock_scope, mock_receive, mock_send) - - # Verify SSE session manager handle_request was called - mock_sse_session_manager.handle_request.assert_called_once_with( - mock_scope, mock_receive, mock_send - ) - def test_generate_stable_server_id(): """ diff --git a/tests/scratch/mcp_test_config.yaml b/tests/scratch/mcp_test_config.yaml deleted file mode 100644 index 7795f48e457..00000000000 --- a/tests/scratch/mcp_test_config.yaml +++ /dev/null @@ -1,13 +0,0 @@ -model_list: - - model_name: groq-model - litellm_params: - model: groq/llama-3.3-70b-versatile - api_key: os.environ/GROQ_API_KEY -litellm_settings: - default_mcp_sampling_model: groq-model -mcp_servers: - test_server: - transport: stdio - command: "c:\\Users\\DELL\\Desktop\\litellm\\.venv\\Scripts\\python.exe" - args: ["c:\\Users\\DELL\\Desktop\\litellm\\tests\\scratch\\custom_mcp_server.py"] - allow_all_keys: true \ No newline at end of file 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 new file mode 100644 index 00000000000..d1884e89d8c --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_auth.py @@ -0,0 +1,71 @@ +import pytest +from unittest.mock import AsyncMock, patch +from litellm.proxy._experimental.mcp_server.server import ( + set_auth_context, + get_active_auth_context, + extract_mcp_auth_context, + get_auth_context, + get_or_extract_auth_context, +) +from litellm.proxy._types import UserAPIKeyAuth + + +@pytest.mark.asyncio +async def test_auth_context_persistence(): + """Test that auth context is correctly set and retrieved.""" + auth_data = UserAPIKeyAuth(api_key="test-key") + + # Set context + set_auth_context(auth_data) + + # Retrieve context + retrieved = get_active_auth_context() + assert retrieved is not None + assert retrieved.user_api_key_auth.api_key == auth_data.api_key + + # Test get_auth_context tuple + auth_tuple = get_auth_context() + assert auth_tuple[0].api_key == auth_data.api_key + + +@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 + + auth_data = UserAPIKeyAuth(api_key="fallback-key") + auth_user = MCPAuthenticatedUser(user_api_key_auth=auth_data) + + # Set on server object + server._litellm_auth_context = auth_user + + # Ensure ContextVar is empty + token = auth_context_var.set(None) + try: + result = await get_or_extract_auth_context() + assert result[0].api_key == "fallback-key" + finally: + auth_context_var.reset(token) + + +@pytest.mark.asyncio +async def test_extract_mcp_auth_context_with_key(): + """Test extract_mcp_auth_context with a valid API key.""" + mock_scope = { + "type": "http", + "headers": [(b"authorization", b"Bearer sk-123")], + "path": "/mcp/sse", + "method": "GET", + "query_string": b"", + } + + mock_user_auth = UserAPIKeyAuth(api_key="sk-123") + + 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") + + # Returns (user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, oauth2_headers, raw_headers, client_ip) + assert result[0].api_key == mock_user_auth.api_key + assert result[5]["authorization"] == "Bearer sk-123" 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 new file mode 100644 index 00000000000..4b0d38ad95d --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_conversions.py @@ -0,0 +1,76 @@ +from unittest.mock import MagicMock +from litellm.proxy._experimental.mcp_server.sampling_handler import ( + _convert_mcp_content_to_openai, + _convert_single_content, + _resolve_model_from_preferences, +) + + +def test_convert_text_content(): + mock_text = MagicMock() + mock_text.type = "text" + mock_text.text = "hello world" + + result = _convert_single_content(mock_text) + assert result == {"type": "text", "text": "hello world"} + + +def test_convert_image_content(): + mock_image = MagicMock() + mock_image.type = "image" + mock_image.data = "base64data" + mock_image.mimeType = "image/jpeg" + + result = _convert_single_content(mock_image) + assert result == {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,base64data"}} + + +def test_convert_audio_content(): + mock_audio = MagicMock() + mock_audio.type = "audio" + mock_audio.data = "audiobase64" + mock_audio.mimeType = "audio/mp3" + + result = _convert_single_content(mock_audio) + assert result == {"type": "input_audio", "input_audio": {"data": "audiobase64", "format": "mp3"}} + + +def test_convert_list_content(): + mock_text = MagicMock() + mock_text.type = "text" + mock_text.text = "text" + + mock_image = MagicMock() + mock_image.type = "image" + mock_image.data = "img" + mock_image.mimeType = "image/png" + + result = _convert_mcp_content_to_openai([mock_text, mock_image]) + assert isinstance(result, list) + assert len(result) == 2 + assert result[0] == {"type": "text", "text": "text"} + assert result[1]["type"] == "image_url" + + +def test_resolve_model_from_hints(): + import litellm.proxy.proxy_server as proxy_server + + mock_prefs = MagicMock() + mock_hint = MagicMock() + mock_hint.name = "claude" + mock_prefs.hints = [mock_hint] + + # Save original + 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"] + result = _resolve_model_from_preferences(mock_prefs) + assert result == "claude-3-5-sonnet" + finally: + proxy_server.llm_router = original_router + + +def test_resolve_model_fallback(): + result = _resolve_model_from_preferences(None, default_model="fallback-model") + assert result == "fallback-model" 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 681be0056d4..cd5e871d480 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 @@ -3,9 +3,11 @@ Unit tests for the MCP Sampling Handler. Tests the sampling/createMessage handler that routes MCP sampling requests through litellm.acompletion(). """ -import json + from unittest.mock import AsyncMock, MagicMock, patch import pytest + + # ───────────────────────────────────────────────────────────── # Helper factories # ───────────────────────────────────────────────────────────── @@ -15,6 +17,8 @@ def _make_text_content(text: str): tc.type = "text" tc.text = text return tc + + def _make_image_content(data: str = "base64data", mime_type: str = "image/png"): """Create a mock ImageContent.""" ic = MagicMock() @@ -22,12 +26,16 @@ def _make_image_content(data: str = "base64data", mime_type: str = "image/png"): ic.data = data ic.mimeType = mime_type return ic + + def _make_sampling_message(role: str, content): """Create a mock SamplingMessage.""" msg = MagicMock() msg.role = role msg.content = content return msg + + def _make_model_preferences(hints=None, cost=None, speed=None, intelligence=None): """Create a mock ModelPreferences.""" prefs = MagicMock() @@ -36,11 +44,15 @@ def _make_model_preferences(hints=None, cost=None, speed=None, intelligence=None prefs.speedPriority = speed prefs.intelligencePriority = intelligence return prefs + + def _make_hint(name: str): """Create a mock model hint.""" hint = MagicMock() hint.name = name return hint + + def _make_params( messages=None, model_preferences=None, @@ -64,9 +76,9 @@ def _make_params( params.toolChoice = tool_choice params.metadata = metadata 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() @@ -78,34 +90,42 @@ def _make_completion_response( response.choices = [choice] response.model = model return response + + # ───────────────────────────────────────────────────────────── # Tests: Message conversion # ───────────────────────────────────────────────────────────── class TestConvertMCPMessagesToOpenAI: """Tests for _convert_mcp_messages_to_openai.""" + def test_should_convert_simple_text_message(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_messages_to_openai, ) + tc = _make_text_content("Hello") msg = _make_sampling_message("user", tc) result = _convert_mcp_messages_to_openai([msg]) assert len(result) == 1 assert result[0]["role"] == "user" + def test_should_add_system_prompt(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_messages_to_openai, ) + tc = _make_text_content("Hello") msg = _make_sampling_message("user", tc) result = _convert_mcp_messages_to_openai([msg], system_prompt="Be helpful") assert len(result) == 2 assert result[0]["role"] == "system" assert result[0]["content"] == "Be helpful" + def test_should_convert_image_content(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_messages_to_openai, ) + ic = _make_image_content("base64imgdata", "image/jpeg") msg = _make_sampling_message("user", ic) result = _convert_mcp_messages_to_openai([msg]) @@ -114,10 +134,12 @@ class TestConvertMCPMessagesToOpenAI: assert isinstance(content, list) assert content[0]["type"] == "image_url" assert "base64imgdata" in content[0]["image_url"]["url"] + def test_should_convert_list_of_mixed_content(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_messages_to_openai, ) + tc = _make_text_content("Describe this image") ic = _make_image_content("imgdata") msg = _make_sampling_message("user", [tc, ic]) @@ -126,69 +148,88 @@ class TestConvertMCPMessagesToOpenAI: content = result[0]["content"] assert isinstance(content, list) assert len(content) == 2 + def test_should_convert_multiple_messages(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_messages_to_openai, ) + 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" assert result[1]["role"] == "assistant" + + # ───────────────────────────────────────────────────────────── # Tests: Model resolution # ───────────────────────────────────────────────────────────── class TestResolveModel: """Tests for _resolve_model_from_preferences.""" + def test_should_use_default_model_when_no_preferences(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _resolve_model_from_preferences, ) + 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) def test_should_fallback_to_gpt4o_mini(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _resolve_model_from_preferences, ) + with patch("litellm.model_list", []): result = _resolve_model_from_preferences(None) assert result == "gpt-4o-mini" + @patch("litellm.model_list", ["gpt-4o", "claude-3.5-sonnet", "gemini-pro"]) + @patch("litellm.proxy.proxy_server.llm_router", None) def test_should_match_hint_by_substring(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _resolve_model_from_preferences, ) + hint = _make_hint("claude") prefs = _make_model_preferences(hints=[hint]) result = _resolve_model_from_preferences(prefs) assert "claude" in result.lower() + @patch("litellm.model_list", ["gpt-4o"]) + @patch("litellm.proxy.proxy_server.llm_router", None) def test_should_use_default_when_no_hint_matches(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _resolve_model_from_preferences, ) + hint = _make_hint("nonexistent-model") prefs = _make_model_preferences(hints=[hint]) result = _resolve_model_from_preferences(prefs, default_model="gpt-4o") assert result == "gpt-4o" + + # ───────────────────────────────────────────────────────────── # Tests: Tool conversion # ───────────────────────────────────────────────────────────── class TestConvertMCPToolsToOpenAI: """Tests for _convert_mcp_tools_to_openai.""" + def test_should_return_none_for_no_tools(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_tools_to_openai, ) + assert _convert_mcp_tools_to_openai(None) is None assert _convert_mcp_tools_to_openai([]) is None + def test_should_convert_mcp_tool_to_openai_format(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_tools_to_openai, ) + tool = MagicMock() tool.name = "get_weather" tool.description = "Get weather for a city" @@ -202,72 +243,86 @@ class TestConvertMCPToolsToOpenAI: assert result[0]["type"] == "function" assert result[0]["function"]["name"] == "get_weather" assert result[0]["function"]["description"] == "Get weather for a city" + + # ───────────────────────────────────────────────────────────── # Tests: Tool choice conversion # ───────────────────────────────────────────────────────────── class TestConvertMCPToolChoiceToOpenAI: """Tests for _convert_mcp_tool_choice_to_openai.""" + def test_should_return_none_for_no_tool_choice(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_tool_choice_to_openai, ) + assert _convert_mcp_tool_choice_to_openai(None) is None + def test_should_convert_auto_mode(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_tool_choice_to_openai, ) + tc = MagicMock() tc.mode = "auto" assert _convert_mcp_tool_choice_to_openai(tc) == "auto" + def test_should_convert_required_mode(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_tool_choice_to_openai, ) + tc = MagicMock() tc.mode = "required" assert _convert_mcp_tool_choice_to_openai(tc) == "required" + def test_should_convert_none_mode(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_mcp_tool_choice_to_openai, ) + tc = MagicMock() tc.mode = "none" assert _convert_mcp_tool_choice_to_openai(tc) == "none" + + # ───────────────────────────────────────────────────────────── # Tests: Response conversion # ───────────────────────────────────────────────────────────── class TestConvertOpenAIResponseToMCPResult: """Tests for _convert_openai_response_to_mcp_result.""" + def test_should_convert_text_response(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_openai_response_to_mcp_result, ) + response = _make_completion_response(content="Hello!", model="gpt-4o") result = _convert_openai_response_to_mcp_result(response, "gpt-4o") assert result.role == "assistant" assert result.model == "gpt-4o" assert result.stopReason == "endTurn" assert result.content.text == "Hello!" + def test_should_set_max_tokens_stop_reason(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _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" + def test_should_convert_tool_calls_response(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_openai_response_to_mcp_result, ) + tc = MagicMock() 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) @@ -275,16 +330,20 @@ class TestConvertOpenAIResponseToMCPResult: tool_use = result.content[0] assert tool_use.type == "tool_use" assert tool_use.name == "get_weather" + + # ───────────────────────────────────────────────────────────── # Tests: Full handler # ───────────────────────────────────────────────────────────── class TestHandleSamplingCreateMessage: """Tests for the main handle_sampling_create_message function.""" + @pytest.mark.asyncio async def test_should_call_litellm_acompletion(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( handle_sampling_create_message, ) + mock_response = _make_completion_response(content="Test response") params = _make_params( messages=[_make_sampling_message("user", _make_text_content("Hello"))], @@ -300,11 +359,13 @@ class TestHandleSamplingCreateMessage: mock_completion.assert_called_once() assert result.role == "assistant" assert result.content.text == "Test response" + @pytest.mark.asyncio async def test_should_include_temperature(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( handle_sampling_create_message, ) + mock_response = _make_completion_response() params = _make_params( messages=[_make_sampling_message("user", _make_text_content("Hi"))], @@ -319,11 +380,13 @@ class TestHandleSamplingCreateMessage: ) call_kwargs = mock_completion.call_args[1] assert call_kwargs["temperature"] == 0.7 + @pytest.mark.asyncio async def test_should_include_stop_sequences(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( handle_sampling_create_message, ) + mock_response = _make_completion_response() params = _make_params( messages=[_make_sampling_message("user", _make_text_content("Hi"))], @@ -338,11 +401,13 @@ class TestHandleSamplingCreateMessage: ) call_kwargs = mock_completion.call_args[1] assert call_kwargs["stop"] == ["STOP", "END"] + @pytest.mark.asyncio async def test_should_include_tools_and_tool_choice(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( handle_sampling_create_message, ) + mock_response = _make_completion_response() tool = MagicMock() tool.name = "search" @@ -365,11 +430,13 @@ class TestHandleSamplingCreateMessage: call_kwargs = mock_completion.call_args[1] assert "tools" in call_kwargs assert call_kwargs["tool_choice"] == "auto" + @pytest.mark.asyncio async def test_should_return_error_on_exception(self): from litellm.proxy._experimental.mcp_server.sampling_handler import ( handle_sampling_create_message, ) + params = _make_params( messages=[_make_sampling_message("user", _make_text_content("Hi"))], ) @@ -382,4 +449,4 @@ class TestHandleSamplingCreateMessage: ) assert hasattr(result, "code") assert result.code == -1 - assert "API error" in result.message \ No newline at end of file + assert "API error" in result.message 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 447c28078ec..43178c0da13 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,18 +35,14 @@ 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 @@ -307,9 +303,7 @@ 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( @@ -337,9 +331,7 @@ 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 @@ -410,9 +402,7 @@ 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] @@ -443,13 +433,11 @@ 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( - server, mcp_auth_header, extra_headers, stdio_env + server, mcp_auth_header, extra_headers, stdio_env, **kwargs ): # pragma: no cover - helper nonlocal captured_extra_headers captured_extra_headers = extra_headers @@ -559,9 +547,7 @@ 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( @@ -692,9 +678,7 @@ 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 @@ -741,9 +725,7 @@ 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) @@ -861,9 +843,7 @@ 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] @@ -930,9 +910,7 @@ 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 @@ -1054,9 +1032,7 @@ 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 @@ -1398,9 +1374,7 @@ 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={}) @@ -1444,13 +1418,8 @@ 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): @@ -1474,9 +1443,7 @@ 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={}) @@ -1520,13 +1487,8 @@ 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): @@ -1550,9 +1512,7 @@ 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={}) @@ -1589,9 +1549,7 @@ 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={}) @@ -1617,10 +1575,7 @@ 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.""" @@ -1651,9 +1606,7 @@ 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" @@ -1688,13 +1641,9 @@ 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.""" @@ -1720,14 +1669,11 @@ class TestMCPServerManager: # Unprefixed resolution resolved_server_unpref = manager._get_mcp_server_from_tool_name("create_zap") - print(resolved_server_unpref) assert resolved_server_unpref is not None 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 @@ -1772,9 +1718,7 @@ 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 @@ -1824,9 +1768,7 @@ 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 @@ -1871,9 +1813,7 @@ 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 @@ -1943,9 +1883,7 @@ 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) @@ -1988,9 +1926,7 @@ 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) @@ -2135,9 +2071,7 @@ 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={}) @@ -2173,13 +2107,8 @@ 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): @@ -2204,9 +2133,7 @@ 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() @@ -2230,9 +2157,7 @@ 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) @@ -2297,9 +2222,7 @@ 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 @@ -2544,10 +2467,7 @@ 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.""" @@ -2555,9 +2475,7 @@ 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.""" @@ -2635,9 +2553,7 @@ 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"] @@ -2661,9 +2577,7 @@ 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"])) == [ @@ -2673,9 +2587,7 @@ 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"] @@ -2687,23 +2599,15 @@ 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"])) == [ @@ -2714,9 +2618,7 @@ 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"] @@ -2726,14 +2628,10 @@ 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"] @@ -2762,18 +2660,14 @@ 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"]} @@ -2792,9 +2686,7 @@ 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"]}) @@ -2807,14 +2699,64 @@ 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"]}) + assert sorted(result["uuid-a"]) == ["read_file", "write_file"] + + @pytest.mark.asyncio + async def test_create_sampling_callback(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _create_sampling_callback, ) - result = manager.expand_tool_permissions( - {"uuid-a": ["read_file"], "alias-a": ["write_file"]} + callback = _create_sampling_callback(user_api_key_auth="test_auth") + + with patch( + "litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", + new_callable=AsyncMock, + ) as mock_handle: + mock_handle.return_value = "mocked_result" + + result = await callback("mock_context", "mock_params") + + mock_handle.assert_called_once() + called_kwargs = mock_handle.call_args.kwargs + assert called_kwargs["context"] == "mock_context" + assert called_kwargs["params"] == "mock_params" + assert called_kwargs["user_api_key_auth"] == "test_auth" + assert result == "mocked_result" + + @pytest.mark.asyncio + async def test_create_elicitation_callback(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _create_elicitation_callback, ) - assert sorted(result["uuid-a"]) == ["read_file", "write_file"] + + callback = _create_elicitation_callback() + + with ( + patch( + "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, + ): + mock_session = MagicMock() + mock_session.capabilities = "test_capabilities" + mock_get_session.return_value = mock_session + + mock_handle.return_value = "mocked_elicitation" + + result = await callback("mock_context", "mock_params") + + mock_handle.assert_called_once() + called_kwargs = mock_handle.call_args.kwargs + assert called_kwargs["context"] == "mock_context" + assert called_kwargs["params"] == "mock_params" + assert called_kwargs["downstream_session"] == mock_session + assert called_kwargs["downstream_capabilities"] == "test_capabilities" + assert result == "mocked_elicitation" if __name__ == "__main__":