fix(mcp-oauth2): store refresh_token, add unit tests, note PKCE gap

- Extract and store refresh_token alongside access_token when the
  provider returns one; stored as JSON blob {"access_token": ...,
  "refresh_token": ...} so users aren't forced back through OAuth2
  consent when the access token expires
- Add _extract_access_token() helper in server.py that transparently
  handles both the new JSON blob format and legacy plain-string
  credentials (backward compatible)
- Add tests: refresh_token stored as JSON, plain token when no
  refresh_token, _extract_access_token with all input shapes
- Add PKCE comment acknowledging RFC 9700 gap; confidential client
  with client_secret is lower risk, full PKCE tracked as follow-up
This commit is contained in:
Ishaan Jaffer 2026-03-06 18:33:17 -08:00
parent 0b6fa2b4b1
commit d52db24931
3 changed files with 238 additions and 4 deletions

View file

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

View file

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

View file

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