mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #36599 from daniel-meismer-zocdoc/bugfix/mcp-oauth-startup-discovery
fix(mcp): oauth discovery must not cause outages
This commit is contained in:
commit
4fd7a73ef5
6 changed files with 1380 additions and 249 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 "<unparseable url>"
|
||||
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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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) != []
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue