From 7670d3f09760a1965e0fe05ea7b28ce280947c51 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Sat, 7 Mar 2026 12:18:31 -0800 Subject: [PATCH] feat: support client_secret_basic token auth, use patch.object for require_byok flag tests MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Add token_endpoint_auth_method field to MCPServer ("post" default, "basic" for Okta/Auth0) - Implement HTTP Basic auth branch in token exchange (client_secret_basic per RFC 6749 §2.3.1) - Use patch.object(litellm, "require_byok_credential_store", True) in 503 tests to avoid global mutation - Add callback_url to all three callback test state entries (exercises primary path, not fallback) - Add test asserting Basic auth creds go in Authorization header, not POST body --- .../mcp_server/openapi_oauth2_endpoints.py | 19 ++- .../types/mcp_server/mcp_server_manager.py | 3 + .../test_openapi_oauth2_endpoints.py | 111 +++++++++++++++--- 3 files changed, 112 insertions(+), 21 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py index 34a62e1d6f5..7dda962bf3e 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py @@ -10,6 +10,7 @@ Endpoints: GET /v1/mcp/server/{server_id}/oauth2/status — check if user is connected """ +import base64 import html as _html_module import json import secrets @@ -346,19 +347,31 @@ async def openapi_oauth2_callback( ) token_request_data = { - "client_id": server.client_id or "", - "client_secret": server.client_secret or "", "code": code, "redirect_uri": callback_url, "grant_type": "authorization_code", } + # RFC 6749 §2.3: client credentials can be sent as POST body + # (client_secret_post, the default) or as HTTP Basic auth + # (client_secret_basic, required by some providers like Okta/Auth0). + auth_method = getattr(server, "token_endpoint_auth_method", "post") + token_headers: Dict[str, str] = {"Accept": "application/json"} + if auth_method == "basic": + basic_creds = base64.b64encode( + f"{server.client_id or ''}:{server.client_secret or ''}".encode() + ).decode() + token_headers["Authorization"] = f"Basic {basic_creds}" + else: + token_request_data["client_id"] = server.client_id or "" + token_request_data["client_secret"] = server.client_secret or "" + try: async with httpx.AsyncClient() as client: response = await client.post( server.token_url, data=token_request_data, - headers={"Accept": "application/json"}, + headers=token_headers, timeout=30.0, ) response.raise_for_status() diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index d94795fda2e..c8b7a8b32ea 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -48,6 +48,9 @@ class MCPServer(BaseModel): authorization_url: Optional[str] = None token_url: Optional[str] = None registration_url: Optional[str] = None + # "post" (default, RFC 6749 §2.3.1 client_secret_post) or + # "basic" (HTTP Basic auth via Authorization header, required by some providers) + token_endpoint_auth_method: str = "post" # Stdio-specific fields command: Optional[str] = None args: Optional[List[str]] = None 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 a3c677b39d7..79c99e64982 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 @@ -141,6 +141,7 @@ async def test_callback_provider_error_in_json_body(): "user_id": "user1", "timestamp": now, "expires_at": now + 600, + "callback_url": "http://localhost:4000/v1/mcp/oauth2/callback", } mock_server = MagicMock() @@ -246,6 +247,7 @@ async def test_callback_stores_refresh_token_as_json(): "user_id": "user1", "timestamp": now, "expires_at": now + 600, + "callback_url": "http://localhost:4000/v1/mcp/oauth2/callback", } mock_server = MagicMock() @@ -322,6 +324,7 @@ async def test_callback_stores_plain_token_when_no_refresh_token(): "user_id": "user1", "timestamp": now, "expires_at": now + 600, + "callback_url": "http://localhost:4000/v1/mcp/oauth2/callback", } mock_server = MagicMock() @@ -381,6 +384,84 @@ async def test_callback_stores_plain_token_when_no_refresh_token(): assert stored_credentials[0] == "ghu_only_access" +@pytest.mark.asyncio +async def test_callback_uses_basic_auth_when_token_endpoint_auth_method_is_basic(): + """When token_endpoint_auth_method='basic', credentials go in HTTP Basic header, not body.""" + import base64 + + from litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints import ( + _pending_oauth2_states, + openapi_oauth2_callback, + ) + + state = "test-state-basic-auth" + now = time.time() + _pending_oauth2_states[state] = { + "server_id": "server1", + "user_id": "user1", + "timestamp": now, + "expires_at": now + 600, + "callback_url": "http://localhost:4000/v1/mcp/oauth2/callback", + } + + mock_server = MagicMock() + mock_server.token_url = "https://provider.example/token" + mock_server.client_id = "my_client_id" + mock_server.client_secret = "my_client_secret" + mock_server.token_endpoint_auth_method = "basic" + mock_server.server_name = "BasicProvider" + mock_server.name = "test" + + mock_response = MagicMock() + mock_response.headers = {"content-type": "application/json"} + mock_response.json.return_value = {"access_token": "tok_basic", "token_type": "bearer"} + mock_response.raise_for_status = MagicMock() + + captured_calls: list = [] + + async def fake_post(url, data=None, headers=None, timeout=None): + captured_calls.append({"url": url, "data": data, "headers": headers}) + return mock_response + + 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.store_user_credential", + side_effect=AsyncMock(), + ), 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(side_effect=fake_post) + 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(captured_calls) == 1 + call = captured_calls[0] + # Credentials must be in the Authorization header + expected_basic = base64.b64encode(b"my_client_id:my_client_secret").decode() + assert call["headers"].get("Authorization") == f"Basic {expected_basic}" + # client_id and client_secret must NOT appear in the POST body + assert "client_id" not in (call["data"] or {}) + assert "client_secret" not in (call["data"] or {}) + + # --------------------------------------------------------------------------- # _extract_access_token (server.py helper) # --------------------------------------------------------------------------- @@ -537,17 +618,14 @@ async def test_check_byok_credential_raises_503_when_no_db(): import litellm - with patch( - "litellm.proxy.proxy_server.prisma_client", - None, - create=True, - ): - litellm.require_byok_credential_store = True # type: ignore[attr-defined] - try: + with patch.object(litellm, "require_byok_credential_store", True): + with patch( + "litellm.proxy.proxy_server.prisma_client", + None, + create=True, + ): with pytest.raises(HTTPException) as exc_info: await _check_byok_credential(mock_server, mock_user) - finally: - litellm.require_byok_credential_store = False # type: ignore[attr-defined] assert exc_info.value.status_code == 503 assert "byok_store_unavailable" in str(exc_info.value.detail) @@ -629,17 +707,14 @@ async def test_get_byok_credential_raises_503_when_no_db(): import litellm - with patch( - "litellm.proxy.proxy_server.prisma_client", - None, - create=True, - ): - litellm.require_byok_credential_store = True # type: ignore[attr-defined] - try: + with patch.object(litellm, "require_byok_credential_store", True): + with patch( + "litellm.proxy.proxy_server.prisma_client", + None, + create=True, + ): with pytest.raises(HTTPException) as exc_info: await _get_byok_credential(mock_server, mock_user) - finally: - litellm.require_byok_credential_store = False # type: ignore[attr-defined] assert exc_info.value.status_code == 503 assert "byok_store_unavailable" in str(exc_info.value.detail)