diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 8911aafa33e..e46e6299277 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1678,10 +1678,17 @@ async def authorize( lookup_name: Final[str | None] = mcp_server_name or client_id client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) mcp_server = ( - global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip) if lookup_name else None + await global_mcp_server_manager.get_resolved_mcp_server_by_name(lookup_name, client_ip=client_ip) + if lookup_name + else None ) if mcp_server is None and mcp_server_name is None: - mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip) + unresolved_server: Final = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip) + mcp_server = ( + await global_mcp_server_manager.ensure_oauth_metadata_discovered(unresolved_server) + if unresolved_server is not None + else None + ) if mcp_server is None: raise HTTPException(status_code=404, detail="MCP server not found") _raise_if_not_oauth2(mcp_server) @@ -1761,9 +1768,14 @@ async def token_endpoint( lookup_name: Final = mcp_server_name or client_id client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) - mcp_server = global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip) + mcp_server = await global_mcp_server_manager.get_resolved_mcp_server_by_name(lookup_name, client_ip=client_ip) if mcp_server is None and mcp_server_name is None: - mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip) + unresolved_server: Final = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip) + mcp_server = ( + await global_mcp_server_manager.ensure_oauth_metadata_discovered(unresolved_server) + if unresolved_server is not None + else None + ) if mcp_server is None: raise HTTPException(status_code=404, detail="MCP server not found") return await exchange_token_with_server( @@ -2565,9 +2577,10 @@ async def register_client(request: Request, mcp_server_name: str | None = None): return await register_aggregate_client(request=request, request_body=data) resolved: Final = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip) if resolved: + resolved_server: Final = await global_mcp_server_manager.ensure_oauth_metadata_discovered(resolved) return await register_client_with_server( request=request, - mcp_server=resolved, + mcp_server=resolved_server, client_name=data.get("client_name", ""), grant_types=data.get("grant_types", []), response_types=data.get("response_types", []), @@ -2577,7 +2590,10 @@ async def register_client(request: Request, mcp_server_name: str | None = None): ) return dummy_return - mcp_server: Final = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip) + mcp_server: Final = await global_mcp_server_manager.get_resolved_mcp_server_by_name( + mcp_server_name, + client_ip=client_ip, + ) if mcp_server is None: return dummy_return return await register_client_with_server( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 97b03b5e60c..7fff6c12fe0 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -15,6 +15,7 @@ import re import time from collections.abc import AsyncIterator, Callable, Mapping, Sequence from contextlib import asynccontextmanager +from dataclasses import dataclass, replace from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypedDict, cast from urllib.parse import ParseResult, urlparse @@ -221,12 +222,43 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: Final[tuple[MCPAuth, ...]] = ( ) -# OAuth discovery retry cooldown for servers whose endpoints stay unresolved. The base is one -# reload cadence so a transient upstream failure recovers immediately; the cap bounds the request -# amplification and log volume of a permanently broken configuration. +_MCP_OAUTH_DISCOVERY_ON_STARTUP_ENV: Final = "LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP" +_TRUE_ENV_VALUES: Final = frozenset(("1", "true", "yes", "on")) +_OAUTH_DISCOVERY_RETRY_DELAYS_SECONDS: Final = (0.05, 0.15) _OAUTH_DISCOVERY_RETRY_BASE_SECONDS: Final = 30.0 _OAUTH_DISCOVERY_RETRY_MAX_SECONDS: Final = 900.0 + +def _oauth_discovery_now() -> float: + return time.monotonic() + + +def _oauth_discovery_retry_delay(consecutive_failures: int) -> float: + backoff_multiplier: Final[int] = 1 << max(consecutive_failures - 1, 0) + return min( + _OAUTH_DISCOVERY_RETRY_BASE_SECONDS * backoff_multiplier, + _OAUTH_DISCOVERY_RETRY_MAX_SECONDS, + ) + + +def _mcp_oauth_discovery_on_startup_enabled() -> bool: + """Return whether remote MCP OAuth metadata is discovered during registration. + + Discovery is deferred until the first admitted request unless explicitly + enabled with ``1``, ``true``, ``yes``, or ``on``. + """ + value: Final = os.getenv(_MCP_OAUTH_DISCOVERY_ON_STARTUP_ENV) + return value is not None and value.strip().lower() in _TRUE_ENV_VALUES + + +def _requires_oauth_discovery( + server_url: str | None, + use_issuer_anchor: bool, + server: MCPServer, +) -> bool: + return _has_oauth_discovery_source(server_url, use_issuer_anchor) and _oauth_endpoints_unresolved(server) + + _StringList: TypeAlias = list[str] _StringMap: TypeAlias = dict[str, str] _ToolParamMap: TypeAlias = dict[str, list[str]] @@ -235,6 +267,34 @@ _InMemoryCacheDict: TypeAlias = dict[str, object] _ToolArguments: TypeAlias = dict[str, object] +@dataclass(frozen=True, slots=True) +class _OAuthDiscoveryResolved: + server: MCPServer + + +@dataclass(frozen=True, slots=True) +class _OAuthDiscoveryFailed: + server_id: str + timed_out: bool + + +@dataclass(frozen=True, slots=True) +class _OAuthDiscoveryStale: + server_id: str + + +_OAuthDiscoveryOutcome: TypeAlias = _OAuthDiscoveryResolved | _OAuthDiscoveryFailed | _OAuthDiscoveryStale + + +@dataclass(frozen=True, slots=True) +class _OAuthDiscoverySlot: + server_id: str + generation: int + task: asyncio.Task[_OAuthDiscoveryOutcome] | None = None + consecutive_failures: int = 0 + retry_not_before: float = 0.0 + + class MCPServerConfig(TypedDict, total=False): """Shape of a single ``mcp_servers`` entry in config.yaml, as consumed by :meth:`MCPServerManager.load_servers_from_config`. Every key is optional: YAML supplies @@ -625,6 +685,7 @@ def _warn_oauth_endpoints_unresolved( server_ref: str, server_url: str | None, discovery_attempted: bool, + discovery_deferred: bool = False, issuer_anchored: bool, metadata: MCPOAuthMetadata | None, needs_authorization_url: bool, @@ -643,7 +704,7 @@ def _warn_oauth_endpoints_unresolved( are needed (client_credentials never needs authorization_url; OBO needs only token_url); the issuer-anchored arm is excluded here because it has its own RFC 8414 ยง3.3 warning. """ - if issuer_anchored: + if discovery_deferred or issuer_anchored: return unresolved: Final = tuple( field @@ -1427,41 +1488,288 @@ class MCPServerManager: # empty result, or failure). Used to throttle re-probes for servers that do # not return instructions, and to apply a short cooldown after failures. self._upstream_initialize_instructions_probed_at: dict[str, float] = {} - # Per-server (consecutive failures, monotonic timestamp) for OAuth discovery retries, so a - # server whose endpoints never resolve backs off instead of re-running the full - # RFC 9728 -> 8414 chain, and re-logging its warning, on every reload forever. - self._oauth_discovery_retry_state: dict[ - str, tuple[int, float] - ] = {} # mutable-ok: retry cooldown cache, keyed per server and pruned on success + self._oauth_discovery_on_startup = _mcp_oauth_discovery_on_startup_enabled() + self._oauth_discovery_generation_counter = 0 + self._oauth_discovery_slots: tuple[_OAuthDiscoverySlot, ...] = () - def _oauth_discovery_retry_due(self, server_id: str) -> bool: - """Whether an unresolved server is due for another discovery attempt. + def _oauth_discovery_slot(self, server_id: str) -> _OAuthDiscoverySlot | None: + return next((slot for slot in self._oauth_discovery_slots if slot.server_id == server_id), None) - The reload fast-path exemption is what retries a failed discovery, so without a cooldown a - permanently unresolvable server re-runs the whole RFC 9728 -> RFC 8414 -> origin-fallback - chain and re-emits its unresolved-endpoints warning on every reload, per server, forever. - Delay doubles per consecutive failure from ``_OAUTH_DISCOVERY_RETRY_BASE_SECONDS`` up to - ``_OAUTH_DISCOVERY_RETRY_MAX_SECONDS``, so a transient outage still recovers on the next - reload while a broken configuration settles to one attempt per cap. - """ - state: Final = self._oauth_discovery_retry_state.get(server_id) - if state is None: - return True - failures, attempted_at = state - backoff_multiplier: Final[int] = 2 ** max(failures - 1, 0) - delay: Final = min( - _OAUTH_DISCOVERY_RETRY_BASE_SECONDS * backoff_multiplier, - _OAUTH_DISCOVERY_RETRY_MAX_SECONDS, + def _remove_oauth_discovery_slot(self, server_id: str) -> None: + self._oauth_discovery_slots = tuple(slot for slot in self._oauth_discovery_slots if slot.server_id != server_id) + + def _store_oauth_discovery_slot(self, slot: _OAuthDiscoverySlot) -> None: + self._oauth_discovery_slots = ( + *(existing for existing in self._oauth_discovery_slots if existing.server_id != slot.server_id), + slot, ) - return (time.monotonic() - attempted_at) >= delay - def _record_oauth_discovery_outcome(self, server: MCPServer) -> None: - """Advance or clear a server's retry cooldown after a rebuild resolved it or did not.""" - if not _oauth_endpoints_unresolved(server): - self._oauth_discovery_retry_state.pop(server.server_id, None) + def _set_oauth_discovery_deferred(self, server_id: str, discovery_deferred: bool) -> None: + previous: Final = self._oauth_discovery_slot(server_id) + self._remove_oauth_discovery_slot(server_id) + if previous is not None and previous.task is not None and not previous.task.done(): + previous.task.cancel() + if discovery_deferred: + self._oauth_discovery_generation_counter += 1 + self._store_oauth_discovery_slot( + _OAuthDiscoverySlot( + server_id=server_id, + generation=self._oauth_discovery_generation_counter, + ) + ) + + def _invalidate_oauth_discovery_state(self, server_id: str) -> None: + previous: Final = self._oauth_discovery_slot(server_id) + self._remove_oauth_discovery_slot(server_id) + if previous is not None and previous.task is not None and not previous.task.done(): + previous.task.cancel() + + def _registered_server(self, server: MCPServer) -> MCPServer: + return self.registry.get(server.server_id) or self.config_mcp_servers.get(server.server_id) or server + + async def _discover_oauth_metadata_for_server(self, server: MCPServer) -> MCPOAuthMetadata | None: + manual_issuer: Final = _blank_to_none(server.issuer) + manual_authorization_url: Final = _blank_to_none(server.authorization_url) + manual_token_url: Final = _blank_to_none(server.token_url) + is_discovery_auth_type: Final = server.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + use_issuer_anchor: Final = server.issuer_is_anchored + obo_needs_discovery: Final = self._obo_needs_endpoint_discovery( + server.auth_type, + server.token_exchange_endpoint, + manual_token_url, + ) + needs_authorization_url: Final = is_discovery_auth_type and server.oauth2_flow != "client_credentials" + needs_token_url: Final = is_discovery_auth_type or obo_needs_discovery + warn_on_empty_discovery: Final = _discovery_failure_leaves_needs_unresolved( + needs_authorization_url=needs_authorization_url, + needs_token_url=needs_token_url, + manual_authorization_url=manual_authorization_url, + manual_token_url=manual_token_url, + ) + metadata: Final = await ( + self._fetch_issuer_anchored_oauth_metadata(manual_issuer, server.url) + if use_issuer_anchor and manual_issuer is not None + else self._descovery_metadata( + server_url=server.url or "", + allow_origin_fallback=is_discovery_auth_type, + warn_when_no_metadata=warn_on_empty_discovery, + ) + ) + if use_issuer_anchor: + return metadata + gated_metadata: Final = ( + _restrict_discovery_to_corroborated_authorization_server( + metadata, + manual_authorization_url, + server.server_id, + server.is_dcr_bridge, + ) + if is_discovery_auth_type + else metadata + ) + _warn_oauth_endpoints_unresolved( + server_ref=server.alias or server.server_name or server.server_id, + server_url=server.url, + discovery_attempted=True, + issuer_anchored=False, + metadata=gated_metadata, + needs_authorization_url=needs_authorization_url, + needs_token_url=needs_token_url, + manual_authorization_url=manual_authorization_url, + manual_token_url=manual_token_url, + ) + return gated_metadata + + @staticmethod + def _merge_discovered_oauth_metadata(server: MCPServer, metadata: MCPOAuthMetadata | None) -> MCPServer: + if metadata is None: + return server + discovered_issuer: Final = metadata.discovered_issuer if not metadata.from_origin_fallback else None + resolved: Final = server.model_copy() + resolved.scopes = server.scopes or metadata.scopes + resolved.issuer = server.issuer or discovered_issuer + resolved.authorization_url = server.authorization_url or metadata.authorization_url + resolved.token_url = server.token_url or metadata.token_url + resolved.registration_url = server.registration_url or metadata.registration_url + return resolved + + def _oauth_discovery_slot_is_current(self, server_id: str, generation: int) -> bool: + slot: Final = self._oauth_discovery_slot(server_id) + return slot is not None and slot.generation == generation + + def _publish_resolved_oauth_server( + self, + server: MCPServer, + generation: int, + ) -> MCPServer | None: + if not self._oauth_discovery_slot_is_current(server.server_id, generation): + return None + if server.server_id in self.registry: + self.registry[server.server_id] = server + elif server.server_id in self.config_mcp_servers: + self.config_mcp_servers[server.server_id] = server + else: + return None + self._remove_oauth_discovery_slot(server.server_id) + return server + + async def _attempt_oauth_metadata_once( + self, + server: MCPServer, + generation: int, + ) -> _OAuthDiscoveryOutcome | None: + if not self._oauth_discovery_slot_is_current(server.server_id, generation): + return _OAuthDiscoveryStale(server_id=server.server_id) + current: Final = self._registered_server(server) + if not _oauth_endpoints_unresolved(current): + published: Final = self._publish_resolved_oauth_server(current, generation) + return ( + _OAuthDiscoveryResolved(server=published) + if published is not None + else _OAuthDiscoveryStale(server_id=server.server_id) + ) + metadata: Final = await self._discover_oauth_metadata_for_server(current) + if not self._oauth_discovery_slot_is_current(server.server_id, generation): + return _OAuthDiscoveryStale(server_id=server.server_id) + candidate: Final = self._merge_discovered_oauth_metadata(self._registered_server(server), metadata) + if _oauth_endpoints_unresolved(candidate): + return None + published_candidate: Final = self._publish_resolved_oauth_server(candidate, generation) + return ( + _OAuthDiscoveryResolved(server=published_candidate) + if published_candidate is not None + else _OAuthDiscoveryStale(server_id=server.server_id) + ) + + async def _attempt_oauth_metadata_resolution( + self, + server: MCPServer, + generation: int, + retry_delays: tuple[float, ...] = _OAUTH_DISCOVERY_RETRY_DELAYS_SECONDS, + ) -> _OAuthDiscoveryOutcome: + outcome: Final = await self._attempt_oauth_metadata_once(server, generation) + if outcome is not None: + return outcome + if not retry_delays: + return _OAuthDiscoveryFailed(server_id=server.server_id, timed_out=False) + await asyncio.sleep(retry_delays[0]) + return await self._attempt_oauth_metadata_resolution(server, generation, retry_delays[1:]) + + async def _run_oauth_metadata_resolution( + self, + server: MCPServer, + generation: int, + ) -> _OAuthDiscoveryOutcome: + try: + outcome: Final = await asyncio.wait_for( + self._attempt_oauth_metadata_resolution(server, generation), + timeout=MCP_METADATA_TIMEOUT, + ) + except asyncio.TimeoutError: + verbose_logger.warning( + "Deferred MCP OAuth discovery timed out after %ss for server %s", + MCP_METADATA_TIMEOUT, + server.server_id, + ) + failure: Final = _OAuthDiscoveryFailed(server_id=server.server_id, timed_out=True) + self._record_oauth_discovery_failure(server.server_id, generation) + return failure + if isinstance(outcome, _OAuthDiscoveryFailed): + self._record_oauth_discovery_failure(server.server_id, generation) + return outcome + + def _record_oauth_discovery_failure(self, server_id: str, generation: int) -> None: + slot: Final = self._oauth_discovery_slot(server_id) + if slot is None or slot.generation != generation: return - failures, _ = self._oauth_discovery_retry_state.get(server.server_id, (0, 0.0)) - self._oauth_discovery_retry_state[server.server_id] = (failures + 1, time.monotonic()) + consecutive_failures: Final = slot.consecutive_failures + 1 + self._store_oauth_discovery_slot( + replace( + slot, + consecutive_failures=consecutive_failures, + retry_not_before=_oauth_discovery_now() + _oauth_discovery_retry_delay(consecutive_failures), + ) + ) + + def _get_or_start_oauth_discovery_task( + self, + server: MCPServer, + ) -> tuple[asyncio.Task[_OAuthDiscoveryOutcome], int] | None: + slot: Final = self._oauth_discovery_slot(server.server_id) + if slot is None: + return None + if slot.task is not None: + if not slot.task.done() or _oauth_discovery_now() < slot.retry_not_before: + return slot.task, slot.generation + task: Final = asyncio.create_task( + self._run_oauth_metadata_resolution(self._registered_server(server), slot.generation) + ) + self._store_oauth_discovery_slot(replace(slot, task=task)) + return task, slot.generation + + def prime_oauth_metadata_discovery(self, server: MCPServer) -> None: + """Start best-effort OAuth metadata discovery for ``server``. + + The call returns immediately and never delays registration. It is a no-op + when the server has no deferred discovery slot. + + Args: + server: The registered MCP server to warm metadata for. + """ + self._get_or_start_oauth_discovery_task(server) + + def _prime_oauth_metadata_discovery_for_servers(self, servers: Sequence[MCPServer]) -> None: + for server in servers: + self.prime_oauth_metadata_discovery(server) + + def _reconcile_oauth_discovery_slots_for_servers(self, servers: Sequence[MCPServer]) -> None: + """Align retry slots after an atomic registry replacement.""" + for server in servers: + should_defer = _requires_oauth_discovery(server.url, server.issuer_is_anchored, server) + has_slot = self._oauth_discovery_slot(server.server_id) is not None + if should_defer != has_slot: + self._set_oauth_discovery_deferred(server.server_id, should_defer) + + async def ensure_oauth_metadata_discovered(self, server: MCPServer) -> MCPServer: + """Join the bounded discovery task and return the resolved server. + + Concurrent callers share one task per server. A failed attempt remains + retryable after a per-server cooldown. + + Args: + server: The MCP server whose OAuth metadata must be resolved. + + Returns: + The resolved server, or the registered server when no discovery is + pending. + + Raises: + HTTPException: Status 503 when discovery times out or returns + incomplete metadata. + """ + acquisition: Final = self._get_or_start_oauth_discovery_task(server) + if acquisition is None: + return self._registered_server(server) + task, generation = acquisition + try: + outcome: Final = await asyncio.shield(task) + except asyncio.CancelledError: + if task.cancelled() and not self._oauth_discovery_slot_is_current(server.server_id, generation): + return await self.ensure_oauth_metadata_discovered(server) + raise + match outcome: + case _OAuthDiscoveryResolved(resolved_server): + return resolved_server + case _OAuthDiscoveryStale(): + return await self.ensure_oauth_metadata_discovered(server) + case _OAuthDiscoveryFailed(timed_out=timed_out): + current: Final = self._registered_server(server) + server_ref: Final = current.alias or current.server_name or current.name or current.server_id + reason: Final = "timed out" if timed_out else "returned incomplete metadata" + raise HTTPException( + status_code=503, + detail=f"OAuth metadata discovery {reason} for MCP server {server_ref!r}", + ) def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None: raw: Final[str | None] = getattr(client, "_last_initialize_instructions", None) @@ -1655,7 +1963,8 @@ class MCPServerManager: manual_authorization_url=manual_authorization_url, manual_token_url=manual_token_url, ) - if not should_discover: + discovery_deferred = should_discover and not self._oauth_discovery_on_startup + if not should_discover or discovery_deferred: mcp_oauth_metadata = None elif use_issuer_anchor and manual_issuer is not None: mcp_oauth_metadata = await self._fetch_issuer_anchored_oauth_metadata(manual_issuer, server_url) @@ -1733,6 +2042,7 @@ class MCPServerManager: server_ref=server_name or server_id, server_url=server_url, discovery_attempted=should_discover, + discovery_deferred=discovery_deferred, issuer_anchored=use_issuer_anchor, metadata=gated_oauth_metadata, needs_authorization_url=needs_authorization_url, @@ -1814,6 +2124,10 @@ class MCPServerManager: self._assign_unique_short_prefix(new_server) _warn_internal_delegate_pkce_if_applicable(new_server, source="config") self.config_mcp_servers[server_id] = new_server + self._set_oauth_discovery_deferred( + server_id, + _requires_oauth_discovery(server_url, use_issuer_anchor, new_server), + ) # Check if this is an OpenAPI-based server spec_path = server_config.get("spec_path", None) @@ -1831,6 +2145,8 @@ class MCPServerManager: await self._hydrate_config_servers_dcr_clients() + self._prime_oauth_metadata_discovery_for_servers(tuple(self.config_mcp_servers.values())) + self.initialize_tool_name_to_mcp_server_name_mapping() async def _hydrate_config_servers_dcr_clients(self) -> None: @@ -2032,6 +2348,7 @@ class MCPServerManager: if evicted is not None: verbose_logger.debug("Removed MCP Server: %s", mcp_server.server_id or mcp_server.server_name) self._cleanup_server_tool_routing_artifacts(evicted) + self._invalidate_oauth_discovery_state(evicted.server_id) else: verbose_logger.warning("Server ID %s not found in registry", mcp_server.server_id) @@ -2063,7 +2380,7 @@ class MCPServerManager: use_issuer_anchor: bool, scopes: list[str] | None, token_exchange_endpoint: str | None, - ) -> MCPOAuthMetadata | None: + ) -> tuple[MCPOAuthMetadata | None, bool]: obo_needs_discovery = self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url) needs_authorization_url: Final = ( is_discovery_auth_type and getattr(mcp_server, "oauth2_flow", None) != "client_credentials" @@ -2079,7 +2396,8 @@ class MCPServerManager: needs_discovery: Final = _has_oauth_discovery_source(server_url, use_issuer_anchor) and ( (is_discovery_auth_type and not has_all_upstream_oauth_fields) or obo_needs_discovery ) - if not needs_discovery: + discovery_deferred: Final = needs_discovery and not self._oauth_discovery_on_startup + if not needs_discovery or discovery_deferred: mcp_oauth_metadata: MCPOAuthMetadata | None = None elif use_issuer_anchor and manual_issuer is not None: mcp_oauth_metadata = await self._fetch_issuer_anchored_oauth_metadata(manual_issuer, server_url) @@ -2090,7 +2408,7 @@ class MCPServerManager: warn_when_no_metadata=warn_on_empty_discovery, ) if use_issuer_anchor: - return mcp_oauth_metadata + return mcp_oauth_metadata, discovery_deferred gated_metadata: Final = ( _restrict_discovery_to_corroborated_authorization_server( mcp_oauth_metadata, @@ -2105,6 +2423,7 @@ class MCPServerManager: server_ref=mcp_server.alias or mcp_server.server_name or mcp_server.server_id, server_url=server_url, discovery_attempted=needs_discovery, + discovery_deferred=discovery_deferred, issuer_anchored=False, metadata=gated_metadata, needs_authorization_url=needs_authorization_url, @@ -2112,7 +2431,7 @@ class MCPServerManager: manual_authorization_url=manual_authorization_url, manual_token_url=manual_token_url, ) - return gated_metadata + return gated_metadata, discovery_deferred async def build_mcp_server_from_table( self, @@ -2220,7 +2539,7 @@ class MCPServerManager: manual_registration_url, mcp_server.alias or mcp_server.server_name or mcp_server.server_id, ) - gated_oauth_metadata: Final = await self._resolve_table_oauth_metadata( + gated_oauth_metadata, _ = await self._resolve_table_oauth_metadata( mcp_server=mcp_server, auth_type=auth_type, server_url=server_url, @@ -2329,6 +2648,10 @@ class MCPServerManager: max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None), ) _warn_internal_delegate_pkce_if_applicable(new_server, source="database") + self._set_oauth_discovery_deferred( + new_server.server_id, + _requires_oauth_discovery(server_url, use_issuer_anchor, new_server), + ) return new_server async def _maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True): @@ -2362,6 +2685,7 @@ class MCPServerManager: self._assign_unique_short_prefix(new_server) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) + self.prime_oauth_metadata_discovery(new_server) verbose_logger.debug("Added MCP Server: %s", new_server.name) except Exception as e: @@ -2378,6 +2702,7 @@ class MCPServerManager: evicted = self.registry.pop(mcp_server.server_name, None) if evicted is not None: self._cleanup_server_tool_routing_artifacts(evicted) + self._invalidate_oauth_discovery_state(evicted.server_id) return try: if mcp_server.server_id in self.registry: @@ -2396,6 +2721,7 @@ class MCPServerManager: self._assign_unique_short_prefix(new_server) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) + self.prime_oauth_metadata_discovery(new_server) verbose_logger.debug("Updated MCP Server: %s", new_server.name) except Exception as e: @@ -3201,7 +3527,8 @@ class MCPServerManager: subject_token: Final = self._extract_bearer_token(oauth2_headers, None) if not subject_token: return - spec: Final = to_server_spec(server) + resolved_server: Final = await self.ensure_oauth_metadata_discovered(server) + spec: Final = to_server_spec(resolved_server) if spec is None or not isinstance(spec.config, TokenExchangeConfig): return match await self._cred_provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec): @@ -3210,7 +3537,7 @@ class MCPServerManager: case Error(err): if err.tag == "unauthorized": raise_token_exchange_challenge( - server, + resolved_server, root_path=get_server_root_path(), claims=err.unauthorized.claims, ) @@ -3246,8 +3573,9 @@ class MCPServerManager: Returns: Configured MCP client instance. """ - transport: Final = server.transport or MCPTransport.sse - spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(server) + resolved_server: Final = await self.ensure_oauth_metadata_discovered(server) + transport: Final = resolved_server.transport or MCPTransport.sse + spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server) provider: Final = cred_provider or self._cred_provider # A caller-supplied per-request override (mcp_auth_header / x-mcp-*) defers to the v1 path # so it wins - except for the modes the v2 resolver owns per-caller (authorization_code's @@ -3266,16 +3594,20 @@ class MCPServerManager: ) ): spec = None - auth_value: Final = await resolve_mcp_auth(server, mcp_auth_header) if spec is None else None + auth_value: Final = await resolve_mcp_auth(resolved_server, mcp_auth_header) 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: Final = _create_elicitation_callback() if server.allow_elicitation else None + sampling_cb = ( + _create_sampling_callback(user_api_key_auth=user_api_key_auth) if resolved_server.allow_sampling else None + ) + elicitation_cb: Final = _create_elicitation_callback() if resolved_server.allow_elicitation else None # Handle stdio transport if transport == MCPTransport.stdio: resolved_env: Final = ( - 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(resolved_server.env) if resolved_server.env is not None else None) ) # Ensure npm-based STDIO MCP servers have a writable cache dir. @@ -3286,8 +3618,8 @@ class MCPServerManager: # Defense-in-depth: block commands not in the allowlist. # The Pydantic validator blocks new servers; this catches legacy # config/DB records predating the allowlist. - if server.command: - base_command: Final = os.path.basename(server.command) + if resolved_server.command: + base_command: Final = os.path.basename(resolved_server.command) # Strip .exe/.cmd/.bat/.com suffix for Windows compatibility base_command_no_ext = base_command.lower() for ext in [".exe", ".cmd", ".bat", ".com"]: @@ -3300,24 +3632,24 @@ class MCPServerManager: ): raise HTTPException( status_code=403, - detail=f"MCP stdio command '{server.command}' is not in the allowlist ({sorted(MCP_STDIO_ALLOWED_COMMANDS)}). " + detail=f"MCP stdio command '{resolved_server.command}' is not in the allowlist ({sorted(MCP_STDIO_ALLOWED_COMMANDS)}). " f"Add it to LITELLM_MCP_STDIO_EXTRA_COMMANDS to allow this command.", ) stdio_config: MCPStdioConfig | None = None - if server.command and server.args is not None: + if resolved_server.command and resolved_server.args is not None: stdio_config = MCPStdioConfig( - command=server.command, - args=server.args, + command=resolved_server.command, + args=resolved_server.args, env=resolved_env, ) return MCPClient( server_url="", # Not used for stdio transport_type=transport, - auth_type=server.auth_type, + auth_type=resolved_server.auth_type, auth_value=auth_value, - timeout=(server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT), + timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT), stdio_config=stdio_config, extra_headers=extra_headers, sampling_callback=sampling_cb, @@ -3325,7 +3657,7 @@ class MCPServerManager: ) else: # For HTTP/SSE transports - server_url: Final = server.url or "" + server_url: Final = resolved_server.url or "" if spec is not None: inbound_token = subject_token @@ -3335,7 +3667,7 @@ class MCPServerManager: if per_server_token is not None: inbound_token = per_server_token resolved_auth, extra_headers = await self._resolve_v2_auth( - server=server, + server=resolved_server, spec=spec, provider=provider, subject_token=inbound_token, @@ -3345,8 +3677,8 @@ class MCPServerManager: return MCPClient( 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), + auth_type=resolved_server.auth_type, + timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT), extra_headers=extra_headers, resolved_auth=resolved_auth, sampling_callback=sampling_cb, @@ -3355,23 +3687,23 @@ class MCPServerManager: # Create SigV4 auth if configured aws_auth = None - if server.auth_type == MCPAuth.aws_sigv4: + if resolved_server.auth_type == MCPAuth.aws_sigv4: aws_auth = MCPSigV4Auth( - aws_access_key_id=server.aws_access_key_id, - aws_secret_access_key=server.aws_secret_access_key, - aws_session_token=server.aws_session_token, - aws_region_name=server.aws_region_name, - aws_service_name=server.aws_service_name, - aws_role_name=server.aws_role_name, - aws_session_name=server.aws_session_name, + aws_access_key_id=resolved_server.aws_access_key_id, + aws_secret_access_key=resolved_server.aws_secret_access_key, + aws_session_token=resolved_server.aws_session_token, + aws_region_name=resolved_server.aws_region_name, + aws_service_name=resolved_server.aws_service_name, + aws_role_name=resolved_server.aws_role_name, + aws_session_name=resolved_server.aws_session_name, ) return MCPClient( server_url=server_url, transport_type=transport, - auth_type=server.auth_type, + auth_type=resolved_server.auth_type, auth_value=auth_value, - timeout=(server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT), + timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT), extra_headers=extra_headers, aws_auth=aws_auth, sampling_callback=sampling_cb, @@ -3827,7 +4159,10 @@ class MCPServerManager: ) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]: origin: Final = _redact_mcp_resource_url(server_url) or "" try: - client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) + client: Final = get_async_httpx_client( + llm_provider=httpxSpecialProvider.MCP, + params={"timeout": MCP_METADATA_TIMEOUT}, # mutable-ok: HTTP client factory requires a dict + ) response: Final = await client.get(server_url) response.raise_for_status() ( @@ -5448,6 +5783,8 @@ class MCPServerManager: Note: This now handles prefixed tool names """ for server in self.get_registry().values(): + if self._oauth_discovery_slot(server.server_id) is not None: + continue if server.needs_user_oauth_token: # Skip OAuth2 servers that rely on user-provided tokens continue @@ -5560,9 +5897,9 @@ class MCPServerManager: and existing_server.updated_at is not None and server.updated_at is not None and existing_server.updated_at == server.updated_at - and not ( - _oauth_endpoints_unresolved(existing_server) - and self._oauth_discovery_retry_due(server.server_id) + and ( + self._oauth_discovery_slot(server.server_id) is not None + or not _oauth_endpoints_unresolved(existing_server) ) ): # Re-use existing server instance to avoid re-running build_mcp_server_from_table() @@ -5581,7 +5918,6 @@ class MCPServerManager: # 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) - self._record_oauth_discovery_outcome(new_server) # 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: @@ -5618,7 +5954,18 @@ class MCPServerManager: e, ) + dropped_registry_keys: Final = previous_registry.keys() - registered_registry.keys() + for registry_key in dropped_registry_keys: + self._invalidate_oauth_discovery_state(previous_registry[registry_key].server_id) + self.registry = registered_registry + # A discovery task may have published into ``previous_registry`` while + # this replacement was being staged. Reconcile every published entry + # synchronously after the swap so a lost publication cannot also leave + # the replacement unresolved with no retry slot. + registered_servers: Final = tuple(registered_registry.values()) + self._reconcile_oauth_discovery_slots_for_servers(registered_servers) + self._prime_oauth_metadata_discovery_for_servers(registered_servers) if registered_openapi_tools: self.initialize_tool_name_to_mcp_server_name_mapping() @@ -5806,6 +6153,14 @@ class MCPServerManager: return server return None + async def get_resolved_mcp_server_by_name( + self, + server_name: str, + client_ip: str | None = None, + ) -> MCPServer | None: + server: Final = self.get_mcp_server_by_name(server_name, client_ip=client_ip) + return await self.ensure_oauth_metadata_discovered(server) if server is not None else None + def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]: """ Get registry filtered by client IP access control. @@ -5900,21 +6255,19 @@ class MCPServerManager: should_skip_health_check = True if not should_skip_health_check: - resolved_static_headers: Final = await self._resolve_static_headers_with_env_vars( - server=server, - user_api_key_auth=None, - raise_on_missing=False, - ) - extra_headers: Final = dict(resolved_static_headers) if resolved_static_headers else {} - - client: Final = await self._create_mcp_client( - server=server, - mcp_auth_header=None, - extra_headers=extra_headers, - stdio_env=None, - ) - try: + resolved_static_headers: Final = await self._resolve_static_headers_with_env_vars( + server=server, + user_api_key_auth=None, + raise_on_missing=False, + ) + extra_headers: Final = dict(resolved_static_headers) if resolved_static_headers else {} + client: Final = await self._create_mcp_client( + server=server, + mcp_auth_header=None, + extra_headers=extra_headers, + stdio_env=None, + ) async def _noop(session): return "ok" diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 17457c3362f..4184fad009c 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -3754,6 +3754,14 @@ if MCP_AVAILABLE: # preemptive challenge and let downstream authorization # return 403. continue + if server is not None and server.auth_type == MCPAuth.oauth2 and server.oauth2_flow == "client_credentials": + # Stamped M2M: the challenge decision below never reads discovered + # metadata, so deferred-discovery failures must not 503 this loop. + # Unstamped rows stay on the discover-first path because filling + # authorization_url/token_url can change their inferred flow. + continue + if server is not None: + server = await global_mcp_server_manager.ensure_oauth_metadata_discovered(server) if server and server.auth_type == MCPAuth.oauth2: # The challenge decision is per oauth2 sub-mode, not per header: # gateway-managed modes (M2M and interactive authorization_code) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 514e34d6241..424b993de85 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -41,6 +41,128 @@ def _mock_callback_request(base_url: str = "http://localhost:3000/"): return req +def _unresolved_oauth_server(): + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + return MCPServer( + server_id="cold-oauth-server", + name="cold_oauth_server", + server_name="cold_oauth_server", + alias="cold_oauth_server", + url="https://mcp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + client_id="client-id", + ) + + +def _resolved_oauth_metadata(): + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata + + return MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + registration_url="https://idp.example.com/register", + scopes=["mcp.read"], + ) + + +@pytest.mark.asyncio +async def test_authorize_resolves_cold_oauth_metadata(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + server = _unresolved_oauth_server() + global_mcp_server_manager.registry[server.server_id] = server + global_mcp_server_manager._set_oauth_discovery_deferred(server.server_id, True) + request = _mock_callback_request("https://litellm.example.com/") + expected = MagicMock() + + with ( + patch.object( + global_mcp_server_manager, + "_discover_oauth_metadata_for_server", + new=AsyncMock(return_value=_resolved_oauth_metadata()), + ) as discovery, + patch.object(discoverable_endpoints, "authorize_with_server", new=AsyncMock(return_value=expected)) as relay, + ): + response = await discoverable_endpoints.authorize( + request=request, + client_id="client-id", + mcp_server_name=server.server_name, + redirect_uri="http://127.0.0.1:60108/callback", + ) + + discovery.assert_awaited_once_with(server) + assert relay.await_args.kwargs["mcp_server"].authorization_url == "https://idp.example.com/authorize" + assert response is expected + + +@pytest.mark.asyncio +async def test_token_resolves_cold_oauth_metadata(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + server = _unresolved_oauth_server() + global_mcp_server_manager.registry[server.server_id] = server + global_mcp_server_manager._set_oauth_discovery_deferred(server.server_id, True) + request = _mock_callback_request("https://litellm.example.com/") + expected = MagicMock() + + with ( + patch.object( + global_mcp_server_manager, + "_discover_oauth_metadata_for_server", + new=AsyncMock(return_value=_resolved_oauth_metadata()), + ) as discovery, + patch.object( + discoverable_endpoints, "exchange_token_with_server", new=AsyncMock(return_value=expected) + ) as relay, + ): + response = await discoverable_endpoints.token_endpoint( + request=request, + grant_type="refresh_token", + client_id="client-id", + refresh_token="refresh-token", + mcp_server_name=server.server_name, + ) + + discovery.assert_awaited_once_with(server) + assert relay.await_args.kwargs["mcp_server"].token_url == "https://idp.example.com/token" + assert response is expected + + +@pytest.mark.asyncio +async def test_register_resolves_cold_oauth_metadata(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + server = _unresolved_oauth_server() + global_mcp_server_manager.registry[server.server_id] = server + global_mcp_server_manager._set_oauth_discovery_deferred(server.server_id, True) + request = _mock_callback_request("https://litellm.example.com/") + expected = MagicMock() + + with ( + patch.object( + global_mcp_server_manager, + "_discover_oauth_metadata_for_server", + new=AsyncMock(return_value=_resolved_oauth_metadata()), + ) as discovery, + patch.object(discoverable_endpoints, "_read_request_body", new=AsyncMock(return_value={})), + patch.object( + discoverable_endpoints, "register_client_with_server", new=AsyncMock(return_value=expected) + ) as relay, + ): + response = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_name) + + discovery.assert_awaited_once_with(server) + assert relay.await_args.kwargs["mcp_server"].registration_url == "https://idp.example.com/register" + assert response is expected + + @pytest.fixture def trust_xff(): """Force ``IPAddressUtils.is_request_from_trusted_proxy`` to True. @@ -8522,6 +8644,8 @@ async def test_authorize_wall_names_the_issuer_for_anchored_servers(): assert "verify the Issuer" in detail_text assert "Servers with no url" not in detail_text assert "idp.example.com" not in detail_text + + def test_passthrough_authorization_code_round_trips_and_rejects_hostile_input(): """The passthrough gateway code seals and recovers the ephemeral DCR client and upstream code, and is total over hostile input: a raw upstream code opens to None, and a tampered or @@ -8930,7 +9054,9 @@ async def test_mint_ephemeral_dcr_client_unusable_registration_response_is_502(p ) from litellm.types.mcp import MCPAuth - server = _bridge_server(auth_type=MCPAuth.true_passthrough, dcr_bridge=None, server_id=server_id, server_name=server_id) + server = _bridge_server( + auth_type=MCPAuth.true_passthrough, dcr_bridge=None, server_id=server_id, server_name=server_id + ) mock_response = MagicMock() mock_response.text = json.dumps(payload) mock_response.raise_for_status = MagicMock() @@ -9008,8 +9134,6 @@ async def test_token_exchange_authenticates_with_the_sealed_clients_own_auth_met assert sent_body["client_secret"] == "mint-secret" - - # --------------------------------------------------------------------------- # LIT-4339: RFC 8707 resource indicators on the upstream OAuth legs # --------------------------------------------------------------------------- @@ -9262,7 +9386,9 @@ def test_upstream_resource_auto_keeps_the_path_because_it_identifies_the_server( sets ``upstream_resource`` explicitly instead of using ``auto``.""" from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource - first = resolve_upstream_resource(_resource_server(url="https://gw.example.com/team-a/mcp", upstream_resource="auto")) + first = resolve_upstream_resource( + _resource_server(url="https://gw.example.com/team-a/mcp", upstream_resource="auto") + ) second = resolve_upstream_resource( _resource_server(url="https://gw.example.com/team-b/mcp", upstream_resource="auto") ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 7df83065865..3392203dbab 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -21,7 +21,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.types.mcp import MCPAuth -from litellm.types.mcp_server.mcp_server_manager import MCPServer +from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer def _rendered_log_message(call): @@ -43,13 +43,19 @@ def cleanup_mcp_global_state(): global_mcp_server_manager, ) - # Clear before test + for slot in global_mcp_server_manager._oauth_discovery_slots: + if slot.task is not None and not slot.task.done(): + slot.task.cancel() global_mcp_server_manager.registry.clear() global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.clear() + global_mcp_server_manager._oauth_discovery_slots = () yield - # Clear after test + for slot in global_mcp_server_manager._oauth_discovery_slots: + if slot.task is not None and not slot.task.done(): + slot.task.cancel() global_mcp_server_manager.registry.clear() global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.clear() + global_mcp_server_manager._oauth_discovery_slots = () except ImportError: # MCP not available, skip cleanup yield @@ -1308,9 +1314,7 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing(): assert result.outcomes["failing2"].tag == "internal" # Verify failure logging for both servers - rendered_exceptions = [ - _rendered_log_message(c) for c in mock_logger.exception.call_args_list if c.args - ] + rendered_exceptions = [_rendered_log_message(c) for c in mock_logger.exception.call_args_list if c.args] assert ( "Error getting tools from server failing_server1: Server failing_server1 connection failed" in rendered_exceptions @@ -5733,13 +5737,17 @@ async def test_delegate_bad_token_gets_connect_time_401(): server = _delegate_auth_mcp_server() scope = _delegate_scope([(b"authorization", b"Bearer bogus-token")]) - with _patch_delegate_resolver(server, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[server]), - ), patch( - "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", - new=AsyncMock(return_value=(401, 'Bearer realm="upstream", error="invalid_token"')), - ) as probe: + with ( + _patch_delegate_resolver(server, "delegate_test"), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, 'Bearer realm="upstream", error="invalid_token"')), + ) as probe, + ): with pytest.raises(HTTPException) as exc_info: await _check_passthrough_upstream_auth( scope=scope, @@ -5751,7 +5759,9 @@ async def test_delegate_bad_token_gets_connect_time_401(): assert exc_info.value.status_code == 401 challenge = exc_info.value.headers["www-authenticate"] assert 'error="invalid_token"' in challenge - assert 'resource_metadata="http://localhost:4000/.well-known/oauth-protected-resource/mcp/delegate_test"' in challenge + assert ( + 'resource_metadata="http://localhost:4000/.well-known/oauth-protected-resource/mcp/delegate_test"' in challenge + ) probe.assert_awaited_once() probe_url, probe_auth = probe.call_args.args assert probe_url == "http://upstream:9401/mcp" @@ -5769,13 +5779,17 @@ async def test_delegate_valid_token_passes_preflight(): server = _delegate_auth_mcp_server() scope = _delegate_scope([(b"authorization", b"Bearer good-token")]) - with _patch_delegate_resolver(server, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[server]), - ), patch( - "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", - new=AsyncMock(return_value=(200, None)), - ) as probe: + with ( + _patch_delegate_resolver(server, "delegate_test"), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(200, None)), + ) as probe, + ): await _check_passthrough_upstream_auth( scope=scope, user_api_key_auth=UserAPIKeyAuth(), @@ -5799,12 +5813,16 @@ async def test_delegate_valid_token_forbidden_returns_403(): server = _delegate_auth_mcp_server() scope = _delegate_scope([(b"authorization", b"Bearer scoped-out-token")]) - with _patch_delegate_resolver(server, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[server]), - ), patch( - "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", - new=AsyncMock(return_value=(403, None)), + with ( + _patch_delegate_resolver(server, "delegate_test"), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(403, None)), + ), ): with pytest.raises(HTTPException) as exc_info: await _check_passthrough_upstream_auth( @@ -5830,13 +5848,17 @@ async def test_delegate_tokenless_request_not_probed(): server = _delegate_auth_mcp_server() scope = _delegate_scope([(b"content-type", b"application/json")]) - with _patch_delegate_resolver(server, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[server]), - ), patch( - "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", - new=AsyncMock(return_value=(401, None)), - ) as probe: + with ( + _patch_delegate_resolver(server, "delegate_test"), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, None)), + ) as probe, + ): await _check_passthrough_upstream_auth( scope=scope, user_api_key_auth=UserAPIKeyAuth(), @@ -5859,13 +5881,17 @@ async def test_delegate_preflight_skipped_on_multi_server_routes(): servers = [_delegate_auth_mcp_server("delegate-1"), _delegate_auth_mcp_server("delegate-2")] scope = _delegate_scope([(b"authorization", b"Bearer bogus-token")]) - with _patch_delegate_resolver(servers[0], "delegate_test", "other_server"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=servers), - ), patch( - "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", - new=AsyncMock(return_value=(401, None)), - ) as probe: + with ( + _patch_delegate_resolver(servers[0], "delegate_test", "other_server"), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=servers), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, None)), + ) as probe, + ): await _check_passthrough_upstream_auth( scope=scope, user_api_key_auth=UserAPIKeyAuth(), @@ -5898,13 +5924,17 @@ async def test_bare_authorization_never_probes_passthrough_servers(): ) scope = _delegate_scope([(b"authorization", b"Bearer ambiguous-token")]) - with _patch_delegate_resolver(passthrough_server, "pt_server"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[passthrough_server]), - ), patch( - "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", - new=AsyncMock(return_value=(401, None)), - ) as probe: + with ( + _patch_delegate_resolver(passthrough_server, "pt_server"), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[passthrough_server]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, None)), + ) as probe, + ): await _check_passthrough_upstream_auth( scope=scope, user_api_key_auth=UserAPIKeyAuth(), @@ -5940,13 +5970,17 @@ async def test_delegate_not_probed_when_named_only_via_server_id(): "headers": [(b"authorization", b"Bearer sk-litellm-proxy-key")], } - with _patch_delegate_resolver(server, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[server]), - ), patch( - "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", - new=AsyncMock(return_value=(401, None)), - ) as probe: + with ( + _patch_delegate_resolver(server, "delegate_test"), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, None)), + ) as probe, + ): await _check_passthrough_upstream_auth( scope=scope, user_api_key_auth=UserAPIKeyAuth(user_id="u1", api_key="hashed-sk"), @@ -5994,12 +6028,16 @@ async def test_delegate_preflight_with_unpatched_probe(): server = _delegate_auth_mcp_server() - with _patch_delegate_resolver(server, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[server]), - ), patch( - "litellm.proxy._experimental.mcp_server.server.get_async_httpx_client", - return_value=mock_client, + with ( + _patch_delegate_resolver(server, "delegate_test"), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.get_async_httpx_client", + return_value=mock_client, + ), ): with pytest.raises(HTTPException) as exc_info: await _check_passthrough_upstream_auth( @@ -6019,7 +6057,9 @@ async def test_delegate_preflight_with_unpatched_probe(): assert exc_info.value.status_code == 401 challenge = exc_info.value.headers["www-authenticate"] assert 'error="invalid_token"' in challenge - assert 'resource_metadata="http://localhost:4000/.well-known/oauth-protected-resource/mcp/delegate_test"' in challenge + assert ( + 'resource_metadata="http://localhost:4000/.well-known/oauth-protected-resource/mcp/delegate_test"' in challenge + ) probed_urls = [call.kwargs["url"] for call in mock_client.post.await_args_list] assert probed_urls == ["http://upstream:9401/mcp", "http://upstream:9401/mcp"] @@ -6044,12 +6084,16 @@ async def test_delegate_challenge_echoes_requested_alias(): "headers": [(b"authorization", b"Bearer bogus-token")], } - with _patch_delegate_resolver(server, "dt-alias"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[server]), - ), patch( - "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", - new=AsyncMock(return_value=(401, 'Bearer error="invalid_token"')), + with ( + _patch_delegate_resolver(server, "dt-alias"), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, 'Bearer error="invalid_token"')), + ), ): with pytest.raises(HTTPException) as exc_info: await _check_passthrough_upstream_auth( @@ -6076,13 +6120,17 @@ async def test_delegate_probe_not_fanned_out_to_access_group_members(): group_member = _delegate_auth_mcp_server() - with _patch_delegate_resolver(group_member, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[group_member]), - ), patch( - "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", - new=AsyncMock(return_value=(401, None)), - ) as probe: + with ( + _patch_delegate_resolver(group_member, "delegate_test"), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[group_member]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, None)), + ) as probe, + ): await _check_passthrough_upstream_auth( scope=_delegate_scope([(b"authorization", b"Bearer bogus-token")]), user_api_key_auth=UserAPIKeyAuth(), @@ -7990,6 +8038,54 @@ class TestPreemptive401ModeAware: client_ip=None, ) + @pytest.mark.asyncio + async def test_deferred_discovery_runs_before_delegate_challenge(self): + from litellm.proxy._experimental.mcp_server import server as server_module + + manager = server_module.global_mcp_server_manager + server = _make_oauth2_server( + "lazy_delegate", + oauth2_flow="authorization_code", + delegate_auth_to_upstream=True, + ) + manager.registry[server.server_id] = server + manager._set_oauth_discovery_deferred(server.server_id, True) + metadata = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + registration_url="https://idp.example.com/register", + ) + + with ( + patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)) as discovery, + pytest.raises(HTTPException) as exc, + ): + await self._run(server, None, has_stored_token=False) + + discovery.assert_awaited_once() + resolved = manager.registry[server.server_id] + assert resolved.authorization_url == "https://idp.example.com/authorize" + assert resolved.token_url == "https://idp.example.com/token" + assert resolved.registration_url == "https://idp.example.com/register" + assert manager._oauth_discovery_slot(server.server_id) is None + assert exc.value.status_code == 401 + + @pytest.mark.asyncio + async def test_stamped_m2m_challenge_skips_deferred_discovery(self): + from litellm.proxy._experimental.mcp_server import server as server_module + + manager = server_module.global_mcp_server_manager + server = _make_oauth2_server("stamped_m2m", oauth2_flow="client_credentials") + + with patch.object( + manager, + "ensure_oauth_metadata_discovered", + new=AsyncMock(side_effect=HTTPException(status_code=503, detail="discovery down")), + ) as discovery: + await self._run(server, None, has_stored_token=False) + + discovery.assert_not_awaited() + @pytest.mark.asyncio async def test_gateway_managed_interactive_no_token_challenges_with_x_litellm_api_key(self): """No stored token, key in x-litellm-api-key (oauth2_headers empty): 401.""" @@ -8318,16 +8414,13 @@ class TestListFiltersHonorThePrefixBoundary: url="http://127.0.0.1:5115/mcp", transport=MCPTransport.http, ) - published = MCPTool( - name=f"{self.SERVER_ID}-read_wiki_contents", description="", inputSchema={"type": "object"} - ) + published = MCPTool(name=f"{self.SERVER_ID}-read_wiki_contents", description="", inputSchema={"type": "object"}) auth = UserAPIKeyAuth(api_key="sk-test") - with patch.object( - MCPRequestHandler, "get_allowed_tools_for_server", AsyncMock(return_value=grants) - ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager" - ) as mock_manager: + with ( + patch.object(MCPRequestHandler, "get_allowed_tools_for_server", AsyncMock(return_value=grants)), + patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager") as mock_manager, + ): mock_manager.get_mcp_server_by_id.return_value = server listed = await filter_tools_by_key_team_permissions([published], self.SERVER_ID, auth) != [] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 99181f0f087..cd1ef7320dd 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -2,11 +2,10 @@ import importlib import asyncio import json import logging -import time import os import sys from datetime import datetime -from typing import Any, Dict, Optional +from typing import Any, Dict, Final, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -33,10 +32,12 @@ from mcp.types import ( ) from mcp.types import Tool as MCPTool +from litellm.constants import MCP_METADATA_TIMEOUT from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, _deserialize_json_dict, _flow_endpoints_missing, + _mcp_oauth_discovery_on_startup_enabled, _oauth_endpoints_unresolved, _deserialize_json_list, _normalize_mcp_server_cost_info, @@ -54,6 +55,7 @@ from litellm.proxy._types import ( MCPTransport, UserAPIKeyAuth, ) +from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.mcp import MCPAuth, MCPAuthType from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer @@ -72,6 +74,11 @@ def _reload_mcp_manager_module(): return reloaded +@pytest.fixture(autouse=True) +def enable_eager_mcp_oauth_discovery(monkeypatch): + monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1") + + class TestMCPServerManager: """Test MCP Server Manager stdio functionality""" @@ -428,6 +435,498 @@ class TestMCPServerManager: base.update(overrides) return {"m2mserver": base} + @pytest.mark.parametrize("value", ["1", "true", "TRUE", "yes", "on"]) + def test_mcp_oauth_discovery_on_startup_true_values(self, value): + with patch.dict(os.environ, {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": value}): + assert _mcp_oauth_discovery_on_startup_enabled() is True + + @pytest.mark.parametrize("value", ["0", "false", "FALSE", "no", "off", "", "invalid"]) + def test_mcp_oauth_discovery_on_startup_non_true_values(self, value): + with patch.dict(os.environ, {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": value}): + assert _mcp_oauth_discovery_on_startup_enabled() is False + + def test_mcp_oauth_discovery_on_startup_defaults_to_disabled(self): + with patch.dict(os.environ, {}, clear=True): + assert _mcp_oauth_discovery_on_startup_enabled() is False + + @pytest.mark.asyncio + async def test_config_oauth_discovery_warmup_is_non_blocking_and_shared(self): + metadata = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + scopes=["mcp.read"], + ) + with patch.dict(os.environ, {}, clear=True): + manager = MCPServerManager() + + started = asyncio.Event() + release = asyncio.Event() + + async def discover(_server): + started.set() + await release.wait() + return metadata + + with ( + patch.object(manager, "_discover_oauth_metadata_for_server", side_effect=discover) as discovery, + patch.object(manager, "initialize_tool_name_to_mcp_server_name_mapping"), + ): + load_task = asyncio.create_task( + manager.load_servers_from_config( + self._oauth2_config( + oauth2_flow="authorization_code", + authorization_url=None, + token_url=None, + ) + ) + ) + await started.wait() + assert load_task.done() + await load_task + + server = next(iter(manager.config_mcp_servers.values())) + waiters = [asyncio.create_task(manager.ensure_oauth_metadata_discovered(server)) for _ in range(10)] + await asyncio.sleep(0) + release.set() + resolved = await asyncio.gather(*waiters) + + discovery.assert_awaited_once_with(server) + assert all(result is resolved[0] for result in resolved) + assert resolved[0].authorization_url == "https://idp.example.com/authorize" + assert resolved[0].token_url == "https://idp.example.com/token" + assert resolved[0].scopes == ["mcp.read"] + assert manager.config_mcp_servers[server.server_id] is resolved[0] + assert server.authorization_url is None + assert manager._oauth_discovery_slot(server.server_id) is None + + @pytest.mark.asyncio + async def test_table_oauth_discovery_can_be_deferred_until_first_use(self): + row = LiteLLM_MCPServerTable( + server_id="lazy-db-1", + alias="lazy_db", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + created_at=datetime.now(), + updated_at=datetime.now(), + ) + metadata = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + ) + with patch.dict(os.environ, {}, clear=True): + manager = MCPServerManager() + + discovery = AsyncMock(return_value=metadata) + with patch.object(manager, "_descovery_metadata", new=discovery): + server = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) + + discovery.assert_not_awaited() + assert manager._oauth_discovery_slot(server.server_id) is not None + manager.registry[server.server_id] = server + + with patch.object(manager, "_descovery_metadata", new=discovery): + resolved = await manager.ensure_oauth_metadata_discovered(server) + + discovery.assert_awaited_once() + assert resolved.authorization_url == "https://idp.example.com/authorize" + assert resolved.token_url == "https://idp.example.com/token" + + @pytest.mark.asyncio + async def test_lazy_oauth_discovery_failure_is_shared_and_retries_after_cooldown(self): + with patch.dict(os.environ, {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": "false"}): + manager = MCPServerManager() + + metadata = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + ) + discovery = AsyncMock(side_effect=[None, None, None, metadata]) + discovery_clock: Final = MagicMock(return_value=100.0) + with ( + patch.object(manager, "_discover_oauth_metadata_for_server", new=discovery), + patch.object(manager, "initialize_tool_name_to_mcp_server_name_mapping"), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager._oauth_discovery_now", + new=discovery_clock, + ), + ): + await manager.load_servers_from_config( + self._oauth2_config( + oauth2_flow="authorization_code", + authorization_url=None, + token_url=None, + ) + ) + + server = next(iter(manager.config_mcp_servers.values())) + failures: Final = await asyncio.gather( + *(manager.ensure_oauth_metadata_discovered(server) for _ in range(10)), + return_exceptions=True, + ) + cooldown_failures: Final = await asyncio.gather( + *(manager.ensure_oauth_metadata_discovered(server) for _ in range(10)), + return_exceptions=True, + ) + discovery_clock.return_value = 130.0 + resolutions: Final = await asyncio.gather( + *(manager.ensure_oauth_metadata_discovered(server) for _ in range(10)) + ) + + assert discovery.await_count == 4 + assert all(isinstance(failure, HTTPException) and failure.status_code == 503 for failure in failures) + assert all(isinstance(failure, HTTPException) and failure.status_code == 503 for failure in cooldown_failures) + assert len({id(resolution) for resolution in resolutions}) == 1 + assert resolutions[0].authorization_url == "https://idp.example.com/authorize" + assert resolutions[0].token_url == "https://idp.example.com/token" + assert manager._oauth_discovery_slot(server.server_id) is None + + @pytest.mark.asyncio + async def test_lazy_oauth_discovery_timeout_is_bounded(self): + manager = MCPServerManager() + server = MCPServer( + server_id="lazy-timeout-1", + name="lazy_timeout", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + manager.registry[server.server_id] = server + manager._set_oauth_discovery_deferred(server.server_id, True) + + async def never_returns(_server): + await asyncio.Future() + + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCP_METADATA_TIMEOUT", + 0.01, + ), + patch.object(manager, "_discover_oauth_metadata_for_server", side_effect=never_returns) as discovery, + ): + with pytest.raises(HTTPException) as exc: + await asyncio.wait_for(manager.ensure_oauth_metadata_discovered(server), timeout=0.2) + + assert exc.value.status_code == 503 + assert "timed out" in str(exc.value.detail) + discovery.assert_awaited_once_with(server) + assert manager._oauth_discovery_slot(server.server_id) is not None + + @pytest.mark.asyncio + async def test_cancelling_one_waiter_does_not_cancel_shared_discovery(self): + manager = MCPServerManager() + server = MCPServer( + server_id="lazy-cancel-1", + name="lazy_cancel", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + manager.registry[server.server_id] = server + manager._set_oauth_discovery_deferred(server.server_id, True) + started = asyncio.Event() + release = asyncio.Event() + metadata = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + ) + + async def discover(_server): + started.set() + await release.wait() + return metadata + + with patch.object(manager, "_discover_oauth_metadata_for_server", side_effect=discover) as discovery: + cancelled_waiter = asyncio.create_task(manager.ensure_oauth_metadata_discovered(server)) + successful_waiter = asyncio.create_task(manager.ensure_oauth_metadata_discovered(server)) + await started.wait() + cancelled_waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await cancelled_waiter + release.set() + resolved = await successful_waiter + + discovery.assert_awaited_once_with(server) + assert resolved.authorization_url == "https://idp.example.com/authorize" + + @pytest.mark.asyncio + async def test_lazy_oauth_discovery_ignores_stale_registration_result(self): + manager = MCPServerManager() + old_server = MCPServer( + server_id="lazy-reload-1", + name="lazy_reload", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + replacement = old_server.model_copy(update={"url": "https://new.example.com/mcp"}) + manager.registry[old_server.server_id] = old_server + manager._set_oauth_discovery_deferred(old_server.server_id, True) + started = asyncio.Event() + metadata = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + ) + + async def discover(candidate): + if candidate.url == old_server.url: + started.set() + await asyncio.Future() + return metadata + + with patch.object(manager, "_discover_oauth_metadata_for_server", side_effect=discover): + old_attempt = asyncio.create_task(manager.ensure_oauth_metadata_discovered(old_server)) + await started.wait() + manager.registry[replacement.server_id] = replacement + manager._set_oauth_discovery_deferred(replacement.server_id, True) + resolved = await old_attempt + + assert resolved is manager.registry[replacement.server_id] + assert resolved.url == replacement.url + assert old_server.authorization_url is None + assert old_server.token_url is None + assert replacement.authorization_url is None + assert replacement.token_url is None + assert resolved.authorization_url == "https://idp.example.com/authorize" + assert resolved.token_url == "https://idp.example.com/token" + assert manager._oauth_discovery_slot(replacement.server_id) is None + + def test_registry_swap_reconcile_keeps_slot_for_issuer_anchored_server_without_url(self): + manager = MCPServerManager() + server = MCPServer( + server_id="anchored-no-url-1", + name="anchored_no_url", + url=None, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + issuer="https://idp.example.com", + issuer_is_anchored=True, + ) + manager.registry[server.server_id] = server + manager._set_oauth_discovery_deferred(server.server_id, True) + + manager._reconcile_oauth_discovery_slots_for_servers([server]) + + assert manager._oauth_discovery_slot(server.server_id) is not None + + resolved = server.model_copy( + update={ + "authorization_url": "https://idp.example.com/authorize", + "token_url": "https://idp.example.com/token", + } + ) + manager.registry[resolved.server_id] = resolved + manager._reconcile_oauth_discovery_slots_for_servers([resolved]) + + assert manager._oauth_discovery_slot(server.server_id) is None + + def _assert_oauth_discovery_state_removed(self, manager, server_id): + assert manager._oauth_discovery_slot(server_id) is None + + @pytest.mark.asyncio + async def test_deactivated_server_clears_lazy_oauth_discovery_state(self): + manager = MCPServerManager() + server = MCPServer( + server_id="lazy-deactivated-1", + name="lazy_deactivated", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + ) + manager.registry[server.server_id] = server + manager._set_oauth_discovery_deferred(server.server_id, True) + record = LiteLLM_MCPServerTable( + server_id=server.server_id, + server_name=server.name, + url=server.url, + transport=MCPTransport.http, + approval_status="rejected", + ) + + await manager.update_server(record) + + assert manager.registry == {} + self._assert_oauth_discovery_state_removed(manager, server.server_id) + + @pytest.mark.asyncio + async def test_database_reload_drop_clears_lazy_oauth_discovery_state(self): + manager = MCPServerManager() + server = MCPServer( + server_id="lazy-dropped-1", + name="lazy_dropped", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + ) + manager.registry[server.server_id] = server + manager._set_oauth_discovery_deferred(server.server_id, True) + repository = MagicMock() + repository.table.find_many = AsyncMock(return_value=[]) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + return_value=repository, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + ): + await manager.reload_servers_from_database() + + assert manager.registry == {} + self._assert_oauth_discovery_state_removed(manager, server.server_id) + + @pytest.mark.asyncio + async def test_database_reload_rearms_discovery_lost_to_registry_swap(self): + """A resolution published into the old registry while reload is staged + must leave the swapped-in unresolved entry with a fresh retry slot. + """ + manager = MCPServerManager() + stamp = datetime.now() + server = MCPServer( + server_id="lazy-swap-1", + name="lazy_swap", + server_name="lazy_swap", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + updated_at=stamp, + ) + manager.registry[server.server_id] = server + previous_registry = manager.registry + manager._set_oauth_discovery_deferred(server.server_id, True) + old_generation = manager._oauth_discovery_slot(server.server_id).generation + resolved = server.model_copy( + update={ + "authorization_url": "https://idp.example.com/authorize", + "token_url": "https://idp.example.com/token", + } + ) + row = LiteLLM_MCPServerTable( + server_id=server.server_id, + server_name=server.server_name, + url=server.url, + transport=server.transport, + auth_type=server.auth_type, + oauth2_flow=server.oauth2_flow, + updated_at=stamp, + ) + raw_row = MagicMock() + raw_row.model_dump.return_value = row.model_dump() + repository = MagicMock() + repository.table.find_many = AsyncMock(return_value=[raw_row]) + + async def publish_while_staged(*_args, **_kwargs): + assert manager.registry is previous_registry + assert manager._publish_resolved_oauth_server(resolved, old_generation) is resolved + + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + return_value=repository, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch.object( + manager, + "_maybe_register_openapi_tools", + new=AsyncMock(side_effect=publish_while_staged), + ), + patch.object(manager, "_prime_oauth_metadata_discovery_for_servers"), + ): + await manager.reload_servers_from_database() + + assert previous_registry[server.server_id] is resolved + assert manager.registry[server.server_id] is server + retry_slot = manager._oauth_discovery_slot(server.server_id) + assert retry_slot is not None + assert retry_slot.generation > old_generation + + @pytest.mark.asyncio + async def test_lazy_oauth_discovery_preserves_manual_authorization_url_gate(self): + with patch.dict(os.environ, {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": "false"}): + manager = MCPServerManager() + + metadata = MCPOAuthMetadata( + authorization_url="https://attacker.example.com/authorize", + token_url="https://attacker.example.com/token", + scopes=["mcp.read"], + ) + discovery = AsyncMock(return_value=metadata) + with ( + patch.object(manager, "_descovery_metadata", new=discovery), + patch.object(manager, "initialize_tool_name_to_mcp_server_name_mapping"), + ): + await manager.load_servers_from_config( + self._oauth2_config( + oauth2_flow="authorization_code", + authorization_url="https://idp.example.com/authorize", + token_url=None, + ) + ) + + server = next(iter(manager.config_mcp_servers.values())) + with ( + patch.object(manager, "_descovery_metadata", new=discovery), + pytest.raises(HTTPException) as exc, + ): + await manager.ensure_oauth_metadata_discovered(server) + + assert exc.value.status_code == 503 + assert manager.config_mcp_servers[server.server_id].authorization_url == "https://idp.example.com/authorize" + assert manager.config_mcp_servers[server.server_id].token_url is None + assert manager.config_mcp_servers[server.server_id].scopes is None + assert manager._oauth_discovery_slot(server.server_id) is not None + + @pytest.mark.asyncio + async def test_create_mcp_client_triggers_deferred_oauth_discovery(self): + manager = MCPServerManager() + server = MCPServer( + server_id="lazy-client-1", + name="lazy_client", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + ) + ensure_oauth_metadata_discovered: Final = AsyncMock(return_value=server) + + with ( + patch.object( + manager, + "ensure_oauth_metadata_discovered", + new=ensure_oauth_metadata_discovered, + ), + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient"), + ): + await manager._create_mcp_client(server) + + ensure_oauth_metadata_discovered.assert_awaited_once_with(server) + + @pytest.mark.asyncio + async def test_startup_tool_mapping_skips_servers_with_deferred_discovery(self): + manager = MCPServerManager() + server = MCPServer( + server_id="lazy-map-1", + name="lazy_map", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, + ) + manager.registry[server.server_id] = server + manager._set_oauth_discovery_deferred(server.server_id, True) + + with patch.object(manager, "_get_tools_from_server", new=AsyncMock()) as get_tools: + await manager._initialize_tool_name_to_mcp_server_name_mapping() + + get_tools.assert_not_awaited() + @pytest.mark.asyncio async def test_load_servers_from_config_requires_oauth2_flow(self): """auth_type oauth2 without an explicit oauth2_flow is a config error: the @@ -1521,7 +2020,9 @@ class TestMCPServerManager: ) resource_rooted = AsyncMock(return_value=MCPOAuthMetadata(token_url="https://attacker.example.com/steal")) with ( - patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved)) as anchored, + patch.object( + manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved) + ) as anchored, patch.object(manager, "_descovery_metadata", new=resource_rooted), ): built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) @@ -1555,7 +2056,9 @@ class TestMCPServerManager: patch.object( manager, "_fetch_single_authorization_server_metadata", new=AsyncMock(return_value=issuer_document) ) as issuer_fetch, - patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=resource_document)) as resource_fetch, + patch.object( + manager, "_descovery_metadata", new=AsyncMock(return_value=resource_document) + ) as resource_fetch, ): result = await manager._fetch_issuer_anchored_oauth_metadata( "https://idp.example.com", "https://up.example.com/mcp" @@ -1982,6 +2485,29 @@ class TestMCPServerManager: await manager.preflight_token_exchange(server=server, oauth2_headers=None, user_api_key_auth=None) assert resolved == ["good-subject"] + @pytest.mark.asyncio + async def test_preflight_token_exchange_skips_discovery_for_other_auth_modes(self): + """Preflight must not make unrelated auth modes depend on OAuth discovery.""" + manager = MCPServerManager() + server = MCPServer( + server_id="plain-preflight", + name="plain_preflight", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + manager.ensure_oauth_metadata_discovered = AsyncMock( + side_effect=AssertionError("non-token-exchange server was resolved") + ) + + await manager.preflight_token_exchange( + server=server, + oauth2_headers={"Authorization": "Bearer subject"}, + user_api_key_auth=None, + ) + + manager.ensure_oauth_metadata_discovered.assert_not_awaited() + @pytest.mark.asyncio async def test_call_regular_mcp_tool_passthrough_strips_authorization_when_admission_consumed_litellm_key( self, @@ -2882,7 +3408,7 @@ class TestMCPServerManager: patch( "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", return_value=mock_client, - ), + ) as get_client, patch.object( manager, "_attempt_well_known_discovery", @@ -2901,6 +3427,10 @@ class TestMCPServerManager: ): result = await manager._descovery_metadata("http://localhost:8001/mcp") + get_client.assert_called_once_with( + llm_provider=httpxSpecialProvider.MCP, + params={"timeout": MCP_METADATA_TIMEOUT}, + ) mock_well_known.assert_awaited_once_with("http://localhost:8001/mcp") mock_fetch_auth.assert_awaited_once_with( ["https://login.microsoftonline.com/test-tenant-id/v2.0"], @@ -3191,7 +3721,9 @@ class TestMCPServerManager: registration_url="https://discovered.example.com/register", ) - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): + async def fake_discovery( + server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False + ): assert server_url == "https://example.com/mcp" # oauth2 (browser flow) keeps the origin fallback; only OBO disables it. assert allow_origin_fallback is True @@ -3441,6 +3973,29 @@ class TestMCPServerManager: assert result.health_check_error == "Connection timeout" assert result.last_health_check is not None + @pytest.mark.asyncio + async def test_health_check_server_contains_client_creation_failure(self): + """Deferred discovery failures are reported unhealthy, not raised.""" + manager = MCPServerManager() + server = MCPServer( + server_id="discovery-failure", + name="discovery-failure", + transport=MCPTransport.http, + auth_type=None, + authentication_token="test-token", + url="https://up.example.com/mcp", + ) + manager.get_mcp_server_by_id = MagicMock(return_value=server) + manager._resolve_static_headers_with_env_vars = AsyncMock(return_value=None) + manager._create_mcp_client = AsyncMock( + side_effect=HTTPException(status_code=503, detail="OAuth discovery unavailable") + ) + + result = await manager.health_check_server(server.server_id) + + assert result.status == "unhealthy" + assert "OAuth discovery unavailable" in (result.health_check_error or "") + @pytest.mark.asyncio async def test_health_check_server_not_found(self): """Test health check for a server that doesn't exist""" @@ -4243,6 +4798,20 @@ class TestMCPServerManager: with pytest.raises(ValueError, match="Tool .* not found"): manager._resolve_mcp_server_for_tool_call("nonexistent", "ghost_tool") + def test_resolve_mcp_server_for_tool_call_unscoped_cached_tool_still_fails(self): + """Without an explicit server, an unmapped tool remains ambiguous.""" + manager = MCPServerManager() + manager.registry = { + "github": MCPServer( + server_id="github", + name="github", + transport=MCPTransport.http, + ) + } + + with pytest.raises(ValueError, match="Tool cached_tool not found"): + manager._resolve_mcp_server_for_tool_call("", "cached_tool") + def test_resolve_mcp_server_for_tool_call_unknown_tool_with_known_server(self): """Server-name match alone must not let unknown tools slip through. @@ -5586,7 +6155,9 @@ class TestMCPServerTimestamps: manager = MCPServerManager() calls: list[bool] = [] - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): + async def fake_discovery( + server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False + ): calls.append(allow_origin_fallback) return MCPOAuthMetadata( scopes=None, @@ -5621,7 +6192,9 @@ class TestMCPServerTimestamps: manager = MCPServerManager() calls: list[str] = [] - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): + async def fake_discovery( + server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False + ): calls.append(server_url) raise AssertionError("discovery must not run when token_exchange_endpoint is configured") @@ -5654,7 +6227,9 @@ class TestMCPServerTimestamps: lives on the in-memory registry entry only, for oauth2 and OBO alike.""" manager = MCPServerManager() - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): + async def fake_discovery( + server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False + ): return MCPOAuthMetadata( scopes=["mcp.read"], authorization_url="https://idp.example.com/authorize", @@ -5736,9 +6311,7 @@ class TestMCPServerTimestamps: assert _flow_endpoints_missing(MCPAuth.oauth2, "client_credentials", None, None) is True assert _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, None) is True assert _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, "https://idp/token") is False - assert ( - _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, None, "https://idp/exchange") is False - ) + assert _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, None, "https://idp/exchange") is False assert _flow_endpoints_missing(MCPAuth.api_key, None, None, None) is False def test_unresolved_check_uses_the_flow_judge_not_the_raw_column(self): @@ -5783,7 +6356,9 @@ class TestMCPServerTimestamps: registration_url=None, ) assert _oauth_endpoints_unresolved(relay_arm) is True - assert _oauth_endpoints_unresolved(relay_arm.model_copy(update={"registration_url": "https://idp/reg"})) is False + assert ( + _oauth_endpoints_unresolved(relay_arm.model_copy(update={"registration_url": "https://idp/reg"})) is False + ) assert _oauth_endpoints_unresolved(relay_arm.model_copy(update={"client_id": "admin-client"})) is False def test_entra_obo_without_scopes_is_unresolved(self): @@ -5804,50 +6379,6 @@ class TestMCPServerTimestamps: assert _oauth_endpoints_unresolved(entra.model_copy(update={"scopes": ["api://app/.default"]})) is False assert _oauth_endpoints_unresolved(entra.model_copy(update={"token_exchange_profile": "rfc8693"})) is False - def test_oauth_discovery_retry_backs_off_per_server(self): - """Without a cooldown the fast-path exemption re-runs the full discovery chain, and re-emits - the unresolved warning, on every reload forever for a server that can never resolve. Delay - doubles per consecutive failure up to the cap, a success clears the state so the next failure - starts from the base delay again, and the cooldown is per server.""" - manager = MCPServerManager() - - def unresolved(server_id): - return MCPServer( - server_id=server_id, - name=server_id, - url="https://up.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - oauth2_flow="authorization_code", - ) - - assert manager._oauth_discovery_retry_due("a") is True - - manager._record_oauth_discovery_outcome(unresolved("a")) - assert manager._oauth_discovery_retry_due("a") is False - assert manager._oauth_discovery_retry_due("b") is True, "cooldown must be per server" - - failures_before, _ = manager._oauth_discovery_retry_state["a"] - manager._record_oauth_discovery_outcome(unresolved("a")) - failures_after, _ = manager._oauth_discovery_retry_state["a"] - assert failures_after == failures_before + 1 - - # An elapsed cooldown lets the retry through, and the delay grows with the failure count - manager._oauth_discovery_retry_state["a"] = (1, time.monotonic() - 31.0) - assert manager._oauth_discovery_retry_due("a") is True - manager._oauth_discovery_retry_state["a"] = (5, time.monotonic() - 31.0) - assert manager._oauth_discovery_retry_due("a") is False - - resolved = unresolved("a").model_copy( - update={ - "authorization_url": "https://idp.example.com/authorize", - "token_url": "https://idp.example.com/token", - } - ) - manager._record_oauth_discovery_outcome(resolved) - assert "a" not in manager._oauth_discovery_retry_state - assert manager._oauth_discovery_retry_due("a") is True - @pytest.mark.asyncio async def test_reload_fast_path_retries_unresolved_oauth_servers(self): """A server whose discovery failed must not be pinned broken by the updated_at fast path: @@ -8640,7 +9171,9 @@ class TestOBOEndpointDiscovery: ) seen = [] - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): + async def fake_discovery( + server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False + ): seen.append((server_url, allow_origin_fallback)) return discovered @@ -8668,7 +9201,9 @@ class TestOBOEndpointDiscovery: async def test_config_obo_with_configured_endpoint_skips_discovery(self): manager = MCPServerManager() - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): + async def fake_discovery( + server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False + ): raise AssertionError("discovery must not run when the endpoint is configured") manager._descovery_metadata = fake_discovery # type: ignore[attr-defined] @@ -9073,7 +9608,9 @@ class TestUrllessIssuerDiscovery: ) resource_rooted = AsyncMock(return_value=None) with ( - patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved)) as anchored, + patch.object( + manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved) + ) as anchored, patch.object(manager, "_descovery_metadata", new=resource_rooted), ): built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) @@ -9121,7 +9658,9 @@ class TestUrllessIssuerDiscovery: resolved = MCPOAuthMetadata(token_url="https://idp.example.com/token") resource_rooted = AsyncMock(return_value=None) with ( - patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved)) as anchored, + patch.object( + manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved) + ) as anchored, patch.object(manager, "_descovery_metadata", new=resource_rooted), ): built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) @@ -9139,9 +9678,7 @@ class TestDiscoveryFailureLogging: def _connect_error_client(self, url: str) -> MagicMock: client = MagicMock() - client.get = AsyncMock( - side_effect=httpx.ConnectError(f"[Errno 8] nodename nor servname provided for {url}") - ) + client.get = AsyncMock(side_effect=httpx.ConnectError(f"[Errno 8] nodename nor servname provided for {url}")) return client @pytest.mark.asyncio @@ -9184,9 +9721,7 @@ class TestDiscoveryFailureLogging: manager = MCPServerManager() url = "https://real-host.example.com/mcp-typo" client = MagicMock() - client.get = AsyncMock( - return_value=httpx.Response(404, request=httpx.Request("GET", url)) - ) + client.get = AsyncMock(return_value=httpx.Response(404, request=httpx.Request("GET", url))) with ( patch( "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",