mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
f2d20baf9d
commit
d3b9c735b7
3 changed files with 67 additions and 6 deletions
|
|
@ -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 "",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue