fix: replace asserts with proper guards, wrap HTTP errors with context

- Replace `assert` statements with `if/raise ValueError` (asserts can be
  disabled with python -O in production)
- Wrap `httpx.HTTPStatusError` to provide a clear error message with
  server_id and status code
- Add tests for HTTP error and non-dict JSON response error paths
- Remove unused imports

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Ishaan Jaffer 2026-02-09 17:19:56 -08:00
parent e1cd2eabca
commit 9b8cdc3d72
2 changed files with 55 additions and 7 deletions

View file

@ -8,6 +8,8 @@ with ``client_id``, ``client_secret``, and ``token_url``.
import asyncio
from typing import TYPE_CHECKING, Dict, Optional, Tuple, Union
import httpx
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
@ -74,9 +76,13 @@ class MCPOAuth2TokenCache(InMemoryCache):
"""
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
assert server.client_id is not None, "client_id must be set"
assert server.client_secret is not None, "client_secret must be set"
assert server.token_url is not None, "token_url must be set"
if not server.client_id or not server.client_secret or not server.token_url:
raise ValueError(
f"MCP server '{server.server_id}' missing required OAuth2 fields: "
f"client_id={bool(server.client_id)}, "
f"client_secret={bool(server.client_secret)}, "
f"token_url={bool(server.token_url)}"
)
data: Dict[str, str] = {
"grant_type": "client_credentials",
@ -91,8 +97,15 @@ class MCPOAuth2TokenCache(InMemoryCache):
server.server_id,
)
response = await client.post(server.token_url, data=data)
response.raise_for_status()
try:
response = await client.post(server.token_url, data=data)
response.raise_for_status()
except httpx.HTTPStatusError as exc:
raise ValueError(
f"OAuth2 token request for MCP server '{server.server_id}' "
f"failed with status {exc.response.status_code}"
) from exc
body = response.json()
if not isinstance(body, dict):

View file

@ -5,10 +5,9 @@ Covers the critical path: resolve_mcp_auth(), token caching, auth priority,
fallback to static token, and the skip-condition property.
"""
import asyncio
import time
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
@ -120,3 +119,39 @@ def test_needs_user_oauth_token_property():
# Non-OAuth2 → never needs user OAuth token
assert _server(auth_type=MCPAuth.bearer_token).needs_user_oauth_token is False
@pytest.mark.asyncio
async def test_http_error_raises_value_error():
"""HTTP errors from the token endpoint are wrapped in a clear ValueError."""
server = _server()
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
"Unauthorized", request=MagicMock(), response=mock_response,
)
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
with patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
return_value=mock_client,
), pytest.raises(ValueError, match="failed with status 401"):
await resolve_mcp_auth(server)
@pytest.mark.asyncio
async def test_non_dict_response_raises_value_error():
"""A non-dict JSON response raises a clear ValueError."""
server = _server()
resp = MagicMock()
resp.json.return_value = ["not", "a", "dict"]
resp.raise_for_status = MagicMock()
mock_client = AsyncMock()
mock_client.post.return_value = resp
with patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
return_value=mock_client,
), pytest.raises(ValueError, match="non-object JSON"):
await resolve_mcp_auth(server)