mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
fix(mcp): explain refused OAuth registration and bound discovery retries
This commit is contained in:
parent
acb9086f29
commit
f04fb748c5
10 changed files with 228 additions and 8 deletions
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 _:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue