From a7024590f84d5d7ecc7e9b99e14bc79539ece36c Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 17:04:33 -0700 Subject: [PATCH] fix(mcp): preserve safe OAuth retries and recovery challenges --- .../mcp_server/auth/user_api_key_auth_mcp.py | 25 +++-- .../mcp_server/discoverable_endpoints.py | 52 ++++++----- .../mcp_server/gateway_dcr_flow.py | 4 +- tests/mcp_tests/test_mcp_server.py | 2 + .../auth/test_user_api_key_auth_mcp.py | 31 +++++-- .../mcp_server/test_discoverable_endpoints.py | 91 ++++++++++++++++++- 6 files changed, 159 insertions(+), 46 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 2656b40ffdb..71ea1810486 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -462,7 +462,8 @@ class MCPRequestHandler: request_route: Final = get_request_route(request) # Only OAuth metadata routes registered under /.well-known/ are public. - if request_route.startswith("/.well-known/"): + is_public_metadata: Final = request_route.startswith("/.well-known/") + if is_public_metadata: validated_user_api_key_auth = UserAPIKeyAuth() elif has_explicit_litellm_key: # An explicit x-litellm-api-key is always a LiteLLM credential, even @@ -552,7 +553,7 @@ class MCPRequestHandler: scope.pop(CONNECTION_SCOPE_KEY, None) connection_header: Final = headers.get("authorization") - if is_connection_credential(connection_header): + if is_connection_credential(connection_header) and not is_public_metadata: scope[CONNECTION_SCOPE_KEY] = await MCPRequestHandler._admit_connection_credential( request=request, request_route=request_route, @@ -627,10 +628,14 @@ class MCPRequestHandler: resource=f"{get_request_base_url(request)}/mcp", ) connection: Final = open_connection_credential(connection_header) - if connection is None: + if connection is None or connection.binding != expected_binding: raise HTTPException( status_code=401, - detail="Invalid or expired MCP connection credential", + detail=( + "Invalid or expired MCP connection credential" + if connection is None + else "Connection credential belongs to a different key or resource" + ), headers=MappingProxyType( { "www-authenticate": connection_challenge(request, expected_binding), @@ -638,8 +643,6 @@ class MCPRequestHandler: } ), ) - if connection.binding != expected_binding: - raise HTTPException(status_code=401, detail="Connection credential belongs to a different key or resource") return connection @staticmethod @@ -1216,7 +1219,13 @@ class MCPRequestHandler: raise HTTPException(status_code=401, detail="Invalid or expired credential") @staticmethod - async def _enforce_admitted_live_policy(admitted: UserAPIKeyAuth, request: Request, route: str) -> None: + async def _enforce_admitted_live_policy( + admitted: UserAPIKeyAuth, + request: Request, + route: str, + *, + request_data: dict[str, object] | None = None, + ) -> None: """Run the standard pipeline's authorization checks over the admitted identity. Mirrors the ``user_api_key_auth`` wrapper between the builder and its return: clear the @@ -1246,7 +1255,7 @@ class MCPRequestHandler: await _run_centralized_common_checks( user_api_key_auth_obj=admitted, request=request, - request_data=await _read_request_body(request=request), + request_data=await _read_request_body(request=request) if request_data is None else request_data, route=route, ) except (HTTPException, ProxyException): diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index b918a687d85..7950dd36a14 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1041,6 +1041,7 @@ async def exchange_token_with_server( scope: str | None = None, client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None, connection_binding: ConnectionBinding | None = None, + connection_claim: tuple[str, int] | None = None, ): _raise_if_not_oauth2(mcp_server) if grant_type not in ("authorization_code", "refresh_token"): @@ -1197,6 +1198,10 @@ async def exchange_token_with_server( ) async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) + if connection_claim is not None: + claimed: Final = await claim_connection_once(*connection_claim) + if claimed is not None: + return claimed try: response: Final = await async_client.post( token_url, @@ -3092,11 +3097,8 @@ async def complete_connection(request: Request, flow: str = Form(...), decision: if opened is None or opened.jti != flow or decision not in ("approve", "deny"): return _oauth_error(400, "invalid_request", "Invalid or expired consent") server: Final = await validate_connection_binding(request, opened.binding) - claim: Final = await claim_connection_once(f"consent:{opened.jti}", opened.exp) - if claim is not None: - return claim - if decision == "deny": - denied: Final = RedirectResponse( + response: Final = ( + RedirectResponse( append_connection_query( opened.redirect_uri, ( @@ -3106,21 +3108,25 @@ async def complete_connection(request: Request, flow: str = Form(...), decision: ), status_code=302, ) - cookie_path, _ = _cookie_path_and_secure(request) - denied.delete_cookie(cookie_name, path=cookie_path) - return denied - response: Final = await authorize_with_server( - request, - server, - opened.client_id, - opened.redirect_uri, - opened.state, - opened.code_challenge, - "S256", - "code", - opened.scope, - connection=opened, + if decision == "deny" + else await authorize_with_server( + request, + server, + opened.client_id, + opened.redirect_uri, + opened.state, + opened.code_challenge, + "S256", + "code", + opened.scope, + connection=opened, + ) ) + if response.status_code >= 400: + return response + claim: Final = await claim_connection_once(f"consent:{opened.jti}", opened.exp) + if claim is not None: + return claim cookie_path, _ = _cookie_path_and_secure(request) response.delete_cookie(cookie_name, path=cookie_path) return response @@ -3153,9 +3159,6 @@ async def exchange_connection_token( or not _pkce_verifier_matches(code_verifier, opened.authorization.code_challenge) ): return _oauth_error(400, "invalid_grant", "Invalid authorization code or PKCE verifier") - claim: Final = await claim_connection_once(f"code:{opened.jti}", opened.exp) - if claim is not None: - return claim return await exchange_token_with_server( request, server, @@ -3167,6 +3170,7 @@ async def exchange_connection_token( code_verifier, scope=opened.authorization.scope, connection_binding=binding, + connection_claim=(f"code:{opened.jti}", opened.exp), ) if grant_type == "refresh_token": refreshed: Final = open_connection_credential(refresh_token or "", refresh=True) @@ -3174,9 +3178,6 @@ async def exchange_connection_token( return _oauth_error(400, "invalid_grant", "Invalid refresh credential") if scope and not frozenset(scope.split()).issubset((refreshed.scope or "").split()): return _oauth_error(400, "invalid_scope", "Refresh cannot expand the granted scopes") - claimed: Final = await claim_connection_once(f"refresh:{refreshed.jti}", refreshed.exp) - if claimed is not None: - return claimed return await exchange_token_with_server( request, server, @@ -3189,5 +3190,6 @@ async def exchange_connection_token( refresh_token=refreshed.token.get_secret_value(), scope=scope or refreshed.scope, connection_binding=binding, + connection_claim=(f"refresh:{refreshed.jti}", refreshed.exp), ) return _oauth_error(400, "unsupported_grant_type", "Unsupported connection grant type") diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index 8ec32ce88ed..8328d66532a 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -1632,7 +1632,9 @@ async def validate_connection_binding(request: Request, binding: ConnectionBindi if binding.resource != f"{get_request_base_url(request)}/mcp": raise HTTPException(status_code=400, detail="Invalid connection resource") key: Final = await MCPRequestHandler._reload_admitted_key(binding.key_hash) # pyright: ignore[reportPrivateUsage] # reuse key revocation and SCIM checks - await MCPRequestHandler._enforce_admitted_live_policy(key.model_copy(), request, "/mcp") # pyright: ignore[reportPrivateUsage] # enforce the same MCP route and budget policy + await MCPRequestHandler._enforce_admitted_live_policy( # pyright: ignore[reportPrivateUsage] # enforce the same MCP route and budget policy + key.model_copy(), request, "/mcp", request_data={} + ) allowed: Final = await MCPRequestHandler.get_allowed_mcp_servers(key) server: Final = global_mcp_server_manager.get_mcp_server_by_id( binding.server_id, client_ip=IPAddressUtils.get_mcp_client_ip(request) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 2b92367f186..7873567231e 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -953,7 +953,9 @@ async def test_get_tools_from_mcp_servers(): client_ip=None, user_api_key_auth=None, oauth2_headers=None, + connection_credential=None, ): + assert connection_credential is None if server.server_id == "server1_id": return [mock_tool_1] return [mock_tool_2] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index d7ab4628cc6..c876e8f28e1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -1817,7 +1817,8 @@ class TestMCPPublicRouteGuard: await MCPRequestHandler.process_mcp_request(scope) assert exc_info.value.status_code == 401 - async def test_legitimate_well_known_path_still_bypasses_auth(self): + @pytest.mark.parametrize("bearer", [None, "llm_caccess_stale", "llm_crefresh_stale"]) + async def test_legitimate_well_known_path_still_bypasses_auth(self, bearer): """ Real OAuth discovery routes registered under /.well-known/ must remain public so unauthenticated clients can fetch them per RFC 8414/9728. @@ -1826,16 +1827,21 @@ class TestMCPPublicRouteGuard: "type": "http", "method": "GET", "path": "/.well-known/oauth-protected-resource", - "headers": [], + "headers": [(b"authorization", f"Bearer {bearer}".encode())] if bearer else [], } # No mock needed — public path should not call user_api_key_auth at all with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", ) as mock_auth: - auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) + auth_result, mcp_header, _, server_headers, oauth_headers, raw_headers = ( + await MCPRequestHandler.process_mcp_request(scope) + ) mock_auth.assert_not_called() assert isinstance(auth_result, UserAPIKeyAuth) + assert "litellm.mcp.connection_grant" not in scope + assert not mcp_header and not server_headers and not oauth_headers + assert "authorization" not in raw_headers @pytest.mark.asyncio @@ -9449,7 +9455,7 @@ class TestScopedSessionAdmission: (None, "connection-target", 401), ], ) -@pytest.mark.parametrize("grant_state", ["valid", "expired", "refresh", "oversized"]) +@pytest.mark.parametrize("grant_state", ["valid", "expired", "refresh", "oversized", "other-resource"]) @pytest.mark.parametrize("server_mode", ["managed", "delegated"]) async def test_connection_credential_requires_exact_key_and_server( monkeypatch, key, selector, expected, grant_state, server_mode @@ -9485,7 +9491,9 @@ async def test_connection_credential_requires_exact_key_and_server( ) monkeypatch.setattr(admission, "user_api_key_auth", AsyncMock(return_value=auth)) binding = ConnectionBinding( - key_hash=hash_token("sk-original"), server_id=server.server_id, resource="https://gateway.example/mcp" + key_hash=hash_token("sk-original"), + server_id=server.server_id, + resource="https://other.example/mcp" if grant_state == "other-resource" else "https://gateway.example/mcp", ) issued = json.loads( flow.mint_connection_tokens( @@ -9528,11 +9536,18 @@ async def test_connection_credential_requires_exact_key_and_server( assert exc.value.status_code == expected_status if ( server_mode == "managed" - and key == "sk-original" + and key is not None and selector == "connection-target" - and grant_state != "valid" ): - assert "resource_metadata=" in exc.value.headers["www-authenticate"] + from urllib.parse import parse_qs, urlparse + + challenge = exc.value.headers["www-authenticate"] + metadata_url = challenge.split('resource_metadata="', 1)[1].split('"', 1)[0] + bootstrap = parse_qs(urlparse(metadata_url).query)["connection"][0] + assert flow.open_connection_bootstrap(bootstrap) == ConnectionBinding( + key_hash=hash_token(key), server_id=server.server_id, resource="https://gateway.example/mcp" + ) + assert exc.value.headers["Cache-Control"] == "no-store" assert flow.CONNECTION_SCOPE_KEY not in scope return _, _, _, server_headers, oauth_headers, raw_headers = await MCPRequestHandler.process_mcp_request(scope) 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 7fad4fb6b02..5c4bdf2943d 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 @@ -12574,6 +12574,79 @@ def test_keyed_connection_headerless_exchange_and_rotating_refresh(keyed_oauth_c harness.vault.assert_not_called() +def test_keyed_connection_consent_can_retry_failed_preparation(keyed_oauth_client, monkeypatch): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + harness = keyed_oauth_client + _, _, handle = _start_keyed_oauth(harness) + prepare = AsyncMock(side_effect=[HTTPException(status_code=503, detail="discovery unavailable"), harness.server]) + monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", prepare) + payload = {"flow": handle, "decision": "approve"} + + failed = harness.client.post("/authorize/connection/complete", data=payload) + assert failed.status_code == 503 + assert "location" not in failed.headers + assert not failed.headers.get("set-cookie") + retried = harness.client.post("/authorize/connection/complete", data=payload) + assert retried.status_code == 303 + assert retried.headers["location"].startswith("https://provider.example/authorize?") + harness.upstream.post.assert_not_awaited() + + +@pytest.mark.parametrize("grant_type", ["authorization_code", "refresh_token"]) +def test_keyed_connection_exchange_can_retry_failed_preparation(keyed_oauth_client, monkeypatch, grant_type): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + harness = keyed_oauth_client + code_payload = _complete_keyed_oauth(harness) + tokens = harness.client.post("/token", data=code_payload).json() if grant_type == "refresh_token" else None + payload = ( + {"grant_type": grant_type, "client_id": code_payload["client_id"], "refresh_token": tokens["refresh_token"]} + if tokens is not None else code_payload + ) + harness.upstream.post.reset_mock() + prepare = AsyncMock(side_effect=[HTTPException(status_code=503, detail="discovery unavailable"), harness.server]) + monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", prepare) + + failed = harness.client.post("/token", data=payload) + assert failed.status_code == 503 + harness.upstream.post.assert_not_awaited() + retried = harness.client.post("/token", data=payload) + assert retried.status_code == 200, retried.text + assert retried.json()["access_token"].startswith("llm_caccess_") + harness.upstream.post.assert_awaited_once() + + +@pytest.mark.parametrize("grant_type", ["authorization_code", "refresh_token"]) +@pytest.mark.parametrize("failure", ["response", "timeout"]) +def test_keyed_connection_dispatched_provider_failure_cannot_redeem_twice(keyed_oauth_client, grant_type, failure): + import httpx + + harness = keyed_oauth_client + code_payload = _complete_keyed_oauth(harness) + tokens = harness.client.post("/token", data=code_payload).json() if grant_type == "refresh_token" else None + payload = ( + {"grant_type": grant_type, "client_id": code_payload["client_id"], "refresh_token": tokens["refresh_token"]} + if tokens is not None else code_payload + ) + harness.upstream.post.reset_mock() + harness.upstream.post.return_value = httpx.Response( + 503, json={"error": "temporarily_unavailable"}, request=httpx.Request("POST", "https://provider.example/token") + ) + + if failure == "timeout": + harness.upstream.post.side_effect = httpx.ReadTimeout("provider response unavailable") + with pytest.raises(httpx.ReadTimeout): + harness.client.post("/token", data=payload) + else: + failed = harness.client.post("/token", data=payload) + assert failed.status_code >= 500 + replayed = harness.client.post("/token", data=payload) + assert replayed.status_code == 400 + assert replayed.json()["error"] == "invalid_grant" + harness.upstream.post.assert_awaited_once() + + @pytest.mark.parametrize("change", ["verifier", "redirect", "resource", "tamper"]) def test_keyed_connection_rejects_bad_exchange_before_upstream(keyed_oauth_client, change): harness = keyed_oauth_client @@ -12672,7 +12745,7 @@ def test_keyed_connection_client_cannot_enter_session_login_flow(keyed_oauth_cli @pytest.mark.parametrize("stage", ["authorize", "exchange", "refresh"]) -@pytest.mark.parametrize("policy", ["blocked", "expired", "denied_server", "denied_route", "allowed"]) +@pytest.mark.parametrize("policy", ["blocked", "expired", "denied_server", "denied_route", "over_budget", "allowed"]) def test_keyed_connection_reloads_live_key_before_provider(keyed_oauth_client, monkeypatch, stage, policy): from litellm.proxy import proxy_server from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow @@ -12702,7 +12775,12 @@ def test_keyed_connection_reloads_live_key_before_provider(keyed_oauth_client, m lookup = AsyncMock(return_value=key) monkeypatch.setattr(auth_checks, "get_key_object", lookup) monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) - monkeypatch.setattr(admission, "_run_centralized_common_checks", AsyncMock()) + import litellm + + common_checks = AsyncMock( + side_effect=litellm.BudgetExceededError(current_cost=2, max_budget=1) if policy == "over_budget" else None + ) + monkeypatch.setattr(admission, "_run_centralized_common_checks", common_checks) monkeypatch.setitem(global_mcp_server_manager.registry, harness.server.server_id, harness.server) if stage == "authorize": challenge = urlsafe_b64encode(hashlib.sha256(payload["code_verifier"].encode()).digest()).rstrip(b"=").decode() @@ -12718,25 +12796,30 @@ def test_keyed_connection_reloads_live_key_before_provider(keyed_oauth_client, m }, ) elif stage == "exchange": - response = harness.client.post("/token", data=payload) + response = harness.client.post("/token", data={**payload, "model": "untrusted\nforged log entry"}) else: assert token is not None and token.status_code == 200 response = harness.client.post( "/token", data={ + "model": "untrusted\nforged log entry", "grant_type": "refresh_token", "client_id": payload["client_id"], "refresh_token": token.json()["refresh_token"], "resource": harness.binding.resource, }, ) - assert response.status_code == (200 if policy == "allowed" else 401 if policy in ("blocked", "expired") else 403), ( + assert response.status_code == (200 if policy == "allowed" else 401 if policy in ("blocked", "expired") else 422 if policy == "over_budget" else 403), ( response.text ) lookup.assert_awaited_once() assert lookup.call_args.kwargs["hashed_token"] == harness.binding.key_hash assert harness.upstream.post.call_count == int(policy == "allowed" and stage != "authorize") harness.vault.assert_not_called() + if policy not in ("blocked", "expired", "denied_route"): + common_checks.assert_awaited_once() + assert common_checks.call_args.kwargs["request_data"] == {} + assert common_checks.call_args.kwargs["route"] == "/mcp" @pytest.mark.parametrize(