mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(mcp): relay upstream OAuth errors from the token endpoint instead of masking as 500
Fixes #29588. The async HTTP handler raises on an upstream 4xx (MaskedHTTPStatusError) and exchange_token_with_server doesn't catch it, so an upstream 400 invalid_grant for an expired or rotated refresh token reaches the client as an opaque 500. The client can't tell its refresh token is dead, so it never starts a new authorization flow and retries the same token; headless clients loop. Relay the upstream status and OAuth error body per RFC 6749 §5.2 via a new relay_oauth_token_error helper in oauth_utils, matching the BYOK endpoint's _oauth_token_error behavior. MaskedHTTPStatusError subclasses httpx.HTTPStatusError and preserves the upstream status and body, so relaying is exact. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
parent
d45e9e4d56
commit
c8866c59c8
3 changed files with 102 additions and 5 deletions
|
|
@ -17,6 +17,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
TOKEN_NO_CACHE_HEADERS,
|
||||
get_request_base_url,
|
||||
relay_oauth_token_error,
|
||||
validate_trusted_redirect_uri,
|
||||
)
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
|
|
@ -456,11 +457,18 @@ async def exchange_token_with_server(
|
|||
token_data["code_verifier"] = code_verifier
|
||||
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response = await async_client.post(
|
||||
mcp_server.token_url,
|
||||
headers={"Accept": "application/json"},
|
||||
data=token_data,
|
||||
)
|
||||
try:
|
||||
response = await async_client.post(
|
||||
mcp_server.token_url,
|
||||
headers={"Accept": "application/json"},
|
||||
data=token_data,
|
||||
)
|
||||
except httpx.HTTPStatusError as e:
|
||||
# The async HTTP handler raises on an upstream 4xx/5xx. Relay the OAuth
|
||||
# error to the client (RFC 6749 §5.2) instead of letting it surface as a
|
||||
# 500; otherwise the client can't see e.g. invalid_grant and start a new
|
||||
# authorization flow, and a headless client loops on a dead refresh token.
|
||||
return relay_oauth_token_error(e.response)
|
||||
if response is None:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
|
|
|
|||
|
|
@ -6,7 +6,9 @@ from ipaddress import ip_address
|
|||
from typing import Any, Dict, List, NoReturn, Optional
|
||||
from urllib.parse import ParseResult, urlparse, urlunparse
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
|
|
@ -65,6 +67,23 @@ def _oauth_invalid_request(
|
|||
raise HTTPException(status_code=400, detail=detail)
|
||||
|
||||
|
||||
def relay_oauth_token_error(response: httpx.Response) -> JSONResponse:
|
||||
"""Relay an upstream OAuth token-endpoint error to the client unchanged
|
||||
(RFC 6749 §5.2). Masking it as a 500 hides the ``error`` field (e.g.
|
||||
``invalid_grant``) that clients rely on to drop a dead token and start a
|
||||
new authorization flow.
|
||||
"""
|
||||
try:
|
||||
content = response.json()
|
||||
except Exception:
|
||||
content = {"error": "invalid_request", "error_description": response.text}
|
||||
return JSONResponse(
|
||||
status_code=response.status_code,
|
||||
content=content,
|
||||
headers=TOKEN_NO_CACHE_HEADERS,
|
||||
)
|
||||
|
||||
|
||||
def _origin_label(scheme: str, netloc: str) -> str:
|
||||
"""Human-readable origin for error messages (scheme + host[:port])."""
|
||||
return f"{scheme}://{netloc}" if netloc else f"{scheme}://"
|
||||
|
|
|
|||
|
|
@ -2661,3 +2661,73 @@ async def test_token_endpoint_sets_no_store_cache_control():
|
|||
|
||||
assert response.headers["cache-control"] == "no-store"
|
||||
assert response.headers["pragma"] == "no-cache"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_token_relays_upstream_oauth_error():
|
||||
"""An upstream OAuth error (e.g. 400 invalid_grant for an expired refresh
|
||||
token) is relayed to the client per RFC 6749 §5.2, not masked as a 500, so
|
||||
the client can drop the dead token and start a new authorization flow."""
|
||||
import json
|
||||
|
||||
import httpx
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
exchange_token_with_server,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server = MCPServer(
|
||||
server_id="t",
|
||||
name="t",
|
||||
server_name="t",
|
||||
alias="t",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="cid",
|
||||
client_secret="cs",
|
||||
authorization_url="https://provider.com/oauth/authorize",
|
||||
token_url="https://provider.com/oauth/token",
|
||||
)
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
# The async HTTP handler raises on an upstream 4xx (it calls
|
||||
# raise_for_status internally; MaskedHTTPStatusError subclasses
|
||||
# httpx.HTTPStatusError and preserves the upstream status + body).
|
||||
upstream = httpx.Response(
|
||||
status_code=400,
|
||||
json={"error": "invalid_grant"},
|
||||
request=httpx.Request("POST", server.token_url),
|
||||
)
|
||||
fake_http_client = MagicMock()
|
||||
fake_http_client.post = AsyncMock(
|
||||
side_effect=httpx.HTTPStatusError(
|
||||
"400 Bad Request", request=upstream.request, response=upstream
|
||||
)
|
||||
)
|
||||
|
||||
with patch(
|
||||
"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,
|
||||
mcp_server=server,
|
||||
grant_type="refresh_token",
|
||||
code=None,
|
||||
redirect_uri=None,
|
||||
client_id="cid",
|
||||
client_secret=None,
|
||||
code_verifier=None,
|
||||
refresh_token="expired-refresh-token",
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body) == {"error": "invalid_grant"}
|
||||
# Error responses must not be cached either.
|
||||
assert response.headers["cache-control"] == "no-store"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue