mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
0b6fa2b4b1
commit
d52db24931
3 changed files with 238 additions and 4 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue