mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
1ce66e98a2
commit
b9673b1e6a
3 changed files with 117 additions and 9 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue