fix(mcp): request offline access from Google upstream OAuth providers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-15 06:48:34 +00:00
parent 1ce66e98a2
commit b9673b1e6a
3 changed files with 117 additions and 9 deletions

View file

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

View file

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

View file

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