From f04fb748c532edea6606e1ee9f7d110a9d82b4b2 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Thu, 10 Sep 2026 21:02:46 -0700 Subject: [PATCH 1/3] fix(mcp): explain refused OAuth registration and bound discovery retries --- README.md | 2 + .../mcp_server/faults/__init__.py | 2 + .../mcp_server/faults/classify.py | 5 +- .../mcp_server/faults/render_oauth.py | 21 ++++- .../_experimental/mcp_server/faults/types.py | 10 ++- .../mcp_server/mcp_server_manager.py | 22 ++++- .../mcp_server/faults/test_classify.py | 29 +++++++ .../mcp_server/faults/test_render_oauth.py | 18 ++++ .../mcp_server/test_discoverable_endpoints.py | 45 ++++++++++ .../mcp_server/test_mcp_server_manager.py | 82 +++++++++++++++++++ 10 files changed, 228 insertions(+), 8 deletions(-) diff --git a/README.md b/README.md index 92757fcbbc1..901cc5b0cea 100644 --- a/README.md +++ b/README.md @@ -262,6 +262,8 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \ } ``` +For MCP OAuth, an upstream may advertise dynamic client registration but refuse requests with HTTP 401 or 403. If the provider requires a pre-registered OAuth app, configure its `credentials.client_id` and, when required, `credentials.client_secret` on the MCP server. This skips dynamic registration in the gateway sign-in flow. The provider must approve the app for MCP access; reaching its authorization page does not establish that login or tool calls will succeed + [**Docs: MCP Gateway**](https://docs.litellm.ai/docs/mcp) diff --git a/litellm/proxy/_experimental/mcp_server/faults/__init__.py b/litellm/proxy/_experimental/mcp_server/faults/__init__.py index 1b9ee77d795..de7ee5c866a 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/__init__.py +++ b/litellm/proxy/_experimental/mcp_server/faults/__init__.py @@ -22,6 +22,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import ( GatewayRejected, UpstreamOAuthFault, UpstreamProtocolFault, + UpstreamRegistrationRefused, UpstreamReportedFault, ) @@ -31,6 +32,7 @@ __all__ = [ "GatewayRejected", "UpstreamOAuthFault", "UpstreamProtocolFault", + "UpstreamRegistrationRefused", "UpstreamReportedFault", "classify_upstream_dcr_rejection", "classify_upstream_token_rejection", diff --git a/litellm/proxy/_experimental/mcp_server/faults/classify.py b/litellm/proxy/_experimental/mcp_server/faults/classify.py index 2162d078c09..bb41436b495 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/classify.py +++ b/litellm/proxy/_experimental/mcp_server/faults/classify.py @@ -21,6 +21,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import ( GatewayRejected, UpstreamOAuthFault, UpstreamProtocolFault, + UpstreamRegistrationRefused, UpstreamReportedFault, ) @@ -122,11 +123,13 @@ def classify_upstream_dcr_rejection(response: httpx.Response, log_context: str) """Classify a dynamic-client-registration rejection. RFC 7591 §3.2.2 errors carry ``error`` / ``error_description`` and go through the same blame assignment as token errors (registration sends no client credentials, so credential codes stay caller-actionable); anything - without a usable ``error`` field is an upstream protocol fault.""" + without a usable ``error`` field is a registration refusal for 401/403 and a protocol fault otherwise.""" parsed: Final = _safe_json(response) fields: Final = parsed if isinstance(parsed, dict) else {} code: Final = _bounded_field(fields.get("error")) if code is None: + if response.status_code == 401 or response.status_code == 403: + return UpstreamRegistrationRefused(status_code=response.status_code) _log_out_of_contract("registration", response, log_context) return UpstreamProtocolFault(note=f"upstream registration failed with HTTP {response.status_code}") return _classify_oauth_error_code( diff --git a/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py b/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py index 64d14140a5b..b3464382142 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py +++ b/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py @@ -11,7 +11,7 @@ from typing import Final from fastapi.responses import JSONResponse from typing_extensions import assert_never -from litellm.proxy._experimental.mcp_server.faults.types import UpstreamOAuthFault +from litellm.proxy._experimental.mcp_server.faults.types import CallerRejected, UpstreamOAuthFault from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS @@ -35,6 +35,14 @@ def _upstream_reported_status_and_description(code: str) -> tuple[int, str]: return 502, "the upstream authorization server reported an internal error" +def _registration_refused_description(status_code: int) -> str: + return ( + f"the upstream authorization server refused dynamic client registration (HTTP {status_code}). " + "This provider may require a pre-registered OAuth client. Configure client_id and, if required " + "by the provider, client_secret for this MCP server to skip dynamic registration" + ) + + def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse: """RFC 6749 §5.2 response for a token-endpoint fault. Caller-actionable rejections relay the upstream's code on the status that code implies (401 for invalid_client per §5.2, else 400); @@ -65,6 +73,13 @@ def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse: content={"error": fault.code, "error_description": description}, headers=TOKEN_NO_CACHE_HEADERS, ) + case "upstream_registration_refused": + return render_token_fault( + CallerRejected( + code="unauthorized_client", + description=_registration_refused_description(fault.status_code), + ) + ) case "upstream_protocol_fault": return JSONResponse( status_code=502, @@ -78,7 +93,7 @@ def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse: def dcr_fault_detail(fault: UpstreamOAuthFault) -> tuple[int, str]: """Status and detail string for a registration fault, raised as HTTPException by the caller. RFC 7591 §3.2.2 defines registration errors as 400, so a contract-conformant rejection is 400 - regardless of the status the upstream chose; everything else is a 502 upstream fault.""" + regardless of the upstream status; a bare 401/403 is a registration refusal rendered as 403.""" match fault.tag: case "caller_rejected": detail: Final = f"{fault.code}: {fault.description}" if fault.description else fault.code @@ -87,6 +102,8 @@ def dcr_fault_detail(fault: UpstreamOAuthFault) -> tuple[int, str]: return 502, _gateway_rejected_description(fault.code) case "upstream_reported_fault": return _upstream_reported_status_and_description(fault.code) + case "upstream_registration_refused": + return 403, _registration_refused_description(fault.status_code) case "upstream_protocol_fault": return 502, fault.note case _: diff --git a/litellm/proxy/_experimental/mcp_server/faults/types.py b/litellm/proxy/_experimental/mcp_server/faults/types.py index 4b9505ad801..d081d9d735e 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/types.py +++ b/litellm/proxy/_experimental/mcp_server/faults/types.py @@ -77,4 +77,12 @@ class UpstreamProtocolFault(BaseModel): note: str -UpstreamOAuthFault: TypeAlias = CallerRejected | GatewayRejected | UpstreamReportedFault | UpstreamProtocolFault +class UpstreamRegistrationRefused(BaseModel): + model_config = ConfigDict(frozen=True) + tag: Literal["upstream_registration_refused"] = "upstream_registration_refused" + status_code: Literal[401, 403] + + +UpstreamOAuthFault: TypeAlias = ( + CallerRejected | GatewayRejected | UpstreamReportedFault | UpstreamProtocolFault | UpstreamRegistrationRefused +) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index af25b0e919a..d5f4fa2b256 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1910,7 +1910,7 @@ class MCPServerManager: elif server.server_id in self.config_mcp_servers: self.config_mcp_servers[server.server_id] = server else: - return None + return server self._remove_oauth_discovery_slot(server.server_id) return server @@ -2002,6 +2002,12 @@ class MCPServerManager: if slot.task is not None: if not slot.task.done() or _oauth_discovery_now() < slot.retry_not_before: return slot.task, slot.generation + if ( + not slot.task.cancelled() + and slot.task.exception() is None + and isinstance(slot.task.result(), _OAuthDiscoveryResolved) + ): + return slot.task, slot.generation task: Final = asyncio.create_task( self._run_oauth_metadata_resolution(self._registered_server(server), slot.generation) ) @@ -2031,7 +2037,7 @@ class MCPServerManager: 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: + async def ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer: """Join the bounded discovery task and return the resolved server. Concurrent callers share one task per server. A failed attempt remains @@ -2058,13 +2064,13 @@ class MCPServerManager: 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) + return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale) raise match outcome: case _OAuthDiscoveryResolved(resolved_server): return resolved_server case _OAuthDiscoveryStale(): - return await self.ensure_oauth_metadata_discovered(server) + return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale) case _OAuthDiscoveryFailed(timed_out=timed_out): current: Final = self._registered_server(server) if current.is_client_forwarded_token: @@ -2076,6 +2082,14 @@ class MCPServerManager: detail=f"OAuth metadata discovery {reason} for MCP server {server_ref!r}", ) + async def _rejoin_oauth_metadata_discovery(self, server: MCPServer, *, retry_stale: bool) -> MCPServer: + if retry_stale: + return await self.ensure_oauth_metadata_discovered(server, _retry_stale=False) + current: Final = self._registered_server(server) + if not _oauth_endpoints_unresolved(current) or current.is_client_forwarded_token: + return current + raise HTTPException(status_code=503, detail="OAuth metadata discovery changed repeatedly; retry shortly") + def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None: raw: Final[str | None] = getattr(client, "_last_initialize_instructions", None) if raw and str(raw).strip(): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_classify.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_classify.py index dc20d664a53..30107db4055 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_classify.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_classify.py @@ -1,6 +1,9 @@ """Classification matrix for upstream OAuth/DCR rejections: who is blamed depends only on the §5.2 code and whose credentials the gateway presented, never on the upstream's HTTP status.""" +from typing import Final + +import pytest import httpx from litellm.proxy._experimental.mcp_server.faults.classify import ( @@ -12,6 +15,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import ( GatewayRejected, UpstreamProtocolFault, UpstreamReportedFault, + UpstreamRegistrationRefused, ) @@ -145,3 +149,28 @@ def test_dcr_server_error_code_is_not_blamed_on_caller(): log_context="srv", ) assert isinstance(fault, UpstreamReportedFault) + + +@pytest.mark.parametrize("status_code", [401, 403]) +@pytest.mark.parametrize("body", ["Forbidden", 'private upstream details', '{"error": ""}', '{"error": 12}']) +def test_dcr_access_refusal_without_oauth_error(status_code: int, body: str) -> None: + fault: Final = classify_upstream_dcr_rejection(_response(status_code, text_body=body), log_context="srv") + assert isinstance(fault, UpstreamRegistrationRefused) + assert fault.status_code == status_code + + +@pytest.mark.parametrize("status_code", [401, 403]) +def test_dcr_access_refusal_preserves_oauth_error(status_code: int) -> None: + fault: Final = classify_upstream_dcr_rejection( + _response(status_code, json_body={"error": "invalid_redirect_uri", "error_description": "not allowed"}), + log_context="srv", + ) + assert fault == CallerRejected(code="invalid_redirect_uri", description="not allowed") + + +@pytest.mark.parametrize("status_code", [401, 403]) +def test_token_access_refusal_remains_protocol_fault(status_code: int) -> None: + fault: Final = classify_upstream_token_rejection( + _response(status_code, text_body="Forbidden"), credential_source="gateway_stored", log_context="srv" + ) + assert isinstance(fault, UpstreamProtocolFault) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_render_oauth.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_render_oauth.py index 78513e315a7..a6807ae1454 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_render_oauth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_render_oauth.py @@ -2,6 +2,9 @@ code can never ship on a server-fault status and gateway-side faults never carry provider prose.""" import json +from typing import Final, Literal + +import pytest from litellm.proxy._experimental.mcp_server.faults.render_oauth import ( dcr_fault_detail, @@ -12,6 +15,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import ( GatewayRejected, UpstreamProtocolFault, UpstreamReportedFault, + UpstreamRegistrationRefused, ) @@ -94,3 +98,17 @@ def test_dcr_upstream_reported_fault_maps_to_5xx(): status_code, detail = dcr_fault_detail(UpstreamReportedFault(code="server_error")) assert status_code == 502 assert "internal error" in detail + + +@pytest.mark.parametrize("upstream_status", [401, 403]) +def test_registration_refusal_gives_configuration_guidance(upstream_status: Literal[401, 403]) -> None: + fault: Final = UpstreamRegistrationRefused(status_code=upstream_status) + status, detail = dcr_fault_detail(fault) + assert status == 403 + assert f"HTTP {upstream_status}" in detail + assert "may require a pre-registered OAuth client" in detail + assert "client_id" in detail and "client_secret" in detail + response: Final = render_token_fault(fault) + assert response.status_code == 400 + assert json.loads(response.body) == {"error": "unauthorized_client", "error_description": detail} + assert response.headers["cache-control"] == "no-store" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 5a99139a67f..18ff50a881a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -11141,3 +11141,48 @@ async def test_enforced_login_warms_verified_token_readable_without_database_loo assert token.identity_binding_proof == proof assert token.refresh_token is None read.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("upstream_status", [401, 403]) +@pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate]) +@pytest.mark.parametrize("dcr_bridge", [False, True]) +@pytest.mark.parametrize("flow", ["register", "mint"]) +async def test_dcr_refusal_is_actionable_without_upstream_body( + upstream_status: int, auth_type: MCPAuth, dcr_bridge: bool, flow: str, monkeypatch: pytest.MonkeyPatch +) -> None: + import httpx + from typing import Final + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + mint_ephemeral_dcr_client, + register_client_with_server, + ) + + server: Final = _bridge_server( + auth_type=auth_type, dcr_bridge=dcr_bridge, server_id=f"refused-{auth_type}-{dcr_bridge}-{flow}-{upstream_status}", + client_id=None, + ) + import respx + + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + with respx.mock as upstream: + registration: Final = upstream.post(server.registration_url).mock( + return_value=httpx.Response(upstream_status, text="Forbidden private upstream details") + ) + operation: Final = ( + mint_ephemeral_dcr_client(_bridge_mock_request(), server) + if flow == "mint" + else register_client_with_server( + request=_bridge_mock_request(), mcp_server=server, client_name="Test client", + grant_types=None, response_types=None, token_endpoint_auth_method=None, + client_redirect_uris=["http://localhost:9999/callback"], + ) + ) + with pytest.raises(HTTPException) as exc: + await operation + assert registration.call_count == 1 + assert exc.value.status_code == 403 + assert f"HTTP {upstream_status}" in str(exc.value.detail) + assert "pre-registered OAuth client" in str(exc.value.detail) + assert "private upstream details" not in str(exc.value.detail) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index adf985e9a21..ecf78359191 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -12676,3 +12676,85 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(transport: Li assert "Credential=AKIDEXAMPLE/" in request.headers["Authorization"] finally: request_ctx.reset(token) + + +@pytest.mark.asyncio +async def test_temporary_server_discovery_reuses_resolved_metadata_without_publishing() -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="temporary-oauth-discovery", name="temporary", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + ) + manager._set_oauth_discovery_deferred(server.server_id, True) + metadata: Final = 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, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery: + resolved: Final = await manager.ensure_oauth_metadata_discovered(server) + repeated: Final = await manager.ensure_oauth_metadata_discovered(server) + assert resolved.authorization_url == metadata.authorization_url + assert resolved.token_url == metadata.token_url + assert resolved.registration_url == metadata.registration_url + assert repeated is resolved + assert server.server_id not in manager.registry + assert server.server_id not in manager.config_mcp_servers + discovery.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("auth_type", [MCPAuth.oauth2, MCPAuth.true_passthrough]) +async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="repeated-stale", name="stale", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=auth_type, oauth2_flow="authorization_code", + ) + manager.registry[server.server_id] = server + manager._set_oauth_discovery_deferred(server.server_id, True) + metadata: Final = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + ) + with ( + patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery, + patch.object(manager, "_publish_resolved_oauth_server", return_value=None), + ): + if auth_type == MCPAuth.true_passthrough: + assert await manager.ensure_oauth_metadata_discovered(server) is server + else: + with pytest.raises(HTTPException) as exc: + await manager.ensure_oauth_metadata_discovered(server) + assert exc.value.status_code == 503 + assert "changed repeatedly" in str(exc.value.detail) + assert discovery.await_count == 2 + + +@pytest.mark.asyncio +async def test_stale_discovery_falls_back_to_resolved_registered_server() -> None: + manager: Final = MCPServerManager() + original: Final = MCPServer( + server_id="resolved-replacement", name="replacement", url="https://old.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + ) + replacement: Final = original.model_copy(update={ + "url": "https://new.example.com/mcp", "authorization_url": "https://new.example.com/authorize", + "token_url": "https://new.example.com/token", + }) + manager.registry[original.server_id] = replacement + assert await manager._rejoin_oauth_metadata_discovery(original, retry_stale=False) is replacement + + +def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: + manager: Final = MCPServerManager() + original: Final = MCPServer( + server_id="stale-publication", name="publication", url="https://old.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + ) + manager._set_oauth_discovery_deferred(original.server_id, True) + original_slot: Final = manager._oauth_discovery_slot(original.server_id) + assert original_slot is not None + replacement: Final = original.model_copy(update={"url": "https://new.example.com/mcp"}) + manager.registry[original.server_id] = replacement + manager._set_oauth_discovery_deferred(original.server_id, True) + assert manager._publish_resolved_oauth_server(original, original_slot.generation) is None + assert manager.registry[original.server_id] is replacement From bb9b4c4aef48a244bbe4ea40cc56a60f9ffbfbce Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Thu, 10 Sep 2026 21:54:06 -0700 Subject: [PATCH 2/3] fix(mcp): render registration refusals without recursion --- .../mcp_server/faults/render_oauth.py | 20 +++++++++++-------- 1 file changed, 12 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py b/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py index b3464382142..3ecf5310482 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py +++ b/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py @@ -43,6 +43,16 @@ def _registration_refused_description(status_code: int) -> str: ) +def _render_caller_rejected(fault: CallerRejected) -> JSONResponse: + content: Final = { + "error": fault.code, + **({"error_description": fault.description} if fault.description else {}), + **({"error_uri": fault.error_uri} if fault.error_uri else {}), + } + status_code: Final = 401 if fault.code == "invalid_client" else 400 + return JSONResponse(status_code=status_code, content=content, headers=TOKEN_NO_CACHE_HEADERS) + + def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse: """RFC 6749 §5.2 response for a token-endpoint fault. Caller-actionable rejections relay the upstream's code on the status that code implies (401 for invalid_client per §5.2, else 400); @@ -50,13 +60,7 @@ def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse: blamed for, or shown the internals of, a failure only the operator can fix.""" match fault.tag: case "caller_rejected": - content: Final = { - "error": fault.code, - **({"error_description": fault.description} if fault.description else {}), - **({"error_uri": fault.error_uri} if fault.error_uri else {}), - } - status_code = 401 if fault.code == "invalid_client" else 400 - return JSONResponse(status_code=status_code, content=content, headers=TOKEN_NO_CACHE_HEADERS) + return _render_caller_rejected(fault) case "gateway_rejected": return JSONResponse( status_code=502, @@ -74,7 +78,7 @@ def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse: headers=TOKEN_NO_CACHE_HEADERS, ) case "upstream_registration_refused": - return render_token_fault( + return _render_caller_rejected( CallerRejected( code="unauthorized_client", description=_registration_refused_description(fault.status_code), From 8d5a675878f22fb387c7e5a8bd642913408f67c8 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Fri, 11 Sep 2026 07:01:09 -0700 Subject: [PATCH 3/3] fix(mcp): expire temporary OAuth discovery results --- .../mcp_server/mcp_server_manager.py | 11 ++++++ .../mcp_server/test_mcp_server_manager.py | 36 +++++++++++++++++++ 2 files changed, 47 insertions(+) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d5f4fa2b256..5f9384c7eea 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -254,6 +254,7 @@ _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 +_OAUTH_TEMPORARY_DISCOVERY_TTL_SECONDS: Final = 300.0 def _oauth_discovery_now() -> float: @@ -1898,6 +1899,10 @@ class MCPServerManager: slot: Final = self._oauth_discovery_slot(server_id) return slot is not None and slot.generation == generation + def _expire_temporary_oauth_discovery(self, server_id: str, generation: int) -> None: + if self._oauth_discovery_slot_is_current(server_id, generation): + self._remove_oauth_discovery_slot(server_id) + def _publish_resolved_oauth_server( self, server: MCPServer, @@ -1910,6 +1915,12 @@ class MCPServerManager: elif server.server_id in self.config_mcp_servers: self.config_mcp_servers[server.server_id] = server else: + asyncio.get_running_loop().call_later( + _OAUTH_TEMPORARY_DISCOVERY_TTL_SECONDS, + self._expire_temporary_oauth_discovery, + server.server_id, + generation, + ) return server self._remove_oauth_discovery_slot(server.server_id) return server diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index ecf78359191..f0998cff583 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -12758,3 +12758,39 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: manager._set_oauth_discovery_deferred(original.server_id, True) assert manager._publish_resolved_oauth_server(original, original_slot.generation) is None assert manager.registry[original.server_id] is replacement + + +@pytest.mark.asyncio +async def test_temporary_oauth_discovery_expires_without_more_requests() -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + ) + manager._set_oauth_discovery_deferred(server.server_id, True) + resolved: Final = await manager.ensure_oauth_metadata_discovered(server) + assert manager._oauth_discovery_slot(server.server_id) is not None + loop: Final = asyncio.get_running_loop() + expired: Final = loop.create_future() + with patch.object(loop, "time", return_value=loop.time() + 301): + loop.call_later(0, expired.set_result, None) + await expired + assert resolved.authorization_url == server.authorization_url + assert manager._oauth_discovery_slot(server.server_id) is None + + +def test_old_temporary_discovery_expiry_preserves_replacement() -> None: + manager: Final = MCPServerManager() + manager._set_oauth_discovery_deferred("reused-session", True) + old_slot: Final = manager._oauth_discovery_slot("reused-session") + assert old_slot is not None + manager._set_oauth_discovery_deferred("reused-session", True) + replacement: Final = manager._oauth_discovery_slot("reused-session") + manager._expire_temporary_oauth_discovery("reused-session", old_slot.generation) + assert manager._oauth_discovery_slot("reused-session") is replacement + assert replacement is not None + manager._expire_temporary_oauth_discovery("reused-session", replacement.generation) + assert manager._oauth_discovery_slot("reused-session") is None + manager._expire_temporary_oauth_discovery("reused-session", replacement.generation) + assert manager._oauth_discovery_slot("reused-session") is None