diff --git a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py index 74fa23c8cf7..a3923dcf924 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py @@ -10,12 +10,9 @@ Endpoints: GET /v1/mcp/server/{server_id}/oauth2/status — check if user is connected """ -import base64 -import hashlib -import hmac import html as _html_module import json -import os +import secrets import time from typing import Dict, Optional from urllib.parse import parse_qs, urlencode @@ -70,16 +67,16 @@ def _purge_expired_states() -> None: def _make_state_token(server_id: str, user_id: str, timestamp: float, master_key: str) -> str: - """HMAC-SHA256 of '{server_id}:{user_id}:{timestamp}:{nonce}' base64url-encoded. + """Return a cryptographically random opaque state token. - A 16-byte random nonce is appended so that concurrent flows for the same - (server_id, user_id) pair — e.g. from multiple open browser tabs — each - produce a unique state token instead of colliding in _pending_oauth2_states. + The token is looked up in `_pending_oauth2_states` by the callback — it is + never used to carry or verify signed data. A 32-byte (256-bit) random value + is therefore sufficient and makes the intent explicit. + + Parameters are accepted for API compatibility with existing tests; they are + not used in the generated token. """ - nonce = base64.urlsafe_b64encode(os.urandom(16)).rstrip(b"=").decode() - message = f"{server_id}:{user_id}:{timestamp}:{nonce}".encode() - digest = hmac.new(master_key.encode(), message, hashlib.sha256).digest() - return base64.urlsafe_b64encode(digest).rstrip(b"=").decode() + return secrets.token_urlsafe(32) def _build_success_html(server_name: str) -> str: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index af7fe97ad70..9f3b8196244 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -61,6 +61,11 @@ from litellm.utils import Rules, client, function_setup # Keyed by (user_id, server_id); value is (credential_or_None, monotonic_timestamp). # Storing the credential value (not just a bool) means _get_byok_credential and # _check_byok_credential share a single DB round-trip per TTL window. +# +# Known trade-off: OAuth2 provider tokens may expire before the TTL elapses, causing +# a brief 401 window (up to _BYOK_CRED_CACHE_TTL seconds) after a token expires at the +# provider before the cache entry is evicted and a fresh DB lookup is made. Mitigating +# this (e.g. by storing and checking `expires_in`) is a future improvement. _byok_cred_cache: Dict[Tuple[str, str], Tuple[Optional[str], float]] = {} _BYOK_CRED_CACHE_TTL = 60 # seconds _BYOK_CRED_CACHE_MAX_SIZE = 4096 # cap to prevent unbounded growth @@ -1859,7 +1864,7 @@ if MCP_AVAILABLE: elif mcp_server: response = await _handle_managed_mcp_tool( server_name=server_name, - name=original_tool_name, # Pass the full name (potentially prefixed) + name=original_tool_name, # Pass the unprefixed tool name to the managed server arguments=arguments, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, 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 9de8261545c..564386502db 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 @@ -28,19 +28,19 @@ def test_make_state_token_returns_string(): assert len(token) > 0 -def test_make_state_token_is_unique_for_same_inputs(): - """Two calls with identical inputs must produce different tokens (random nonce).""" +def test_make_state_token_is_unique(): + """Every call produces a unique token (cryptographically random).""" ts = time.time() t1 = _make_state_token("server1", "user1", ts, "master-key") t2 = _make_state_token("server1", "user1", ts, "master-key") - assert t1 != t2, "Tokens should differ due to random nonce" + assert t1 != t2, "Tokens should differ on each call" -def test_make_state_token_differs_by_server(): - ts = time.time() - t1 = _make_state_token("server1", "user1", ts, "master-key") - t2 = _make_state_token("server2", "user1", ts, "master-key") - assert t1 != t2 +def test_make_state_token_has_sufficient_entropy(): + """Token must be at least 32 url-safe characters (≥192 bits of entropy).""" + token = _make_state_token("server1", "user1", time.time(), "master-key") + # secrets.token_urlsafe(32) produces at least 43 url-safe characters + assert len(token) >= 32 # ---------------------------------------------------------------------------