From b9673b1e6a02845a4c7a82a6b328a12ebc111dac Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 15 Sep 2026 06:48:34 +0000 Subject: [PATCH] fix(mcp): request offline access from Google upstream OAuth providers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 13 +-- .../_experimental/mcp_server/oauth_utils.py | 16 ++- .../mcp_server/test_discoverable_endpoints.py | 97 +++++++++++++++++++ 3 files changed, 117 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index bafe33d0a6b..72958b21c6a 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -64,6 +64,7 @@ from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( TOKEN_NO_CACHE_HEADERS, + build_upstream_authorize_url, build_upstream_oauth2_token_request, get_request_base_url, resolve_upstream_resource, @@ -799,9 +800,9 @@ def _redirect_to_upstream_authorize( **({"scope": scope_value} if scope_value else {}), **({"resource": upstream_resource} if upstream_resource else {}), } - parsed_auth_url: Final = urlparse(mcp_server.effective_authorization_url or "") - merged_params: Final = {**dict(parse_qsl(parsed_auth_url.query)), **passthrough_params} - return RedirectResponse(urlunparse(parsed_auth_url._replace(query=urlencode(merged_params)))) + return RedirectResponse( + build_upstream_authorize_url(mcp_server.effective_authorization_url or "", passthrough_params) + ) def _bridge_access_denied_redirect(redirect_uri: str, state: str, mcp_server: MCPServer) -> RedirectResponse: @@ -972,11 +973,7 @@ async def authorize_with_server( if upstream_resource: params["resource"] = upstream_resource - parsed_auth_url: Final = urlparse(resolved_server.effective_authorization_url) - existing_params: Final = dict(parse_qsl(parsed_auth_url.query)) - existing_params.update(params) - final_url: Final = urlunparse(parsed_auth_url._replace(query=urlencode(existing_params))) - response: Final = RedirectResponse(final_url) + response: Final = RedirectResponse(build_upstream_authorize_url(resolved_server.effective_authorization_url, params)) _set_oauth_state_cookie(response, request, relay_state, encoded_state) return response diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 39865a35ec6..d8be6f812ee 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -2,9 +2,11 @@ (BYOK + discoverable / pass-through OAuth proxy).""" import os +from collections.abc import Mapping from ipaddress import ip_address +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, NoReturn -from urllib.parse import ParseResult, urlparse, urlsplit, urlunparse, urlunsplit +from urllib.parse import ParseResult, parse_qsl, urlencode, urlparse, urlsplit, urlunparse, urlunsplit from fastapi import HTTPException, Request from starlette.types import Scope @@ -54,6 +56,9 @@ _DEFAULT_NATIVE_REDIRECT_URIS: Final[list[str]] = [ "cursor://anysphere.cursor-mcp/oauth/callback", ] +_GOOGLE_AUTHORIZATION_HOSTS: Final = frozenset({"accounts.google.com"}) +_GOOGLE_OFFLINE_ACCESS_PARAMS: Final = MappingProxyType({"access_type": "offline", "prompt": "consent"}) + _warned_invalid_proxy_base_url: str | None = None @@ -107,6 +112,15 @@ def _redact_mcp_resource_url(url: str | None) -> str | None: return urlunsplit((parts.scheme, netloc, "", "", "")) or None +def build_upstream_authorize_url(authorization_url: str, params: Mapping[str, str]) -> str: + parsed: Final = urlparse(authorization_url) + provider_defaults: Final = ( + _GOOGLE_OFFLINE_ACCESS_PARAMS if parsed.hostname in _GOOGLE_AUTHORIZATION_HOSTS else MappingProxyType({}) + ) + merged: Final = {**provider_defaults, **dict(parse_qsl(parsed.query)), **params} + return urlunparse(parsed._replace(query=urlencode(merged))) + + def _resolve_proxy_base_url_env() -> str | None: global _warned_invalid_proxy_base_url configured: Final = os.environ.get("PROXY_BASE_URL", "").strip() 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 9ea870d3210..75dce7fa9c6 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 @@ -570,6 +570,103 @@ async def test_authorize_endpoint_preserves_existing_query_params(): assert "scope=read+write" in location +async def _authorize_and_get_location_query(authorization_url: str): + from urllib.parse import parse_qs, urlparse + + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + global_mcp_server_manager.registry.clear() + + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url=authorization_url, + token_url="https://oauth2.googleapis.com/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + with patch( # test-quality-ok: real encryption needs a signing secret the unit env lacks + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper" + ) as mock_encrypt: + mock_encrypt.return_value = "mocked_encrypted_state" + + response = await authorize( + request=mock_request, + client_id="test_client_id", + mcp_server_name="test_oauth", + redirect_uri="http://127.0.0.1:60108/callback", + state="test_state", + ) + + return parse_qs(urlparse(response.headers["location"]).query) + + +@pytest.mark.asyncio +async def test_authorize_endpoint_requests_google_offline_access(): + """Google upstreams must get access_type=offline + prompt=consent so a refresh_token is issued""" + try: + import litellm.proxy._experimental.mcp_server.discoverable_endpoints # noqa: F401 + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + query = await _authorize_and_get_location_query("https://accounts.google.com/o/oauth2/v2/auth") + + assert query.get("access_type") == ["offline"] + assert query.get("prompt") == ["consent"] + assert query.get("client_id") == ["test_client_id"] + + +@pytest.mark.asyncio +async def test_authorize_endpoint_lets_configured_url_override_google_prompt(): + """An operator-set prompt on authorization_url wins over the Google default, offline access stays""" + try: + import litellm.proxy._experimental.mcp_server.discoverable_endpoints # noqa: F401 + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + query = await _authorize_and_get_location_query( + "https://accounts.google.com/o/oauth2/v2/auth?prompt=select_account" + ) + + assert query.get("prompt") == ["select_account"] + assert query.get("access_type") == ["offline"] + + +@pytest.mark.asyncio +async def test_authorize_endpoint_omits_offline_params_for_non_google_provider(): + """Non-Google upstreams get no access_type/prompt injected""" + try: + import litellm.proxy._experimental.mcp_server.discoverable_endpoints # noqa: F401 + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + query = await _authorize_and_get_location_query("https://provider.com/oauth/authorize") + + assert "access_type" not in query + assert "prompt" not in query + + @pytest.mark.asyncio async def test_authorize_endpoint_forwards_pkce_parameters(): """Test that authorize endpoint forwards PKCE parameters (code_challenge and code_challenge_method)"""