mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
a00628bfc4
commit
7670d3f097
3 changed files with 112 additions and 21 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue