From d2aee7e6592d2328b83a2d32539459c7dc3d92fb Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 1 Jul 2026 11:36:15 +0530 Subject: [PATCH] style: fix ruff format and prettier formatting Co-authored-by: Cursor --- .../mcp_server/mcp_server_manager.py | 818 +++++------------- .../mcp_tools/create_mcp_server.tsx | 22 +- 2 files changed, 211 insertions(+), 629 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b7d2aaf78bf..3d5b7a32580 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -166,9 +166,7 @@ def invalidate_user_env_vars_cache(user_id: str, server_id: str) -> None: _user_env_vars_cache.pop((user_id, server_id), None) -def _write_user_env_vars_cache( - user_id: str, server_id: str, values: Dict[str, str] -) -> None: +def _write_user_env_vars_cache(user_id: str, server_id: str, values: Dict[str, str]) -> None: cache_key = (user_id, server_id) # Re-insert at the tail so eviction drops the oldest-written entry, not a # freshly refreshed one, and only sheds a single entry instead of wiping the @@ -210,10 +208,7 @@ def _should_strip_caller_authorization( """ if mcp_server.has_client_credentials: return True - if ( - mcp_server.auth_type == MCPAuth.oauth2 - and to_server_spec(mcp_server) is not None - ): + if mcp_server.auth_type == MCPAuth.oauth2 and to_server_spec(mcp_server) is not None: # Migrated per-user OAuth (authorization_code): the v2 resolver injects the # stored token, so a caller-forwarded Authorization must not be forwarded # upstream — it would override another user's stored credential. Delegate and @@ -222,12 +217,8 @@ def _should_strip_caller_authorization( if not mcp_server.is_oauth_passthrough: return False - normalized_raw_headers = { - str(k).lower(): v for k, v in (raw_headers or {}).items() if isinstance(k, str) - } - has_explicit_litellm_admission_header = ( - normalized_raw_headers.get("x-litellm-api-key") is not None - ) + normalized_raw_headers = {str(k).lower(): v for k, v in (raw_headers or {}).items() if isinstance(k, str)} + has_explicit_litellm_admission_header = normalized_raw_headers.get("x-litellm-api-key") is not None admission_consumed_authorization_as_litellm_key = ( user_api_key_auth is not None and bool(getattr(user_api_key_auth, "api_key", None)) @@ -292,10 +283,7 @@ def _extract_upstream_auth_failure( if current.__cause__ is not None: stack.append(current.__cause__) - if ( - current.__context__ is not None - and current.__context__ is not current.__cause__ - ): + if current.__context__ is not None and current.__context__ is not current.__cause__: stack.append(current.__context__) return None @@ -314,9 +302,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, @@ -329,9 +315,7 @@ def _warn_on_server_name_fields( _warn("server_name", server_name) -def _warn_internal_delegate_pkce_if_applicable( - server: MCPServer, *, source: str -) -> None: +def _warn_internal_delegate_pkce_if_applicable(server: MCPServer, *, source: str) -> None: """Surface internal + upstream PKCE delegate in logs for operators.""" if server.auth_type != MCPAuth.oauth2: return @@ -393,10 +377,7 @@ def _deserialize_json_list(data: Any) -> Optional[List[Dict[str, Any]]]: data = parsed if not isinstance(data, list): return None - return [ - item.model_dump(mode="json") if hasattr(item, "model_dump") else item - for item in data - ] + return [item.model_dump(mode="json") if hasattr(item, "model_dump") else item for item in data] def _normalize_mcp_server_cost_info(mcp_info: MCPInfo) -> None: @@ -464,9 +445,7 @@ def _create_sampling_callback(user_api_key_auth: Optional[Any] = None): ) auth_context = get_active_auth_context() - resolved_auth = user_api_key_auth or ( - auth_context.user_api_key_auth if auth_context else None - ) + resolved_auth = user_api_key_auth or (auth_context.user_api_key_auth if auth_context else None) # Forward original HTTP headers and client IP so that # header-dependent guardrails, tag-based routing, trace # correlation, and forward_llm_provider_auth_headers work @@ -505,11 +484,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, @@ -541,9 +516,7 @@ class MCPServerManager: unless authorization_url is present (interactive OAuth). """ if oauth2_flow in ("client_credentials", "authorization_code"): - return cast( - Literal["client_credentials", "authorization_code"], oauth2_flow - ) + return cast(Literal["client_credentials", "authorization_code"], oauth2_flow) if oauth2_flow: # Ignore unknown/untyped values and continue legacy inference. return None @@ -589,18 +562,12 @@ class MCPServerManager: # not return instructions, and to apply a short cooldown after failures. self._upstream_initialize_instructions_probed_at: Dict[str, float] = {} - 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() - async def _ensure_upstream_initialize_instructions_cached( - self, server: MCPServer - ) -> None: + async def _ensure_upstream_initialize_instructions_cached(self, server: MCPServer) -> None: """ Open one upstream session and cache InitializeResult.instructions if missing. @@ -635,20 +602,13 @@ class MCPServerManager: ): return - last_probed_at = self._upstream_initialize_instructions_probed_at.get( - server.server_id - ) - if ( - last_probed_at is not None - and (time.monotonic() - last_probed_at) < MCP_HEALTH_CHECK_TIMEOUT - ): + last_probed_at = self._upstream_initialize_instructions_probed_at.get(server.server_id) + if last_probed_at is not None and (time.monotonic() - last_probed_at) < MCP_HEALTH_CHECK_TIMEOUT: return # Record the attempt up-front so that a failure / empty response does not # cause every subsequent initialize request to re-open the upstream session. - self._upstream_initialize_instructions_probed_at[server.server_id] = ( - time.monotonic() - ) + self._upstream_initialize_instructions_probed_at[server.server_id] = time.monotonic() try: resolved_static_headers = await self._resolve_static_headers_with_env_vars( @@ -656,9 +616,7 @@ class MCPServerManager: user_api_key_auth=None, raise_on_missing=False, ) - extra_headers: Optional[Dict[str, str]] = ( - dict(resolved_static_headers) if resolved_static_headers else None - ) + extra_headers: Optional[Dict[str, str]] = dict(resolved_static_headers) if resolved_static_headers else None client = await self._create_mcp_client( server=server, mcp_auth_header=None, @@ -669,9 +627,7 @@ class MCPServerManager: async def _noop(_session): return "ok" - 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) except Exception as e: verbose_logger.debug( @@ -724,15 +680,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 @@ -767,9 +718,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 ) @@ -805,15 +754,11 @@ class MCPServerManager: authorization_url=resolved_authorization_url, token_url=resolved_token_url, registration_url=resolved_registration_url, - token_endpoint_auth_method=server_config.get( - "token_endpoint_auth_method", None - ), + token_endpoint_auth_method=server_config.get("token_endpoint_auth_method", None), # 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), @@ -823,12 +768,8 @@ class MCPServerManager: static_headers=server_config.get("static_headers", None), env_vars=server_config.get("env_vars", 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) - ), - delegate_auth_to_upstream=bool( - server_config.get("delegate_auth_to_upstream", False) - ), + available_on_public_internet=bool(server_config.get("available_on_public_internet", True)), + delegate_auth_to_upstream=bool(server_config.get("delegate_auth_to_upstream", False)), oauth_passthrough=bool(server_config.get("oauth_passthrough", False)), # AWS SigV4 fields aws_access_key_id=server_config.get("aws_access_key_id", None), @@ -840,9 +781,7 @@ class MCPServerManager: aws_session_name=server_config.get("aws_session_name", None), instructions=server_config.get("instructions", None), # Token Exchange (OBO) fields - token_exchange_endpoint=server_config.get( - "token_exchange_endpoint", None - ), + token_exchange_endpoint=server_config.get("token_exchange_endpoint", None), audience=server_config.get("audience", None), subject_token_type=server_config.get( "subject_token_type", @@ -859,24 +798,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. @@ -912,9 +845,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) @@ -968,20 +899,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( @@ -993,9 +918,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 @@ -1008,26 +931,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 _cleanup_server_tool_routing_artifacts(self, server: MCPServer) -> None: @@ -1058,9 +971,7 @@ class MCPServerManager: owned_normalized = {normalize_server_name(x) for x in owned_raw} stale_mapping_keys: List[str] = [] - for tool_name, mapped_server in list( - self.tool_name_to_mcp_server_name_mapping.items() - ): + for tool_name, mapped_server in list(self.tool_name_to_mcp_server_name_mapping.items()): if mapped_server in owned_raw: stale_mapping_keys.append(tool_name) elif normalize_server_name(str(mapped_server)) in owned_normalized: @@ -1077,14 +988,10 @@ class MCPServerManager: if evicted is None and mcp_server.server_name: evicted = self.registry.pop(mcp_server.server_name, None) if evicted is not None: - verbose_logger.debug( - "Removed MCP Server: %s", mcp_server.server_id or mcp_server.server_name - ) + verbose_logger.debug("Removed MCP Server: %s", mcp_server.server_id or mcp_server.server_name) self._cleanup_server_tool_routing_artifacts(evicted) 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") def _resolve_env_vars_list( self, @@ -1110,20 +1017,14 @@ 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) - ) + static_headers_dict = _deserialize_json_dict(getattr(mcp_server, "static_headers", None)) env_vars_list = self._resolve_env_vars_list( mcp_server, env_vars_are_encrypted=( - credentials_are_encrypted - if env_vars_are_encrypted is None - else env_vars_are_encrypted + credentials_are_encrypted if env_vars_are_encrypted is None else env_vars_are_encrypted ), ) - credentials_dict = _deserialize_json_dict( - getattr(mcp_server, "credentials", None) - ) + credentials_dict = _deserialize_json_dict(getattr(mcp_server, "credentials", None)) encrypted_auth_value: Optional[str] = None encrypted_client_id: Optional[str] = None @@ -1170,9 +1071,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: @@ -1180,9 +1079,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: @@ -1193,20 +1090,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, @@ -1223,30 +1114,22 @@ class MCPServerManager: static_headers=static_headers_dict, env_vars=env_vars_list, 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=self._resolve_oauth2_flow( auth_type=auth_type, oauth2_flow=getattr(mcp_server, "oauth2_flow", None), - token_url=mcp_server.token_url - or getattr(mcp_oauth_metadata, "token_url", None), + token_url=mcp_server.token_url or getattr(mcp_oauth_metadata, "token_url", None), authorization_url=mcp_server.authorization_url or getattr(mcp_oauth_metadata, "authorization_url", None), 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), ), 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), token_endpoint_auth_method=( - credentials_dict.get("token_endpoint_auth_method") - if credentials_dict - else None + credentials_dict.get("token_endpoint_auth_method") if credentials_dict else None ), command=getattr(mcp_server, "command", None), args=getattr(mcp_server, "args", None) or [], @@ -1255,21 +1138,13 @@ 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) - ), - delegate_auth_to_upstream=bool( - getattr(mcp_server, "delegate_auth_to_upstream", False) - ), + available_on_public_internet=bool(getattr(mcp_server, "available_on_public_internet", True)), + delegate_auth_to_upstream=bool(getattr(mcp_server, "delegate_auth_to_upstream", False)), oauth_passthrough=bool(getattr(mcp_server, "oauth_passthrough", False)), 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), @@ -1284,29 +1159,19 @@ class MCPServerManager: aws_session_name=aws_creds.get("aws_session_name"), instructions=mcp_server.instructions, # Token Exchange (OBO) fields — read from credentials JSON blob - token_exchange_endpoint=( - credentials_dict.get("token_exchange_endpoint") - if credentials_dict - else None - ), + token_exchange_endpoint=(credentials_dict.get("token_exchange_endpoint") if credentials_dict else None), audience=(credentials_dict.get("audience") if credentials_dict else None), - subject_token_type=( - credentials_dict.get("subject_token_type") if credentials_dict else None - ) + subject_token_type=(credentials_dict.get("subject_token_type") if credentials_dict else None) or "urn:ietf:params:oauth:token-type:access_token", timeout=getattr(mcp_server, "timeout", None), ) _warn_internal_delegate_pkce_if_applicable(new_server, source="database") return new_server - async def _maybe_register_openapi_tools( - self, server: MCPServer, *, initialize_mapping: bool = True - ): + async def _maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True): """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, @@ -1330,9 +1195,7 @@ class MCPServerManager: # `credentials` field is the only one still encrypted here). # Re-decrypting plaintext would zero the values, so build with # env_vars_are_encrypted=False. - new_server = await self.build_mcp_server_from_table( - mcp_server, env_vars_are_encrypted=False - ) + new_server = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False) self._assign_unique_short_prefix(new_server) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) @@ -1357,9 +1220,7 @@ class MCPServerManager: if mcp_server.server_id in self.registry: # See add_server: db.py helpers already decrypted env var # values, so don't decrypt them a second time here. - new_server = await self.build_mcp_server_from_table( - mcp_server, env_vars_are_encrypted=False - ) + new_server = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False) # Carry the previously-resolved short prefix across so the # tool names stay stable for clients holding cached lists. existing_prefix = self.registry[mcp_server.server_id].short_prefix @@ -1383,15 +1244,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. @@ -1407,12 +1262,9 @@ class MCPServerManager: try: # The key explicitly opted out of every MCP server. Return zero before # layering on allow_all_keys servers so the opt-out is absolute. - key_object_permission = ( - user_api_key_auth.object_permission if user_api_key_auth else None - ) + key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None if key_object_permission is not None and ( - SpecialMCPServerNames.no_mcp_servers.value - in (key_object_permission.mcp_servers or []) + SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or []) ): return [] @@ -1427,23 +1279,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 @@ -1486,11 +1328,7 @@ class MCPServerManager: # BYOM: approved submissions stay visible to the submitter even when # allow_all_keys=false and no access groups were configured at approval. # ponytail: 60-second TTL cache — upgrade to user-level flag if BYOM adoption grows - submitter_user_id = ( - getattr(user_api_key_auth, "user_id", None) - if user_api_key_auth - else None - ) + submitter_user_id = getattr(user_api_key_auth, "user_id", None) if user_api_key_auth else None if submitter_user_id: from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 get_active_submitted_mcp_server_ids_for_user, @@ -1498,14 +1336,10 @@ class MCPServerManager: from litellm.proxy.proxy_server import prisma_client, user_api_key_cache byom_cache_key = f"byom_submitted_servers:{submitter_user_id}" - submitted_server_ids = await user_api_key_cache.async_get_cache( - key=byom_cache_key - ) + submitted_server_ids = await user_api_key_cache.async_get_cache(key=byom_cache_key) if submitted_server_ids is None: submitted_server_ids = ( - await get_active_submitted_mcp_server_ids_for_user( - prisma_client, submitter_user_id - ) + await get_active_submitted_mcp_server_ids_for_user(prisma_client, submitter_user_id) if prisma_client is not None else [] ) @@ -1515,15 +1349,11 @@ class MCPServerManager: ttl=60, ) combined_servers.update( - server_id - for server_id in submitted_server_ids - if self.get_mcp_server_by_id(server_id) is not None + server_id for server_id in submitted_server_ids if self.get_mcp_server_by_id(server_id) is not None ) 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: # noqa: BLE001 verbose_logger.exception( @@ -1603,15 +1433,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, @@ -1648,18 +1475,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=get_management_object_ttl(user_api_key_cache), ) 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. @@ -1701,9 +1522,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( @@ -1772,9 +1591,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 ######################################################### @@ -1896,9 +1713,7 @@ class MCPServerManager: # the user hasn't filled it in -- only vars without a global fallback do. referenced = collect_env_var_references(strings=(static_headers or {}).values()) referenced_user_vars = referenced & user_var_names - required_user_vars = { - name for name in referenced_user_vars if name not in global_values - } + required_user_vars = {name for name in referenced_user_vars if name not in global_values} user_values: Dict[str, str] = {} if required_user_vars: @@ -1918,21 +1733,15 @@ class MCPServerManager: ) if raise_on_missing: - missing = sorted( - name for name in required_user_vars if not user_values.get(name) - ) + missing = sorted(name for name in required_user_vars if not user_values.get(name)) if missing: # A cached negative must never produce a 412: cache # invalidation is process-local, so a user who just stored # values on another worker would otherwise be told their # credentials are missing until the entry expires. Confirm # against the DB before raising. - user_values = await self._load_user_env_vars( - server, user_api_key_auth, force_refresh=True - ) - missing = sorted( - name for name in required_user_vars if not user_values.get(name) - ) + user_values = await self._load_user_env_vars(server, user_api_key_auth, force_refresh=True) + missing = sorted(name for name in required_user_vars if not user_values.get(name)) if missing: raise MCPMissingUserEnvVarsError( server_id=server.server_id, @@ -1944,9 +1753,7 @@ class MCPServerManager: # Only honor stored user values for currently user-scoped vars, and let # admin globals win, so a stale row from when a var was user-scoped can # never override the global value the admin set after switching it. - scoped_user_values = { - name: value for name, value in user_values.items() if name in user_var_names - } + scoped_user_values = {name: value for name, value in user_values.items() if name in user_var_names} merged_vars: Dict[str, str] = {**scoped_user_values, **global_values} if not static_headers: return static_headers @@ -2041,34 +1848,20 @@ class MCPServerManager: # caller must not be able to substitute another user's stored credential, so we keep the v2 # spec and ignore the override there; the REST tools preview supplies its not-yet-persisted # token through the resolver (cred_provider), never this path. - if ( - spec is not None - and mcp_auth_header - and not isinstance(spec.config, AuthorizationCodeConfig) - ): + if spec is not None and mcp_auth_header and not isinstance(spec.config, AuthorizationCodeConfig): spec = None auth_value = ( - await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token) - if spec is None - else None + await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token) if spec is None else None ) # Create sampling and elicitation callbacks for this client - sampling_cb = ( - _create_sampling_callback(user_api_key_auth=user_api_key_auth) - if server.allow_sampling - else None - ) - elicitation_cb = ( - _create_elicitation_callback() if server.allow_elicitation else None - ) + sampling_cb = _create_sampling_callback(user_api_key_auth=user_api_key_auth) if server.allow_sampling else None + elicitation_cb = _create_elicitation_callback() if server.allow_elicitation else None # Handle stdio transport if transport == MCPTransport.stdio: resolved_env = ( - stdio_env - if stdio_env is not None - else (dict(server.env) if server.env is not None else None) + stdio_env if stdio_env is not None else (dict(server.env) if server.env is not None else None) ) # Ensure npm-based STDIO MCP servers have a writable cache dir. @@ -2110,9 +1903,7 @@ class MCPServerManager: transport_type=transport, auth_type=server.auth_type, auth_value=auth_value, - timeout=( - server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT - ), + timeout=(server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT), stdio_config=stdio_config, extra_headers=extra_headers, sampling_callback=sampling_cb, @@ -2123,9 +1914,7 @@ class MCPServerManager: server_url = server.url or "" if spec is not None: - match await provider.resolve_credentials( - to_subject(user_api_key_auth, subject_token), spec - ): + match await provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec): case Ok(auth): resolved_auth = auth # Do not override an Authorization already supplied via extra_headers @@ -2136,10 +1925,7 @@ class MCPServerManager: if ( header_name and extra_headers - and any( - key.lower() == header_name.lower() - for key in extra_headers - ) + and any(key.lower() == header_name.lower() for key in extra_headers) ): resolved_auth = None case Error(err): @@ -2152,11 +1938,7 @@ class MCPServerManager: server_url=server_url, transport_type=transport, auth_type=server.auth_type, - timeout=( - server.timeout - if server.timeout is not None - else MCP_CLIENT_TIMEOUT - ), + timeout=(server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT), extra_headers=extra_headers, resolved_auth=resolved_auth, sampling_callback=sampling_cb, @@ -2181,9 +1963,7 @@ class MCPServerManager: transport_type=transport, auth_type=server.auth_type, auth_value=auth_value, - timeout=( - server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT - ), + timeout=(server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT), extra_headers=extra_headers, aws_auth=aws_auth, sampling_callback=sampling_cb, @@ -2250,12 +2030,10 @@ class MCPServerManager: static_headers = server.static_headers or {} has_static_authorization = any( - isinstance(k, str) and k.lower() == "authorization" - for k in static_headers.keys() + isinstance(k, str) and k.lower() == "authorization" for k in static_headers.keys() ) has_extra_authorization = bool(extra_headers) and any( - isinstance(k, str) and k.lower() == "authorization" - for k in (extra_headers or {}).keys() + isinstance(k, str) and k.lower() == "authorization" for k in (extra_headers or {}).keys() ) if ( @@ -2285,12 +2063,8 @@ class MCPServerManager: if server.spec_path: # OpenAPI tools were stored in the registry under the prefix # active at registration time — fetch by that same prefix. - _tools = global_mcp_tool_registry.list_tools( - tool_prefix=get_server_prefix(server) - ) - tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type( - _tools - ) + _tools = global_mcp_tool_registry.list_tools(tool_prefix=get_server_prefix(server)) + 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 @@ -2300,9 +2074,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 ) @@ -2310,14 +2082,10 @@ class MCPServerManager: ] return tools else: - tools = await self._fetch_tools_with_timeout( - client, server.name, server=server - ) + tools = await self._fetch_tools_with_timeout(client, server.name, server=server) 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 @@ -2327,9 +2095,7 @@ class MCPServerManager: # aggregator catches this explicitly to keep absorbing. raise 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( @@ -2373,16 +2139,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( @@ -2417,16 +2179,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( @@ -2468,9 +2226,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( @@ -2591,15 +2347,8 @@ class MCPServerManager: authorization_servers, resource_scopes, ) = await self._attempt_well_known_discovery(server_url) - metadata = await self._fetch_authorization_server_metadata( - authorization_servers, server_url - ) - if ( - metadata is None - and not resource_scopes - and authorization_servers - and response.status_code == 200 - ): + metadata = await self._fetch_authorization_server_metadata(authorization_servers, server_url) + if metadata is None and not resource_scopes and authorization_servers and response.status_code == 200: verbose_logger.warning( "MCP OAuth discovery for %s received 200 OK without RFC 9728 challenge and no discoverable authorization metadata.", server_url, @@ -2618,13 +2367,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 = [] resource_scopes = None @@ -2632,9 +2379,7 @@ class MCPServerManager: ( authorization_servers, resource_scopes, - ) = await self._fetch_oauth_metadata_from_resource( - resource_metadata_url, server_url - ) + ) = await self._fetch_oauth_metadata_from_resource(resource_metadata_url, server_url) else: ( authorization_servers, @@ -2646,16 +2391,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, server_url - ) + metadata = await self._fetch_authorization_server_metadata(authorization_servers, server_url) preferred_scopes = scopes or resource_scopes if metadata is None and preferred_scopes: @@ -2665,14 +2406,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 @@ -2681,8 +2418,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") @@ -2700,9 +2436,7 @@ class MCPServerManager: return [], None try: - response = await self._fetch_oauth_discovery_url( - resource_metadata_url, server_url - ) + response = await self._fetch_oauth_discovery_url(resource_metadata_url, server_url) response.raise_for_status() data = response.json() except SSRFError as exc: @@ -2724,23 +2458,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: @@ -2772,9 +2498,7 @@ class MCPServerManager: self, authorization_servers: List[str], server_url: str ) -> Optional[MCPOAuthMetadata]: for issuer in authorization_servers: - metadata = await self._fetch_single_authorization_server_metadata( - issuer, server_url - ) + metadata = await self._fetch_single_authorization_server_metadata(issuer, server_url) if metadata is not None: return metadata return None @@ -2795,13 +2519,9 @@ class MCPServerManager: candidate_urls: List[str] = [] if path: - candidate_urls.append( - f"{base}/.well-known/oauth-authorization-server/{path}" - ) + candidate_urls.append(f"{base}/.well-known/oauth-authorization-server/{path}") candidate_urls.append(f"{base}/.well-known/openid-configuration/{path}") - candidate_urls.append( - f"{issuer_url.rstrip('/')}/.well-known/openid-configuration" - ) + candidate_urls.append(f"{issuer_url.rstrip('/')}/.well-known/openid-configuration") candidate_urls.append(f"{base}/.well-known/oauth-authorization-server") candidate_urls.append(f"{base}/.well-known/openid-configuration") candidate_urls.append(issuer_url.rstrip("/")) @@ -2852,14 +2572,8 @@ class MCPServerManager: def _build_azure_authorization_server_metadata( parsed_issuer_url: Any, ) -> Optional[MCPOAuthMetadata]: - path_parts = [ - part for part in (parsed_issuer_url.path or "").split("/") if part - ] - if ( - parsed_issuer_url.netloc not in _AZURE_ENTRA_HOSTS - or len(path_parts) != 2 - or path_parts[1] != "v2.0" - ): + path_parts = [part for part in (parsed_issuer_url.path or "").split("/") if part] + if parsed_issuer_url.netloc not in _AZURE_ENTRA_HOSTS or len(path_parts) != 2 or path_parts[1] != "v2.0": return None tenant = path_parts[0] @@ -2969,32 +2683,24 @@ class MCPServerManager: ) try: with anyio.fail_after(MCP_TOOL_LISTING_TIMEOUT): - tools = await client.list_tools( - raise_on_error=should_surface_upstream_auth - ) + tools = await client.list_tools(raise_on_error=should_surface_upstream_auth) verbose_logger.debug(f"Tools from {server_name}: {tools}") return tools except TimeoutError: 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: if should_surface_upstream_auth: auth_info = _extract_upstream_auth_failure(e) if auth_info is not None: status_code, www_authenticate = auth_info - verbose_logger.info( - f"Upstream auth failure from MCP server {server_name}: HTTP {status_code}" - ) + verbose_logger.info(f"Upstream auth failure from MCP server {server_name}: HTTP {status_code}") raise MCPUpstreamAuthError( status_code=status_code, www_authenticate=www_authenticate, @@ -3065,9 +2771,7 @@ class MCPServerManager: "attempts; the 3-character prefix space is too crowded." ) - 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. @@ -3101,9 +2805,7 @@ class MCPServerManager: qualified = add_server_prefix_to_name(original_name, known_prefix) self.tool_name_to_mcp_server_name_mapping[qualified] = 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( @@ -3130,9 +2832,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( @@ -3144,17 +2844,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( @@ -3170,9 +2864,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) @@ -3193,20 +2885,14 @@ class MCPServerManager: if server_applies_tool_allowlist(server): if not server.allowed_tools: return False - 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. @@ -3233,18 +2919,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( @@ -3407,42 +3089,22 @@ class MCPServerManager: "name": name, "arguments": arguments, "server_name": server_name, - "mcp_rate_limit_server_name": server.alias - or server.server_name - or server.name, + "mcp_rate_limit_server_name": server.alias or server.server_name or 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: @@ -3454,11 +3116,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"): @@ -3503,9 +3161,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( @@ -3601,9 +3257,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)} strip_caller_authorization = _should_strip_caller_authorization( mcp_server=mcp_server, raw_headers=raw_headers, @@ -3624,9 +3278,7 @@ class MCPServerManager: # MCPMissingUserEnvVarsError when the calling user has not filled in # a required per-user variable — the REST layer converts that into # a friendly 412 with a setup URL. - resolved_static_headers = await self._resolve_static_headers_with_env_vars( - mcp_server, user_api_key_auth - ) + resolved_static_headers = await self._resolve_static_headers_with_env_vars(mcp_server, user_api_key_auth) if resolved_static_headers: if extra_headers is None: extra_headers = {} @@ -3678,21 +3330,13 @@ class MCPServerManager: ) async def _call_tool_via_client(client, params): - return await client.call_tool( - params, host_progress_callback=host_progress_callback - ) + return await client.call_tool(params, host_progress_callback=host_progress_callback) - tasks.append( - asyncio.create_task(_call_tool_via_client(client, call_tool_params)) - ) + tasks.append(asyncio.create_task(_call_tool_via_client(client, call_tool_params))) - _timeout = ( - mcp_server.timeout if mcp_server.timeout is not None else MCP_CLIENT_TIMEOUT - ) + _timeout = mcp_server.timeout if mcp_server.timeout is not None else MCP_CLIENT_TIMEOUT try: - mcp_responses = await asyncio.wait_for( - asyncio.gather(*tasks), timeout=_timeout - ) + mcp_responses = await asyncio.wait_for(asyncio.gather(*tasks), timeout=_timeout) except asyncio.TimeoutError: raise HTTPException( status_code=504, @@ -3706,9 +3350,7 @@ class MCPServerManager: GuardrailRaisedException, HTTPException, ) as e: - 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 @@ -3736,9 +3378,7 @@ class MCPServerManager: candidate.server_name, candidate.name, ): - if identifier and normalize_server_name(identifier) == ( - normalized_server_name - ): + if identifier and normalize_server_name(identifier) == (normalized_server_name): return True return False @@ -3750,9 +3390,7 @@ class MCPServerManager: break if mcp_server is None: fallback = self._get_mcp_server_from_tool_name(name) - if fallback is not None and ( - not server_name or _candidate_matches_server_name(fallback) - ): + if fallback is not None and (not server_name or _candidate_matches_server_name(fallback)): mcp_server = fallback if mcp_server is None: raise ValueError(f"Tool {name} not found") @@ -3767,9 +3405,7 @@ class MCPServerManager: return mcp_server - async def has_user_oauth_token( - self, server: MCPServer, user_api_key_auth: Optional[UserAPIKeyAuth] - ) -> bool: + async def has_user_oauth_token(self, server: MCPServer, user_api_key_auth: Optional[UserAPIKeyAuth]) -> bool: """Whether the v2 resolver can produce a per-user token for this server right now. This is the preemptive 401's existence check, routed through the same resolver that drives @@ -3779,9 +3415,7 @@ class MCPServerManager: spec = to_server_spec(server) if spec is None: return False - return await self._cred_provider.has_user_token( - to_subject(user_api_key_auth, None), spec - ) + return await self._cred_provider.has_user_token(to_subject(user_api_key_auth, None), spec) async def _resolve_oauth2_headers_for_tool_call( self, @@ -3790,11 +3424,7 @@ class MCPServerManager: user_api_key_auth: Optional[UserAPIKeyAuth], ) -> Optional[Dict[str, str]]: """Look up per-user OAuth headers when the client did not supply a token.""" - if ( - not mcp_server.needs_user_oauth_token - or oauth2_headers - or user_api_key_auth is None - ): + if not mcp_server.needs_user_oauth_token or oauth2_headers or user_api_key_auth is None: return oauth2_headers if to_server_spec(mcp_server) is not None: @@ -3842,9 +3472,7 @@ class MCPServerManager: GuardrailRaisedException, HTTPException, ) as e: - 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 async def call_tool( @@ -3911,15 +3539,11 @@ class MCPServerManager: ) tasks.append(during_hook_task) - oauth2_headers = await self._resolve_oauth2_headers_for_tool_call( - mcp_server, oauth2_headers, user_api_key_auth - ) + oauth2_headers = await self._resolve_oauth2_headers_for_tool_call(mcp_server, oauth2_headers, user_api_key_auth) # 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 " @@ -3928,11 +3552,7 @@ class MCPServerManager: "transport to enable hook header injection.", server_name, ) - tasks.append( - asyncio.create_task( - self._call_openapi_tool_handler(mcp_server, name, arguments) - ) - ) + tasks.append(asyncio.create_task(self._call_openapi_tool_handler(mcp_server, name, arguments))) else: return await self._call_regular_mcp_tool( mcp_server=mcp_server, @@ -3961,9 +3581,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)}" @@ -4033,9 +3651,7 @@ class MCPServerManager: # If not found and tool name is prefixed, extract the prefix and # match against any known form. - if is_tool_name_prefixed( - tool_name, known_server_prefixes=set(prefix_to_server.keys()) - ): + if is_tool_name_prefixed(tool_name, known_server_prefixes=set(prefix_to_server.keys())): ( original_tool_name, server_name_from_prefix, @@ -4061,9 +3677,7 @@ class MCPServerManager: self._upstream_initialize_instructions_probed_at.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 @@ -4105,16 +3719,12 @@ 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})") # raw_rows come straight from the DB, so their global env var # values (like credentials) are still encrypted here, unlike the # already-decrypted records add_server/update_server are handed. # Decrypt them while building the registry entry. - new_server = await self.build_mcp_server_from_table( - server, env_vars_are_encrypted=True - ) + new_server = await self.build_mcp_server_from_table(server, env_vars_are_encrypted=True) # Carry the cached short_prefix from the previous registry entry # (if any) so the prefix is stable across reloads. if existing_server is not None and existing_server.short_prefix: @@ -4138,9 +3748,7 @@ class MCPServerManager: # Register OpenAPI tools *after* the final short prefix is assigned # so the tools are stored in the global registry under the same # prefix that lookups will use. - await self._maybe_register_openapi_tools( - new_server, initialize_mapping=False - ) + await self._maybe_register_openapi_tools(new_server, initialize_mapping=False) registered_registry[server_id] = new_server if new_server.spec_path: registered_openapi_tools = True @@ -4156,9 +3764,7 @@ class MCPServerManager: if registered_openapi_tools: self.initialize_tool_name_to_mcp_server_name_mapping() - verbose_logger.debug( - "MCP registry refreshed (%s servers in registry)", len(registered_registry) - ) + verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry)) def get_mcp_servers_from_ids(self, server_ids: List[str]) -> List[MCPServer]: servers = [] @@ -4180,9 +3786,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. @@ -4200,9 +3804,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]: @@ -4236,11 +3838,7 @@ class MCPServerManager: if litellm.public_mcp_servers is None: return [] public_ids = set(litellm.public_mcp_servers) - return [ - server - for server in self.get_registry().values() - if server.server_id in public_ids - ] + return [server for server in self.get_registry().values() if server.server_id in public_ids] public_ids = set(litellm.public_mcp_servers or []) return [ @@ -4273,9 +3871,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) @@ -4315,9 +3911,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. @@ -4352,9 +3946,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. @@ -4365,11 +3957,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, @@ -4398,9 +3986,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")) @@ -4466,9 +4052,7 @@ class MCPServerManager: user_api_key_auth=None, raise_on_missing=False, ) - extra_headers = ( - dict(resolved_static_headers) if resolved_static_headers else {} - ) + extra_headers = dict(resolved_static_headers) if resolved_static_headers else {} client = await self._create_mcp_client( server=server, @@ -4483,15 +4067,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" @@ -4504,9 +4084,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, @@ -4605,9 +4183,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, @@ -4673,9 +4249,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/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index e912ecd4204..34491352af4 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -91,13 +91,21 @@ const CreateMCPServer: React.FC = ({ const [oauthDocsUrl, setOauthDocsUrl] = useState(null); // Single hook call shared by MCPConnectionStatus and MCPToolConfiguration to avoid duplicate requests. - const { tools, isLoadingTools, toolsError, toolsErrorStatus, toolsErrorStackTrace, canFetchTools, fetchTools, clearTools } = - useTestMCPConnection({ - accessToken, - oauthAccessToken, - formValues, - enabled: true, - }); + const { + tools, + isLoadingTools, + toolsError, + toolsErrorStatus, + toolsErrorStackTrace, + canFetchTools, + fetchTools, + clearTools, + } = useTestMCPConnection({ + accessToken, + oauthAccessToken, + formValues, + enabled: true, + }); const authType = formValues.auth_type as string | undefined; const shouldShowAuthValueField = authType ? AUTH_TYPES_REQUIRING_AUTH_VALUE.includes(authType) : false;