diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d42c1c6b879..a271cf6c81a 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1,6 +1,7 @@ import asyncio import html as _html import json +import os import secrets import time from collections.abc import Callable, Mapping @@ -442,6 +443,10 @@ def _resolve_encoded_oauth_state(request: Request, state: str) -> str: return cookie_value if cookie_value else state +def _oauth_state_cookie_present(request: Request, state: str) -> bool: + return _oauth_state_cookie_name(state) in request.cookies + + def _clear_oauth_state_cookie(response: Response, request: Request, state: str) -> None: cookie_name: Final = _oauth_state_cookie_name(state) if cookie_name not in request.cookies: @@ -2250,9 +2255,31 @@ async def callback( ) # 3. Successful authorization response. + encoded_state = _resolve_encoded_oauth_state(request, state) try: - encoded_state = _resolve_encoded_oauth_state(request, state) state_data = decode_state_hash(encoded_state) + except Exception: # noqa: BLE001 # any decode failure means the session is unusable; surface it, never crash the callback + cookie_present: Final = _oauth_state_cookie_present(request, state) + verbose_logger.warning( + "MCP /callback could not decode OAuth state (state_cookie_present=%s, request_base_url=%s, " + "PROXY_BASE_URL_set=%s). If the cookie is absent, /authorize and /callback were served from " + "different origins; set PROXY_BASE_URL to the public origin or configure mcp_trusted_proxy_ranges.", + cookie_present, + get_request_base_url(request), + bool(os.environ.get("PROXY_BASE_URL", "").strip()), + ) + description: Final = ( + "The OAuth session cookie set when authorization started did not arrive at the callback. " + "This usually means the authorize and callback requests used different origins. " + "Ask the gateway operator to set PROXY_BASE_URL to the public URL of this gateway " + "(or configure mcp_trusted_proxy_ranges), then retry." + if not cookie_present + else "The OAuth session could not be decoded. Start the authorization again." + ) + response = _render_oauth_error_html("invalid_request", description) + _clear_oauth_state_cookie(response, request, state) + return response + try: original_state = state_data["original_state"] # Re-validate the client redirect URI at the sink. /authorize 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 f9a0075e530..af74bd38bf9 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 @@ -4264,20 +4264,78 @@ async def test_oauth_callback_handles_invalid_state(): except ImportError: pytest.skip("MCP discoverable endpoints not available") - # Mock state decoding to raise an exception - with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash") as mock_decode: - mock_decode.side_effect = Exception("Failed to decrypt state") + # A state blob that cannot be decrypted takes the same path as the + # LIT-4197 relay handle echoed back without its per-flow cookie + response = await callback( + request=_mock_callback_request(), + code="test_code", + state="invalid_encrypted_state", + ) - # Call callback endpoint with invalid state + # Should return a 400 error page that points at the missing state cookie + assert response.status_code == 400 + body = response.body.decode() + assert "did not arrive" in body + assert "PROXY_BASE_URL" in body + + +@pytest.mark.asyncio +async def test_oauth_callback_missing_state_cookie_logs_origin_hint(monkeypatch, caplog): + """A /callback whose per-flow state cookie never arrived must log a warning + carrying the resolved base URL and whether PROXY_BASE_URL is set, so the + origin-mismatch behind an ingress is diagnosable from logs.""" + import logging + + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + callback, + ) + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + monkeypatch.delenv("PROXY_BASE_URL", raising=False) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): response = await callback( request=_mock_callback_request(), code="test_code", - state="invalid_encrypted_state", + state="relay-handle-without-cookie", ) - # Should return HTML error page - assert response.status_code == 200 - assert "Authentication incomplete" in response.body.decode() + assert response.status_code == 400 + matching = [r.getMessage() for r in caplog.records if "could not decode OAuth state" in r.getMessage()] + assert len(matching) == 1 + assert "state_cookie_present=False" in matching[0] + assert "request_base_url=http://localhost:3000" in matching[0] + assert "PROXY_BASE_URL_set=False" in matching[0] + + +@pytest.mark.asyncio +async def test_oauth_callback_undecodable_state_with_cookie_present(): + """When the state cookie DID arrive but still does not decode, the error + page must not claim the cookie was lost in transit.""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _oauth_state_cookie_name, + callback, + ) + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + state = "relay-handle" + request = _mock_callback_request() + request.cookies = {_oauth_state_cookie_name(state): "garbage"} + + response = await callback( + request=request, + code="test_code", + state=state, + ) + + assert response.status_code == 400 + body = response.body.decode() + assert "could not be decoded" in body + assert "did not arrive" not in body @pytest.mark.asyncio