From c5e40e6345ff607b11b4b029559a504c2927374a Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 1 Jul 2026 11:16:08 +0530 Subject: [PATCH] perf(mcp): cache BYOM submitter server lookup with 60s TTL Co-authored-by: Cursor --- .../mcp_server/mcp_server_manager.py | 839 +++++++++++++----- 1 file changed, 636 insertions(+), 203 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 740cc13f736..b7d2aaf78bf 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -166,7 +166,9 @@ 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 @@ -208,7 +210,10 @@ 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 @@ -217,8 +222,12 @@ 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)) @@ -283,7 +292,10 @@ 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 @@ -302,7 +314,9 @@ def _warn_on_server_name_fields( if result.is_valid: return - warning_text = "; ".join(result.warnings) if result.warnings else "Validation failed" + warning_text = ( + "; ".join(result.warnings) if result.warnings else "Validation failed" + ) verbose_logger.warning( "MCP server '%s' has invalid %s '%s': %s", server_id, @@ -315,7 +329,9 @@ 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 @@ -377,7 +393,10 @@ 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: @@ -445,7 +464,9 @@ def _create_sampling_callback(user_api_key_auth: Optional[Any] = None): ) auth_context = get_active_auth_context() - resolved_auth = user_api_key_auth or (auth_context.user_api_key_auth if auth_context else None) + resolved_auth = user_api_key_auth or ( + auth_context.user_api_key_auth if auth_context else None + ) # Forward original HTTP headers and client IP so that # header-dependent guardrails, tag-based routing, trace # correlation, and forward_llm_provider_auth_headers work @@ -484,7 +505,11 @@ def _create_elicitation_callback(): # In Gateway mode, we relay the elicitation request to the downstream client # that triggered the current operation. downstream_session = get_active_mcp_session() - downstream_capabilities = getattr(downstream_session, "capabilities", None) if downstream_session else None + downstream_capabilities = ( + getattr(downstream_session, "capabilities", None) + if downstream_session + else None + ) return await handle_elicitation_request( context=context, @@ -516,7 +541,9 @@ 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 @@ -562,12 +589,18 @@ 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. @@ -602,13 +635,20 @@ 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( @@ -616,7 +656,9 @@ 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, @@ -627,7 +669,9 @@ 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( @@ -680,10 +724,15 @@ class MCPServerManager: if mcp_aliases and alias is None: # Check if this server_name has an alias in mcp_aliases for alias_name, target_server_name in mcp_aliases.items(): - if target_server_name == server_name and alias_name not in used_aliases: + if ( + target_server_name == server_name + and alias_name not in used_aliases + ): alias = alias_name used_aliases.add(alias_name) - verbose_logger.debug(f"Mapped alias '{alias_name}' to server '{server_name}'") + verbose_logger.debug( + f"Mapped alias '{alias_name}' to server '{server_name}'" + ) break # Create a temporary server object to use with get_server_prefix utility @@ -718,7 +767,9 @@ class MCPServerManager: else: mcp_oauth_metadata = None - resolved_scopes = server_config.get("scopes") or (mcp_oauth_metadata.scopes if mcp_oauth_metadata else None) + resolved_scopes = server_config.get("scopes") or ( + mcp_oauth_metadata.scopes if mcp_oauth_metadata else None + ) resolved_authorization_url = server_config.get("authorization_url") or ( mcp_oauth_metadata.authorization_url if mcp_oauth_metadata else None ) @@ -754,11 +805,15 @@ 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), @@ -768,8 +823,12 @@ 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), @@ -781,7 +840,9 @@ 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", @@ -798,18 +859,24 @@ class MCPServerManager: # Check if this is an OpenAPI-based server spec_path = server_config.get("spec_path", None) if spec_path: - verbose_logger.info(f"Loading OpenAPI spec from {spec_path} for server {server_name}") + verbose_logger.info( + f"Loading OpenAPI spec from {spec_path} for server {server_name}" + ) await self._register_openapi_tools( spec_path=spec_path, server=new_server, base_url=server_config.get("url", ""), ) - verbose_logger.debug(f"Loaded MCP Servers: {json.dumps(self.config_mcp_servers, indent=4, default=str)}") + verbose_logger.debug( + f"Loaded MCP Servers: {json.dumps(self.config_mcp_servers, indent=4, default=str)}" + ) self.initialize_tool_name_to_mcp_server_name_mapping() - async def _register_openapi_tools(self, spec_path: str, server: MCPServer, base_url: str): + async def _register_openapi_tools( + self, spec_path: str, server: MCPServer, base_url: str + ): """ Register tools from an OpenAPI specification for a given server. @@ -845,7 +912,9 @@ class MCPServerManager: # Use base_url from config if provided, otherwise extract from spec if not base_url: base_url = get_openapi_base_url(spec, spec_path) - verbose_logger.info(f"Registering OpenAPI tools for server {server.name} with base URL: {base_url}") + verbose_logger.info( + f"Registering OpenAPI tools for server {server.name} with base URL: {base_url}" + ) # Get server prefix for tool naming server_prefix = get_server_prefix(server) @@ -899,14 +968,20 @@ class MCPServerManager: operation = path_item[method] # Resolve $ref params and merge path-level params into the operation. - resolved_operation = resolve_operation_params(operation, path_item, components) + resolved_operation = resolve_operation_params( + operation, path_item, components + ) # Generate tool name (without prefix initially) - operation_id = operation.get("operationId", f"{method}_{path.replace('/', '_')}") + operation_id = operation.get( + "operationId", f"{method}_{path.replace('/', '_')}" + ) base_tool_name = operation_id.replace(" ", "_").lower() # Add server prefix to tool name - prefixed_tool_name = add_server_prefix_to_name(base_tool_name, server_prefix) + prefixed_tool_name = add_server_prefix_to_name( + base_tool_name, server_prefix + ) # Get description description = operation.get( @@ -918,7 +993,9 @@ class MCPServerManager: input_schema = build_input_schema(resolved_operation) # Create tool function with headers using imported function - tool_func = create_tool_function(path, method, resolved_operation, base_url, headers=headers) + tool_func = create_tool_function( + path, method, resolved_operation, base_url, headers=headers + ) tool_func.__name__ = prefixed_tool_name tool_func.__doc__ = description @@ -931,16 +1008,26 @@ class MCPServerManager: ) # Update tool name to server name mapping (for both prefixed and base names) - self.tool_name_to_mcp_server_name_mapping[base_tool_name] = server_prefix - self.tool_name_to_mcp_server_name_mapping[prefixed_tool_name] = server_prefix + self.tool_name_to_mcp_server_name_mapping[base_tool_name] = ( + server_prefix + ) + self.tool_name_to_mcp_server_name_mapping[prefixed_tool_name] = ( + server_prefix + ) registered_count += 1 - verbose_logger.debug(f"Registered OpenAPI tool: {prefixed_tool_name} for server {server.name}") + verbose_logger.debug( + f"Registered OpenAPI tool: {prefixed_tool_name} for server {server.name}" + ) - verbose_logger.info(f"Successfully registered {registered_count} OpenAPI tools for server {server.name}") + verbose_logger.info( + f"Successfully registered {registered_count} OpenAPI tools for server {server.name}" + ) except Exception as e: - verbose_logger.error(f"Failed to register OpenAPI tools for server {server.name}: {str(e)}") + verbose_logger.error( + f"Failed to register OpenAPI tools for server {server.name}: {str(e)}" + ) raise e def _cleanup_server_tool_routing_artifacts(self, server: MCPServer) -> None: @@ -971,7 +1058,9 @@ 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: @@ -988,10 +1077,14 @@ 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, @@ -1017,14 +1110,20 @@ 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 @@ -1071,7 +1170,9 @@ class MCPServerManager: client_secret_value = encrypted_client_secret # AWS SigV4 credential fields - aws_creds = self._extract_aws_credentials(credentials_dict, credentials_are_encrypted) + aws_creds = self._extract_aws_credentials( + credentials_dict, credentials_are_encrypted + ) scopes: Optional[List[str]] = None if credentials_dict: @@ -1079,7 +1180,9 @@ class MCPServerManager: if scopes_value is not None: scopes = self._extract_scopes(scopes_value) - name_for_prefix = mcp_server.alias or mcp_server.server_name or mcp_server.server_id + name_for_prefix = ( + mcp_server.alias or mcp_server.server_name or mcp_server.server_id + ) mcp_info: MCPInfo = _mcp_info.copy() if "server_name" not in mcp_info: @@ -1090,14 +1193,20 @@ class MCPServerManager: auth_type = cast(MCPAuthType, mcp_server.auth_type) server_url = mcp_server.url - needs_discovery = bool(server_url) and auth_type == MCPAuth.oauth2 and not mcp_server.authorization_url + needs_discovery = ( + bool(server_url) + and auth_type == MCPAuth.oauth2 + and not mcp_server.authorization_url + ) mcp_oauth_metadata = ( await self._descovery_metadata(server_url=server_url) # type: ignore[arg-type] if needs_discovery else None ) - resolved_scopes = scopes or (mcp_oauth_metadata.scopes if mcp_oauth_metadata else None) + resolved_scopes = scopes or ( + mcp_oauth_metadata.scopes if mcp_oauth_metadata else None + ) new_server = MCPServer( server_id=mcp_server.server_id, @@ -1114,22 +1223,30 @@ 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 [], @@ -1138,13 +1255,21 @@ 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), @@ -1159,19 +1284,29 @@ 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, @@ -1195,7 +1330,9 @@ 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) @@ -1220,7 +1357,9 @@ 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 @@ -1244,9 +1383,15 @@ class MCPServerManager: def get_allow_all_keys_server_ids(self) -> List[str]: """Return server IDs that bypass per-key restrictions.""" - return [server.server_id for server in self.get_registry().values() if server.allow_all_keys is True] + return [ + server.server_id + for server in self.get_registry().values() + if server.allow_all_keys is True + ] - async def get_allowed_mcp_servers(self, user_api_key_auth: Optional[UserAPIKeyAuth] = None) -> List[str]: + async def get_allowed_mcp_servers( + self, user_api_key_auth: Optional[UserAPIKeyAuth] = None + ) -> List[str]: """ Get the allowed MCP Servers for the user. @@ -1262,9 +1407,12 @@ 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 [] @@ -1279,13 +1427,23 @@ class MCPServerManager: ) # If admin but NO explicit object permission, get all servers - if user_api_key_auth and _user_has_admin_view(user_api_key_auth) and not has_explicit_object_permission: - verbose_logger.debug("Admin user without explicit object_permission - returning all servers") + if ( + user_api_key_auth + and _user_has_admin_view(user_api_key_auth) + and not has_explicit_object_permission + ): + verbose_logger.debug( + "Admin user without explicit object_permission - returning all servers" + ) return list(self.get_registry().keys()) # Get allowed servers from object permissions (respects object_permission even for admins) - allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth) - verbose_logger.debug(f"Allowed MCP Servers for user api key auth: {allowed_mcp_servers}") + allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers( + user_api_key_auth + ) + verbose_logger.debug( + f"Allowed MCP Servers for user api key auth: {allowed_mcp_servers}" + ) combined_servers = set(allowed_mcp_servers) # Only skip allow_all_keys servers when the request is inside a toolset # scope. toolset_mcp_route / dynamic_mcp_route set _mcp_active_toolset_id @@ -1327,26 +1485,45 @@ class MCPServerManager: # BYOM: approved submissions stay visible to the submitter even when # allow_all_keys=false and no access groups were configured at approval. - submitter_user_id = getattr(user_api_key_auth, "user_id", None) if user_api_key_auth else None + # 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 + ) if submitter_user_id: from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 get_active_submitted_mcp_server_ids_for_user, ) - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache - if prisma_client is not None: - submitted_server_ids = await get_active_submitted_mcp_server_ids_for_user( - prisma_client, - submitter_user_id, + 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 + ) + if submitted_server_ids is None: + submitted_server_ids = ( + await get_active_submitted_mcp_server_ids_for_user( + prisma_client, submitter_user_id + ) + if prisma_client is not None + else [] ) - combined_servers.update( - server_id - for server_id in submitted_server_ids - if self.get_mcp_server_by_id(server_id) is not None + await user_api_key_cache.async_set_cache( + key=byom_cache_key, + value=submitted_server_ids, + 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 + ) 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( @@ -1426,12 +1603,15 @@ class MCPServerManager: keys_to_remove = [ k for k in cache_dict - if (k.startswith("toolset_perms:") and toolset_id in k) or k.startswith("toolset_name:") + if (k.startswith("toolset_perms:") and toolset_id in k) + or k.startswith("toolset_name:") ] for k in keys_to_remove: cache_dict.pop(k, None) except Exception as e: - verbose_logger.warning(f"invalidate_toolset_cache: failed to evict in-memory entries: {e}") + verbose_logger.warning( + f"invalidate_toolset_cache: failed to evict in-memory entries: {e}" + ) async def get_toolset_by_name_cached( self, @@ -1468,12 +1648,18 @@ class MCPServerManager: toolset = await get_mcp_toolset_by_name(prisma_client, toolset_name) await user_api_key_cache.async_set_cache( key=cache_key, - value=(toolset.model_dump(mode="json") if toolset is not None else "__not_found__"), + value=( + toolset.model_dump(mode="json") + if toolset is not None + else "__not_found__" + ), ttl=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. @@ -1515,7 +1701,9 @@ class MCPServerManager: return [] return await self._get_tools_from_server(server) except Exception as e: - verbose_logger.warning(f"Failed to get tools from server {server_id}: {str(e)}") + verbose_logger.warning( + f"Failed to get tools from server {server_id}: {str(e)}" + ) return [] async def list_tools( @@ -1584,7 +1772,9 @@ class MCPServerManager: # Flatten results into single list list_tools_result: List[MCPTool] = [tool for tools in results for tool in tools] - verbose_logger.info(f"Successfully fetched {len(list_tools_result)} tools total from all servers") + verbose_logger.info( + f"Successfully fetched {len(list_tools_result)} tools total from all servers" + ) return list_tools_result ######################################################### @@ -1706,7 +1896,9 @@ 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: @@ -1726,15 +1918,21 @@ 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, @@ -1746,7 +1944,9 @@ 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 @@ -1841,20 +2041,34 @@ 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. @@ -1896,7 +2110,9 @@ 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, @@ -1907,7 +2123,9 @@ 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 @@ -1918,7 +2136,10 @@ 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): @@ -1931,7 +2152,11 @@ 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, @@ -1956,7 +2181,9 @@ 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, @@ -2023,10 +2250,12 @@ 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 ( @@ -2056,8 +2285,12 @@ 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 @@ -2067,7 +2300,9 @@ class MCPServerManager: sep = MCP_TOOL_PREFIX_SEPARATOR tools = [ ( - t.model_copy(update={"name": t.name[len(prefix) + len(sep) :]}) + t.model_copy( + update={"name": t.name[len(prefix) + len(sep) :]} + ) if t.name.startswith(f"{prefix}{sep}") else t ) @@ -2075,10 +2310,14 @@ 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 @@ -2088,7 +2327,9 @@ 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( @@ -2132,12 +2373,16 @@ class MCPServerManager: prompts = await client.list_prompts() - prefixed_or_original_prompts = self._create_prefixed_prompts(prompts, server, add_prefix=add_prefix) + prefixed_or_original_prompts = self._create_prefixed_prompts( + prompts, server, add_prefix=add_prefix + ) return prefixed_or_original_prompts except Exception as e: - verbose_logger.warning(f"Failed to get prompts from server {server.name}: {str(e)}") + verbose_logger.warning( + f"Failed to get prompts from server {server.name}: {str(e)}" + ) return [] async def get_resources_from_server( @@ -2172,12 +2417,16 @@ class MCPServerManager: resources = await client.list_resources() - prefixed_resources = self._create_prefixed_resources(resources, server, add_prefix=add_prefix) + prefixed_resources = self._create_prefixed_resources( + resources, server, add_prefix=add_prefix + ) return prefixed_resources except Exception as e: - verbose_logger.warning(f"Failed to get resources from server {server.name}: {str(e)}") + verbose_logger.warning( + f"Failed to get resources from server {server.name}: {str(e)}" + ) return [] async def get_resource_templates_from_server( @@ -2219,7 +2468,9 @@ class MCPServerManager: return prefixed_templates except Exception as e: - verbose_logger.warning(f"Failed to get resource templates from server {server.name}: {str(e)}") + verbose_logger.warning( + f"Failed to get resource templates from server {server.name}: {str(e)}" + ) return [] async def read_resource_from_server( @@ -2340,8 +2591,15 @@ 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, @@ -2360,11 +2618,13 @@ class MCPServerManager: header_value: Optional[str] = None if exc.response is not None: - header_value = exc.response.headers.get("WWW-Authenticate") or exc.response.headers.get( - "www-authenticate" - ) + header_value = exc.response.headers.get( + "WWW-Authenticate" + ) or exc.response.headers.get("www-authenticate") - resource_metadata_url, scopes = self._parse_www_authenticate_header(header_value) + resource_metadata_url, scopes = self._parse_www_authenticate_header( + header_value + ) authorization_servers = [] resource_scopes = None @@ -2372,7 +2632,9 @@ 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, @@ -2384,12 +2646,16 @@ class MCPServerManager: try: parsed_url = urlparse(server_url) if parsed_url.scheme and parsed_url.netloc: - authorization_servers = [f"{parsed_url.scheme}://{parsed_url.netloc}"] + authorization_servers = [ + f"{parsed_url.scheme}://{parsed_url.netloc}" + ] except Exception: authorization_servers = [] if authorization_servers: - metadata = await self._fetch_authorization_server_metadata(authorization_servers, 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: @@ -2399,10 +2665,14 @@ class MCPServerManager: return metadata except Exception as exc: # pragma: no cover - network/transient issues - verbose_logger.debug("MCP OAuth discovery failed for %s: %s", server_url, exc) + verbose_logger.debug( + "MCP OAuth discovery failed for %s: %s", server_url, exc + ) return None - def _parse_www_authenticate_header(self, header_value: Optional[str]) -> Tuple[Optional[str], Optional[List[str]]]: + def _parse_www_authenticate_header( + self, header_value: Optional[str] + ) -> Tuple[Optional[str], Optional[List[str]]]: if not header_value: return None, None @@ -2411,7 +2681,8 @@ class MCPServerManager: param_pattern = re.compile(r"([a-zA-Z0-9_]+)\s*=\s*\"?([^\",]+)\"?") params: Dict[str, str] = { - match.group(1).lower(): match.group(2).strip() for match in param_pattern.finditer(params_section) + match.group(1).lower(): match.group(2).strip() + for match in param_pattern.finditer(params_section) } resource_metadata_url = params.get("resource_metadata") @@ -2429,7 +2700,9 @@ 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: @@ -2451,15 +2724,23 @@ class MCPServerManager: raw_servers = data.get("authorization_servers") if isinstance(raw_servers, list): - authorization_servers = [entry for entry in raw_servers if isinstance(entry, str) and entry.strip() != ""] + authorization_servers = [ + entry + for entry in raw_servers + if isinstance(entry, str) and entry.strip() != "" + ] else: authorization_servers = [] - scopes = self._extract_scopes(data.get("scopes_supported") or data.get("scopes")) + scopes = self._extract_scopes( + data.get("scopes_supported") or data.get("scopes") + ) return authorization_servers, scopes - async def _attempt_well_known_discovery(self, server_url: str) -> Tuple[List[str], Optional[List[str]]]: + async def _attempt_well_known_discovery( + self, server_url: str + ) -> Tuple[List[str], Optional[List[str]]]: try: parsed = urlparse(server_url) except Exception: @@ -2491,7 +2772,9 @@ 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 @@ -2512,9 +2795,13 @@ 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("/")) @@ -2565,8 +2852,14 @@ 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] @@ -2676,24 +2969,32 @@ 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, @@ -2764,7 +3065,9 @@ 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. @@ -2798,7 +3101,9 @@ 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( @@ -2825,7 +3130,9 @@ class MCPServerManager: prompt.name = name_to_use prefixed_prompts.append(prompt) - verbose_logger.info(f"Successfully fetched {len(prefixed_prompts)} prompts from server {server.name}") + verbose_logger.info( + f"Successfully fetched {len(prefixed_prompts)} prompts from server {server.name}" + ) return prefixed_prompts def _create_prefixed_resources( @@ -2837,11 +3144,17 @@ class MCPServerManager: prefix = get_server_prefix(server) for resource in resources: - name_to_use = add_server_prefix_to_name(resource.name, prefix) if add_prefix else resource.name + name_to_use = ( + add_server_prefix_to_name(resource.name, prefix) + if add_prefix + else resource.name + ) resource.name = name_to_use prefixed_resources.append(resource) - verbose_logger.info(f"Successfully fetched {len(prefixed_resources)} resources from server {server.name}") + verbose_logger.info( + f"Successfully fetched {len(prefixed_resources)} resources from server {server.name}" + ) return prefixed_resources def _create_prefixed_resource_templates( @@ -2857,7 +3170,9 @@ class MCPServerManager: for resource_template in resource_templates: name_to_use = ( - add_server_prefix_to_name(resource_template.name, prefix) if add_prefix else resource_template.name + add_server_prefix_to_name(resource_template.name, prefix) + if add_prefix + else resource_template.name ) resource_template.name = name_to_use prefixed_templates.append(resource_template) @@ -2878,14 +3193,20 @@ 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. @@ -2912,14 +3233,18 @@ class MCPServerManager: unprefixed_tool_name, _ = split_server_prefix_from_name(tool_name) # Check both prefixed and unprefixed tool names - allowed_params_list = server.allowed_params.get(tool_name) or server.allowed_params.get(unprefixed_tool_name) + allowed_params_list = server.allowed_params.get( + tool_name + ) or server.allowed_params.get(unprefixed_tool_name) # If this tool doesn't have allowed_params specified, allow all params if allowed_params_list is None: return None # Filter arguments to only include allowed parameters - disallowed_params = [param for param in arguments.keys() if param not in allowed_params_list] + disallowed_params = [ + param for param in arguments.keys() if param not in allowed_params_list + ] if disallowed_params: raise HTTPException( @@ -3082,22 +3407,42 @@ 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_end_user_id": ( - getattr(user_api_key_auth, "end_user_id", None) if user_api_key_auth else None + "user_api_key_user_id": ( + getattr(user_api_key_auth, "user_id", None) + if user_api_key_auth + else None + ), + "user_api_key_team_id": ( + getattr(user_api_key_auth, "team_id", None) + if user_api_key_auth + else None + ), + "user_api_key_end_user_id": ( + getattr(user_api_key_auth, "end_user_id", None) + if user_api_key_auth + else None + ), + "user_api_key_hash": ( + getattr(user_api_key_auth, "api_key_hash", None) + if user_api_key_auth + else None ), - "user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None), "incoming_bearer_token": incoming_bearer_token, } # Create MCP request object for processing - mcp_request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs) + mcp_request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs( + pre_hook_kwargs + ) # Convert to LLM format for existing guardrail compatibility - synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs) + synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format( + mcp_request_obj, pre_hook_kwargs + ) hook_result: Dict[str, Any] = {} try: @@ -3109,7 +3454,11 @@ class MCPServerManager: ) if modified_data: # Convert response back to MCP format and apply modifications - modified_kwargs = proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs) + modified_kwargs = ( + proxy_logging_obj._convert_mcp_hook_response_to_kwargs( + modified_data, pre_hook_kwargs + ) + ) if modified_kwargs.get("arguments") != arguments: hook_result["arguments"] = modified_kwargs["arguments"] if modified_kwargs.get("extra_headers"): @@ -3154,7 +3503,9 @@ class MCPServerManager: "user_api_key_auth": user_api_key_auth, } - synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, during_hook_kwargs) + synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format( + request_obj, during_hook_kwargs + ) return asyncio.create_task( proxy_logging_obj.during_call_hook( @@ -3250,7 +3601,9 @@ class MCPServerManager: if extra_headers is None: extra_headers = {} - normalized_raw_headers = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} + normalized_raw_headers = { + str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) + } strip_caller_authorization = _should_strip_caller_authorization( mcp_server=mcp_server, raw_headers=raw_headers, @@ -3271,7 +3624,9 @@ 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 = {} @@ -3323,13 +3678,21 @@ 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, @@ -3343,7 +3706,9 @@ 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 @@ -3371,7 +3736,9 @@ 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 @@ -3383,7 +3750,9 @@ 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") @@ -3398,7 +3767,9 @@ 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 @@ -3408,7 +3779,9 @@ 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, @@ -3417,7 +3790,11 @@ 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: @@ -3465,7 +3842,9 @@ 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( @@ -3532,11 +3911,15 @@ 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 " @@ -3545,7 +3928,11 @@ class MCPServerManager: "transport to enable hook header injection.", server_name, ) - tasks.append(asyncio.create_task(self._call_openapi_tool_handler(mcp_server, name, arguments))) + tasks.append( + asyncio.create_task( + self._call_openapi_tool_handler(mcp_server, name, arguments) + ) + ) else: return await self._call_regular_mcp_tool( mcp_server=mcp_server, @@ -3574,7 +3961,9 @@ class MCPServerManager: """ try: if asyncio.get_running_loop(): - asyncio.create_task(self._initialize_tool_name_to_mcp_server_name_mapping()) + asyncio.create_task( + self._initialize_tool_name_to_mcp_server_name_mapping() + ) except RuntimeError as e: # no running event loop verbose_logger.exception( f"No running event loop - skipping tool name to MCP server name mapping initialization: {str(e)}" @@ -3644,7 +4033,9 @@ 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, @@ -3670,7 +4061,9 @@ 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 @@ -3712,12 +4105,16 @@ 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: @@ -3741,7 +4138,9 @@ 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 @@ -3757,7 +4156,9 @@ 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 = [] @@ -3779,7 +4180,9 @@ class MCPServerManager: # Fallback if proxy_server not available return {} - def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: Optional[str]) -> bool: + def _is_server_accessible_from_ip( + self, server: MCPServer, client_ip: Optional[str] + ) -> bool: """ Check if a server is accessible from the given client IP. @@ -3797,7 +4200,9 @@ class MCPServerManager: return True # Non-public server: only accessible from internal IPs general_settings = self._get_general_settings() - internal_networks = IPAddressUtils.parse_internal_networks(general_settings.get("mcp_internal_ip_ranges")) + internal_networks = IPAddressUtils.parse_internal_networks( + general_settings.get("mcp_internal_ip_ranges") + ) return IPAddressUtils.is_internal_ip(client_ip, internal_networks) def get_mcp_server_by_id(self, server_id: str) -> Optional[MCPServer]: @@ -3831,7 +4236,11 @@ 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 [ @@ -3864,7 +4273,9 @@ class MCPServerManager: matches: List[str] = [ server_id for server_id, server in registry.items() - if server.alias == identifier or server.server_name == identifier or server.name == identifier + if server.alias == identifier + or server.server_name == identifier + or server.name == identifier ] if matches: expanded.update(matches) @@ -3904,7 +4315,9 @@ class MCPServerManager: result.setdefault(server_id, []).extend(tools or []) return result - def get_mcp_server_by_name(self, server_name: str, client_ip: Optional[str] = None) -> Optional[MCPServer]: + def get_mcp_server_by_name( + self, server_name: str, client_ip: Optional[str] = None + ) -> Optional[MCPServer]: """ Get the MCP Server from the server name. @@ -3939,7 +4352,9 @@ class MCPServerManager: return server return None - def get_filtered_registry(self, client_ip: Optional[str] = None) -> Dict[str, MCPServer]: + def get_filtered_registry( + self, client_ip: Optional[str] = None + ) -> Dict[str, MCPServer]: """ Get registry filtered by client IP access control. @@ -3950,7 +4365,11 @@ class MCPServerManager: registry = self.get_registry() if client_ip is None: return registry - return {k: v for k, v in registry.items() if self._is_server_accessible_from_ip(v, client_ip)} + return { + k: v + for k, v in registry.items() + if self._is_server_accessible_from_ip(v, client_ip) + } def _generate_stable_server_id( self, @@ -3979,7 +4398,9 @@ class MCPServerManager: A deterministic server ID string """ # Create a string from all the identifying parameters - params_string = f"{server_name}|{url}|{transport}|{auth_type or ''}|{alias or ''}" + params_string = ( + f"{server_name}|{url}|{transport}|{auth_type or ''}|{alias or ''}" + ) # Generate SHA-256 hash hash_object = hashlib.sha256(params_string.encode("utf-8")) @@ -4045,7 +4466,9 @@ 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, @@ -4060,11 +4483,15 @@ class MCPServerManager: return "ok" # Add timeout wrapper to prevent hanging - await asyncio.wait_for(client.run_with_session(_noop), timeout=MCP_HEALTH_CHECK_TIMEOUT) + await asyncio.wait_for( + client.run_with_session(_noop), timeout=MCP_HEALTH_CHECK_TIMEOUT + ) self._remember_upstream_initialize_instructions(server, client) status = "healthy" except asyncio.TimeoutError: - health_check_error = f"Health check timed out after {MCP_HEALTH_CHECK_TIMEOUT} seconds" + health_check_error = ( + f"Health check timed out after {MCP_HEALTH_CHECK_TIMEOUT} seconds" + ) status = "unhealthy" except asyncio.CancelledError: health_check_error = "Health check was cancelled" @@ -4077,7 +4504,9 @@ class MCPServerManager: server_id=server.server_id, server_name=server.server_name, alias=server.alias, - description=(server.mcp_info.get("description") if server.mcp_info else None), + description=( + server.mcp_info.get("description") if server.mcp_info else None + ), url=server.url, transport=server.transport, auth_type=server.auth_type, @@ -4176,7 +4605,9 @@ class MCPServerManager: server_id=server.server_id, server_name=server.server_name, alias=server.alias, - description=(server.mcp_info.get("description") if server.mcp_info else None), + description=( + server.mcp_info.get("description") if server.mcp_info else None + ), url=server.url, spec_path=server.spec_path, transport=server.transport, @@ -4242,7 +4673,9 @@ class MCPServerManager: return await self._run_health_checks(target_server_ids) - async def _run_health_checks(self, target_server_ids: List[str]) -> List[LiteLLM_MCPServerTable]: + async def _run_health_checks( + self, target_server_ids: List[str] + ) -> List[LiteLLM_MCPServerTable]: if not target_server_ids: return []