Merge pull request #40679 from BerriAI/litellm_fix_mcp_oauth_registration_7498

fix(mcp): explain refused OAuth registration and bound discovery retries
This commit is contained in:
joshua-berri 2026-09-11 10:51:31 -07:00 committed by GitHub
commit 3f81ba3d30
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 287 additions and 15 deletions

View file

@ -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)
</details>

View file

@ -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",

View file

@ -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(

View file

@ -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,24 @@ 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_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);
@ -42,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,
@ -65,6 +77,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_caller_rejected(
CallerRejected(
code="unauthorized_client",
description=_registration_refused_description(fault.status_code),
)
)
case "upstream_protocol_fault":
return JSONResponse(
status_code=502,
@ -78,7 +97,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 +106,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 _:

View file

@ -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
)

View file

@ -255,6 +255,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:
@ -1947,6 +1948,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,
@ -1959,7 +1964,13 @@ class MCPServerManager:
elif server.server_id in self.config_mcp_servers:
self.config_mcp_servers[server.server_id] = server
else:
return None
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
@ -2051,6 +2062,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)
)
@ -2080,7 +2097,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
@ -2107,13 +2124,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:
@ -2125,6 +2142,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():

View file

@ -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", '<html>private upstream details</html>', '{"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)

View file

@ -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"

View file

@ -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)

View file

@ -12787,6 +12787,125 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(transport: Li
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
@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
@pytest.mark.asyncio
async def test_openapi_health_coalesces_concurrent_checks_and_reuses_results(respx_mock, monkeypatch):
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")