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.
This commit is contained in:
Ishaan Jaffer 2026-03-07 11:33:31 -08:00
parent f2d20baf9d
commit d3b9c735b7
3 changed files with 67 additions and 6 deletions

View file

@ -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 "",

View file

@ -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()

View file

@ -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);