feat: support client_secret_basic token auth, use patch.object for require_byok flag tests

- 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
This commit is contained in:
Ishaan Jaffer 2026-03-07 12:18:31 -08:00
parent a00628bfc4
commit 7670d3f097
3 changed files with 112 additions and 21 deletions

View file

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

View file

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

View file

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