mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(mcp): keep pooled httpx client for upstream /token, catch wrapped HTTPStatusError
This commit is contained in:
parent
44f6da9693
commit
6f7aeb94dd
2 changed files with 69 additions and 81 deletions
|
|
@ -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<code> / 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<code> / 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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = "<html>upstream gateway error</html>"
|
||||
fake_request = httpx.Request("POST", server.token_url)
|
||||
fake_response = httpx.Response(
|
||||
502,
|
||||
content=b"<html>upstream gateway error</html>",
|
||||
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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue