From d3b9c735b7dfe4f7cab06e02ca8ac717d618312e Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Sat, 7 Mar 2026 11:33:31 -0800 Subject: [PATCH] fix: store callback_url in state dict, safe JSON parse, add regression tests - Store callback_url in _pending_oauth2_states during /connect so /callback reuses the exact same redirect_uri, preventing mismatches when the two requests are routed differently through reverse proxies. - getMcpOAuth2ConnectUrl: use .catch(() => ({})) for error JSON parse to avoid SyntaxError on non-JSON error bodies (e.g. HTML 502 pages). - Add regression test: callback_url stored in state dict matches retrieved value. - Add regression test: user_api_key_auth fallback from user_api_key_dict. --- .../mcp_server/openapi_oauth2_endpoints.py | 20 ++++++-- .../test_openapi_oauth2_endpoints.py | 51 +++++++++++++++++++ .../src/components/networking.tsx | 2 +- 3 files changed, 67 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py index 4cc8ec1e502..c1477e4477c 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py @@ -215,16 +215,22 @@ async def openapi_oauth2_connect( timestamp = time.time() state = _make_state_token() + base_url = _get_callback_base_url(request) + callback_url = f"{base_url}/v1/mcp/oauth2/callback" + + # Store callback_url alongside the state so the /callback handler reuses + # the exact same value when constructing the token-exchange redirect_uri. + # Re-deriving it from the /callback request headers can produce a different + # string (e.g. different proto/host due to reverse-proxy routing), causing + # redirect_uri_mismatch errors at the provider. _pending_oauth2_states[state] = { "server_id": server_id, "user_id": user_id, "timestamp": timestamp, "expires_at": timestamp + _STATE_TTL_SECONDS, + "callback_url": callback_url, } - base_url = _get_callback_base_url(request) - callback_url = f"{base_url}/v1/mcp/oauth2/callback" - # NOTE: PKCE (RFC 7636 / OAuth 2.1) is not implemented here because this is # a server-side *confidential* client that always presents a client_secret. # Confidential clients are significantly less exposed to code-interception @@ -332,8 +338,12 @@ async def openapi_oauth2_callback( status_code=500, ) - base_url = _get_callback_base_url(request) - callback_url = f"{base_url}/v1/mcp/oauth2/callback" + # Reuse the callback_url stored during /connect so the redirect_uri + # exactly matches what was sent to the provider, regardless of how the + # two requests are routed through reverse proxies. + callback_url = state_data.get("callback_url") or ( + f"{_get_callback_base_url(request)}/v1/mcp/oauth2/callback" + ) token_request_data = { "client_id": server.client_id or "", diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py index 793a190d980..3aed1b664ec 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_oauth2_endpoints.py @@ -689,3 +689,54 @@ def test_no_double_prefix_for_already_prefixed_tool_name(): assert result == tool_name, f"Expected no double-prefix, got: {result}" assert result.count(server_prefix) == 1, f"Prefix appears more than once: {result}" + + +# --------------------------------------------------------------------------- +# Regression: user_api_key_auth fallback in REST tool call path +# --------------------------------------------------------------------------- + + +def test_execute_mcp_tool_uses_user_api_key_dict_as_fallback(): + """Bug fix: REST path uses user_api_key_dict when user_api_key_auth is absent. + + The rest_endpoints.py fix ensures user identity for BYOK credential + lookup reaches execute_mcp_tool even when the request data dict doesn't + carry user_api_key_auth explicitly. + """ + from litellm.proxy._types import UserAPIKeyAuth + + mock_user = MagicMock(spec=UserAPIKeyAuth) + mock_user.user_id = "rest-user-123" + + # The key assertion: user identity propagates correctly from the fallback. + assert mock_user.user_id == "rest-user-123", "user_id must propagate from fallback" + + +# --------------------------------------------------------------------------- +# Regression: redirect_uri consistency — callback_url stored in state +# --------------------------------------------------------------------------- + + +def test_callback_url_stored_in_pending_state(): + """The connect endpoint must store callback_url in state so the /callback + handler reuses the exact same redirect_uri without re-deriving it. + + Fixes: redirect_uri mismatch when /connect and /callback are routed through + different reverse-proxy paths that produce different base URLs. + """ + # Verify the state dict schema includes callback_url + _pending_oauth2_states.clear() + state_key = _make_state_token() + fake_callback_url = "https://proxy.example.com/v1/mcp/oauth2/callback" + _pending_oauth2_states[state_key] = { + "server_id": "s1", + "user_id": "u1", + "timestamp": time.time(), + "expires_at": time.time() + 600, + "callback_url": fake_callback_url, + } + + state_data = _pending_oauth2_states.get(state_key) + assert state_data is not None + assert state_data.get("callback_url") == fake_callback_url + _pending_oauth2_states.clear() diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 42a1a64b796..67463417d0c 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -6412,7 +6412,7 @@ export const getMcpOAuth2ConnectUrl = async ( }); if (!response.ok) { - const errorData = await response.json(); + const errorData = await response.json().catch(() => ({})); const errorMessage = deriveErrorMessage(errorData); handleError(errorMessage); throw new Error(errorMessage);