diff --git a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py index 426b072b473..226ca7520e3 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py @@ -14,6 +14,7 @@ import base64 import hashlib import hmac import html as _html_module +import json import os import time from typing import Dict, Optional @@ -200,6 +201,11 @@ async def openapi_oauth2_connect( base_url = get_request_base_url(request) callback_url = f"{base_url}/v1/mcp/oauth2/callback" + # NOTE: PKCE (RFC 7636 / OAuth 2.1) is not implemented here because this is + # a server-side *confidential* client that always presents a client_secret. + # Confidential clients are significantly less exposed to code-interception + # attacks than public clients. PKCE support for public/SPAs is tracked as + # a follow-up improvement. params: dict = { "client_id": server.client_id, "redirect_uri": callback_url, @@ -352,6 +358,7 @@ async def openapi_oauth2_callback( # Parse response: try JSON first, fall back to URL-encoded form (GitHub can return either) # Some providers return HTTP 200 with an error body, so check for error fields explicitly. access_token: Optional[str] = None + refresh_token: Optional[str] = None provider_error: Optional[str] = None content_type = response.headers.get("content-type", "") if "application/json" in content_type: @@ -363,6 +370,7 @@ async def openapi_oauth2_callback( provider_error = f"{err}: {err_desc}" if err_desc else err else: access_token = token_data.get("access_token") + refresh_token = token_data.get("refresh_token") except Exception: pass if access_token is None and provider_error is None: @@ -378,6 +386,9 @@ async def openapi_oauth2_callback( tokens = form_data.get("access_token", []) if tokens: access_token = tokens[0] + refresh_tokens = form_data.get("refresh_token", []) + if refresh_tokens: + refresh_token = refresh_tokens[0] except Exception: pass @@ -424,12 +435,21 @@ async def openapi_oauth2_callback( status_code=500, ) + # Persist the access token (and refresh token if the provider returned one). + # Stored as a JSON blob so the token retrieval path can surface the refresh + # token for future renewal without a schema change. + credential_to_store = ( + json.dumps({"access_token": access_token, "refresh_token": refresh_token}) + if refresh_token + else access_token + ) + try: await store_user_credential( prisma_client=prisma_client, user_id=user_id, server_id=server_id, - credential=access_token, + credential=credential_to_store, ) from litellm.proxy._experimental.mcp_server.server import ( _invalidate_byok_cred_cache, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index f728582d5c6..3d566ca84b5 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -5,6 +5,7 @@ LiteLLM MCP Server Routes import asyncio import contextlib +import json import time import traceback import uuid @@ -64,6 +65,25 @@ _BYOK_CRED_CACHE_TTL = 60 # seconds _BYOK_CRED_CACHE_MAX_SIZE = 4096 # cap to prevent unbounded growth +def _extract_access_token(credential: Optional[str]) -> Optional[str]: + """Extract the access_token from a stored credential. + + OAuth2 callbacks may store a JSON blob of the form + ``{"access_token": "...", "refresh_token": "..."}`` when the provider + returns a refresh token. This helper transparently handles both the JSON + format and plain-string credentials (e.g. static API keys or older entries). + """ + if credential is None: + return None + try: + data = json.loads(credential) + if isinstance(data, dict) and "access_token" in data: + return data["access_token"] + except (json.JSONDecodeError, ValueError): + pass + return credential + + def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None: """Remove a (user_id, server_id) entry from the BYOK credential cache. @@ -1555,11 +1575,16 @@ if MCP_AVAILABLE: if prisma_client is None: return None - credential = await get_user_credential( + raw = await get_user_credential( prisma_client=prisma_client, user_id=user_id, server_id=mcp_server.server_id, ) + # Credentials stored by the OAuth2 callback may be a JSON blob of the + # form {"access_token": "...", "refresh_token": "..."} when the provider + # returned a refresh token. Extract just the access_token so the rest of + # the auth-injection path continues to receive a plain string. + credential = _extract_access_token(raw) _write_byok_cred_cache(user_id, mcp_server.server_id, credential) return credential diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py index 7994fe8a5cc..cbe08a9fc1c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py @@ -1,5 +1,6 @@ """Unit tests for openapi_oauth2_endpoints.py""" +import json import sys import time from unittest.mock import AsyncMock, MagicMock, patch @@ -216,9 +217,197 @@ async def test_status_no_prisma_returns_not_connected(): mock_mgr.get_mcp_server_by_id.return_value = mock_server result = await openapi_oauth2_status("server1", mock_user) - import json - raw = result.body body = json.loads(raw.decode() if isinstance(raw, (bytes, bytearray)) else str(raw)) assert body["connected"] is False assert body["server_id"] == "server1" + + +# --------------------------------------------------------------------------- +# Refresh token — stored as JSON blob when provider returns one +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_callback_stores_refresh_token_as_json(): + """When the provider returns a refresh_token, it is stored as a JSON blob.""" + from litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints import ( + openapi_oauth2_callback, + ) + + state = "test-state-refresh" + now = time.time() + _pending_oauth2_states[state] = { + "server_id": "server1", + "user_id": "user1", + "timestamp": now, + "expires_at": now + 600, + } + + mock_server = MagicMock() + mock_server.token_url = "https://provider.example/token" + mock_server.client_id = "cid" + mock_server.client_secret = "csecret" + mock_server.server_name = "TestProvider" + mock_server.name = "test" + + mock_response = MagicMock() + mock_response.headers = {"content-type": "application/json"} + mock_response.json.return_value = { + "access_token": "ghu_accesstoken123", + "refresh_token": "ghr_refreshtoken456", + "token_type": "bearer", + } + mock_response.raise_for_status = MagicMock() + + stored_credentials: list = [] + + async def fake_store(prisma_client, user_id, server_id, credential): + stored_credentials.append(credential) + + with patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.global_mcp_server_manager" + ) as mock_mgr, patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.get_request_base_url", + return_value="http://localhost:4000", + ), patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.store_user_credential", + side_effect=fake_store, + ), patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(), + create=True, + ), patch( + "litellm.proxy._experimental.mcp_server.server._invalidate_byok_cred_cache", + MagicMock(), + ), patch( + "httpx.AsyncClient" + ) as mock_client_cls: + mock_mgr.get_mcp_server_by_id.return_value = mock_server + mock_async_client = AsyncMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + mock_client_cls.return_value.__aenter__ = AsyncMock(return_value=mock_async_client) + mock_client_cls.return_value.__aexit__ = AsyncMock(return_value=None) + + await openapi_oauth2_callback( + request=MagicMock(), + code="auth-code", + state=state, + error=None, + error_description=None, + ) + + assert len(stored_credentials) == 1 + stored = stored_credentials[0] + parsed = json.loads(stored) + assert parsed["access_token"] == "ghu_accesstoken123" + assert parsed["refresh_token"] == "ghr_refreshtoken456" + + +@pytest.mark.asyncio +async def test_callback_stores_plain_token_when_no_refresh_token(): + """When the provider does not return a refresh_token, the plain access_token is stored.""" + from litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints import ( + openapi_oauth2_callback, + ) + + state = "test-state-no-refresh" + now = time.time() + _pending_oauth2_states[state] = { + "server_id": "server1", + "user_id": "user1", + "timestamp": now, + "expires_at": now + 600, + } + + mock_server = MagicMock() + mock_server.token_url = "https://provider.example/token" + mock_server.client_id = "cid" + mock_server.client_secret = "csecret" + mock_server.server_name = "TestProvider" + mock_server.name = "test" + + mock_response = MagicMock() + mock_response.headers = {"content-type": "application/json"} + mock_response.json.return_value = { + "access_token": "ghu_only_access", + "token_type": "bearer", + } + mock_response.raise_for_status = MagicMock() + + stored_credentials: list = [] + + async def fake_store(prisma_client, user_id, server_id, credential): + stored_credentials.append(credential) + + with patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.global_mcp_server_manager" + ) as mock_mgr, patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.get_request_base_url", + return_value="http://localhost:4000", + ), patch( + "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.store_user_credential", + side_effect=fake_store, + ), patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(), + create=True, + ), patch( + "litellm.proxy._experimental.mcp_server.server._invalidate_byok_cred_cache", + MagicMock(), + ), patch( + "httpx.AsyncClient" + ) as mock_client_cls: + mock_mgr.get_mcp_server_by_id.return_value = mock_server + mock_async_client = AsyncMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + mock_client_cls.return_value.__aenter__ = AsyncMock(return_value=mock_async_client) + mock_client_cls.return_value.__aexit__ = AsyncMock(return_value=None) + + await openapi_oauth2_callback( + request=MagicMock(), + code="auth-code", + state=state, + error=None, + error_description=None, + ) + + assert len(stored_credentials) == 1 + # Plain string — NOT a JSON blob + assert stored_credentials[0] == "ghu_only_access" + + +# --------------------------------------------------------------------------- +# _extract_access_token (server.py helper) +# --------------------------------------------------------------------------- + + +def test_extract_access_token_plain_string(): + """Plain token strings are returned unchanged.""" + from litellm.proxy._experimental.mcp_server.server import _extract_access_token + + assert _extract_access_token("ghp_plaintoken") == "ghp_plaintoken" + + +def test_extract_access_token_json_blob(): + """JSON blob with access_token + refresh_token → access_token returned.""" + from litellm.proxy._experimental.mcp_server.server import _extract_access_token + + blob = json.dumps({"access_token": "ghu_access", "refresh_token": "ghr_refresh"}) + assert _extract_access_token(blob) == "ghu_access" + + +def test_extract_access_token_none(): + """None input returns None.""" + from litellm.proxy._experimental.mcp_server.server import _extract_access_token + + assert _extract_access_token(None) is None + + +def test_extract_access_token_json_without_access_token_key(): + """JSON object without 'access_token' key is returned as-is (treated as plain string).""" + from litellm.proxy._experimental.mcp_server.server import _extract_access_token + + blob = json.dumps({"some_other_key": "value"}) + # Falls back to raw string since there's no access_token + assert _extract_access_token(blob) == blob