diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 246e0d397d3..ad9e069f28f 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -382,23 +382,44 @@ async def authorize_with_server( async def _post_to_upstream_token_endpoint(token_url: str, token_data: Dict[str, Any]): """POST to the upstream IdP /token endpoint. - Uses a raw httpx client (instead of get_async_httpx_client) so we can inspect - 4xx/5xx responses and surface them as proper OAuth2 error JSON per RFC 6749 - §5.2. The wrapped client raises MaskedHTTPStatusError on raise_for_status() - before downstream code can read the response body, so the actual - AADSTS / error_description from Entra/Okta is lost and the client sees - a generic 500. + Uses the pooled get_async_httpx_client and catches the wrapped + HTTPStatusError so the upstream `error` / `error_description` payload + (RFC 6749 §5.2) reaches the client. The wrapper turns any 4xx/5xx into + a MaskedHTTPStatusError before downstream code can read the response + body, which would otherwise surface as a generic 500 and hide the + actual AADSTS / error_description from Entra/Okta. - Returns the parsed JSON dict on success, or a JSONResponse with the upstream - error body on 4xx/5xx (or a transport failure). + Returns the parsed JSON dict on success, or a JSONResponse with the + upstream error body on 4xx/5xx (or a transport failure). """ + async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) try: - async with httpx.AsyncClient(timeout=30.0) as raw_client: - response = await raw_client.post( - token_url, - headers={"Accept": "application/json"}, - data=token_data, + response = await async_client.post( + token_url, + headers={"Accept": "application/json"}, + data=token_data, + ) + except httpx.HTTPStatusError as exc: + upstream_response = exc.response + try: + err_body = upstream_response.json() + except ValueError: + # MaskedHTTPStatusError (the wrapped form) carries a sanitized body + # on `.text`; fall back to the raw response text otherwise. + description = ( + getattr(exc, "text", None) + or upstream_response.text + or "upstream token endpoint error" ) + err_body = { + "error": "invalid_request", + "error_description": description, + } + return JSONResponse( + err_body, + status_code=upstream_response.status_code, + headers=TOKEN_NO_CACHE_HEADERS, + ) except httpx.HTTPError as exc: return JSONResponse( {"error": "server_error", "error_description": str(exc)}, @@ -406,20 +427,6 @@ async def _post_to_upstream_token_endpoint(token_url: str, token_data: Dict[str, headers=TOKEN_NO_CACHE_HEADERS, ) - if response.status_code >= 400: - try: - err_body = response.json() - except ValueError: - err_body = { - "error": "invalid_request", - "error_description": response.text or "upstream token endpoint error", - } - return JSONResponse( - err_body, - status_code=response.status_code, - headers=TOKEN_NO_CACHE_HEADERS, - ) - return response.json() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index ef23c258154..f80ff3123d3 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -3,6 +3,7 @@ import json from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from fastapi import HTTPException @@ -296,13 +297,10 @@ async def test_token_endpoint_forwards_code_verifier(): mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) - fake_cm = MagicMock() - fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client) - fake_cm.__aexit__ = AsyncMock(return_value=None) with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient", - return_value=fake_cm, + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, ): # Call token endpoint with code_verifier response = await token_endpoint( @@ -641,13 +639,10 @@ async def test_token_endpoint_respects_x_forwarded_proto(): mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) - fake_cm = MagicMock() - fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client) - fake_cm.__aexit__ = AsyncMock(return_value=None) with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient", - return_value=fake_cm, + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, ): await token_endpoint( request=mock_request, @@ -947,13 +942,10 @@ async def test_token_endpoint_respects_x_forwarded_host(): mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) - fake_cm = MagicMock() - fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client) - fake_cm.__aexit__ = AsyncMock(return_value=None) with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient", - return_value=fake_cm, + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, ): await token_endpoint( request=mock_request, @@ -1578,14 +1570,11 @@ async def test_token_root_resolves_single_oauth2_server(): mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) - fake_cm = MagicMock() - fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client) - fake_cm.__aexit__ = AsyncMock(return_value=None) try: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient", - return_value=fake_cm, + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, ): # Call /token WITHOUT mcp_server_name response = await token_endpoint( @@ -2101,13 +2090,10 @@ async def test_token_endpoint_refresh_token_grant(): mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) - fake_cm = MagicMock() - fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client) - fake_cm.__aexit__ = AsyncMock(return_value=None) with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient", - return_value=fake_cm, + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, ): response = await token_endpoint( request=mock_request, @@ -2398,13 +2384,10 @@ async def test_token_endpoint_sets_no_store_cache_control(): } fake_http_client = MagicMock() fake_http_client.post = AsyncMock(return_value=fake_http_response) - fake_cm = MagicMock() - fake_cm.__aenter__ = AsyncMock(return_value=fake_http_client) - fake_cm.__aexit__ = AsyncMock(return_value=None) with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient", - return_value=fake_cm, + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=fake_http_client, ): response = await exchange_token_with_server( request=mock_request, @@ -2646,13 +2629,10 @@ async def test_exchange_token_omits_placeholder_client_secret(placeholder): mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) - fake_cm = MagicMock() - fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client) - fake_cm.__aexit__ = AsyncMock(return_value=None) with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient", - return_value=fake_cm, + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, ): await exchange_token_with_server( request=mock_request, @@ -2706,19 +2686,18 @@ async def test_exchange_token_surfaces_upstream_4xx_as_oauth_error_json(): "error": "invalid_grant", "error_description": "AADSTS70008: The provided authorization code or refresh token has expired", } - mock_response = MagicMock() - mock_response.status_code = 400 - mock_response.json.return_value = upstream_body + fake_request = httpx.Request("POST", server.token_url) + fake_response = httpx.Response(400, json=upstream_body, request=fake_request) + raised = httpx.HTTPStatusError( + "400 Bad Request", request=fake_request, response=fake_response + ) mock_async_client = MagicMock() - mock_async_client.post = AsyncMock(return_value=mock_response) - fake_cm = MagicMock() - fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client) - fake_cm.__aexit__ = AsyncMock(return_value=None) + mock_async_client.post = AsyncMock(side_effect=raised) with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient", - return_value=fake_cm, + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, ): response = await exchange_token_with_server( request=mock_request, @@ -2769,20 +2748,22 @@ async def test_exchange_token_handles_upstream_non_json_4xx(): mock_request.base_url = "https://proxy.litellm.example/" mock_request.headers = {} - mock_response = MagicMock() - mock_response.status_code = 502 - mock_response.json.side_effect = ValueError("not json") - mock_response.text = "upstream gateway error" + fake_request = httpx.Request("POST", server.token_url) + fake_response = httpx.Response( + 502, + content=b"upstream gateway error", + request=fake_request, + ) + raised = httpx.HTTPStatusError( + "502 Bad Gateway", request=fake_request, response=fake_response + ) mock_async_client = MagicMock() - mock_async_client.post = AsyncMock(return_value=mock_response) - fake_cm = MagicMock() - fake_cm.__aenter__ = AsyncMock(return_value=mock_async_client) - fake_cm.__aexit__ = AsyncMock(return_value=None) + mock_async_client.post = AsyncMock(side_effect=raised) with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints.httpx.AsyncClient", - return_value=fake_cm, + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, ): response = await exchange_token_with_server( request=mock_request,