diff --git a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py index a3923dcf924..ab226fcfac7 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py @@ -66,15 +66,12 @@ def _purge_expired_states() -> None: del _pending_oauth2_states[k] -def _make_state_token(server_id: str, user_id: str, timestamp: float, master_key: str) -> str: +def _make_state_token() -> str: """Return a cryptographically random opaque state token. - 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. + The token is a dict key in `_pending_oauth2_states` — it is never used to + carry or verify signed data, so a 32-byte (256-bit) random value is + sufficient and makes the intent explicit. """ return secrets.token_urlsafe(32) @@ -187,7 +184,7 @@ async def openapi_oauth2_connect( raise HTTPException(status_code=503, detail="Too many pending OAuth2 flows") timestamp = time.time() - state = _make_state_token(server_id, user_id, timestamp, master_key) + state = _make_state_token() _pending_oauth2_states[state] = { "server_id": server_id, "user_id": user_id, @@ -212,7 +209,9 @@ async def openapi_oauth2_connect( if server.scopes: params["scope"] = " ".join(server.scopes) - authorization_url = f"{server.authorization_url}?{urlencode(params)}" + # Use "&" if the base URL already contains query parameters, otherwise "?" + sep = "&" if "?" in server.authorization_url else "?" + authorization_url = f"{server.authorization_url}{sep}{urlencode(params)}" server_name = server.server_name or server.name or server_id verbose_proxy_logger.debug( @@ -506,6 +505,33 @@ async def openapi_oauth2_status( {"connected": False, "server_id": server_id, "server_name": server_name} ) + # Check the shared credential cache before issuing a DB query. + # The frontend polls this endpoint every 2 s; the cache avoids a raw DB + # hit on each poll. The callback explicitly invalidates the cache entry + # via _invalidate_byok_cred_cache, so the next poll after a successful + # authorization always falls through to the DB and returns connected=True. + try: + import time as _time + + from litellm.proxy._experimental.mcp_server.server import ( + _BYOK_CRED_CACHE_TTL, + _byok_cred_cache, + ) + + cached = _byok_cred_cache.get((user_id, server_id)) + if cached is not None: + cached_cred, ts = cached + if _time.monotonic() - ts < _BYOK_CRED_CACHE_TTL: + return JSONResponse( + { + "connected": cached_cred is not None, + "server_id": server_id, + "server_name": server_name, + } + ) + except Exception: + pass # If cache import fails, fall through to DB + try: connected = await has_user_credential( prisma_client=prisma_client, 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 3ffed9b225f..62408d52bc8 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 @@ -23,22 +23,21 @@ from litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints import ( def test_make_state_token_returns_string(): - token = _make_state_token("server1", "user1", time.time(), "master-key") + token = _make_state_token() assert isinstance(token, str) assert len(token) > 0 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") + t1 = _make_state_token() + t2 = _make_state_token() assert t1 != t2, "Tokens should differ on each call" 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") + token = _make_state_token() # secrets.token_urlsafe(32) produces at least 43 url-safe characters assert len(token) >= 32