From bd0a7d6113e3e9b60219c8a7c780766e87dd3855 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Sat, 7 Mar 2026 15:24:00 -0800 Subject: [PATCH] fix: remove "" cache sentinel, fix state-expiry ordering, fix fragile client_secret test - Remove "" sentinel write from status endpoint: only real token fetches (via _get_byok_credential / _check_byok_credential) write to cache. Status reads cache when real token present, falls through to DB otherwise. Eliminates dual-interpretation of same cache key across status/auth paths. - Revert _get_byok_credential and _check_byok_credential: no longer need special "" handling since "" can no longer appear in cache - Fix _purge_expired_states ordering in callback: check state first so expired tokens get "State expired" error, not misleading "Invalid state" - Fix test_connect_missing_client_secret_raises_400: replace dead-code master_key mock with prisma_client mock to make validation-order-independent --- .../mcp_server/openapi_oauth2_endpoints.py | 30 ++++++++----------- .../proxy/_experimental/mcp_server/server.py | 17 ++--------- .../test_openapi_oauth2_endpoints.py | 8 ++--- 3 files changed, 18 insertions(+), 37 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py index eda0f3ac294..246feae19b0 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_oauth2_endpoints.py @@ -291,7 +291,8 @@ async def openapi_oauth2_callback( # noqa: PLR0915 status_code=400, ) - _purge_expired_states() + # Check state before purging so expired tokens get a descriptive "State expired" + # error instead of a misleading "Invalid state" error. state_data = _pending_oauth2_states.get(state) if state_data is None: return HTMLResponse( @@ -303,6 +304,7 @@ async def openapi_oauth2_callback( # noqa: PLR0915 ) if time.time() > state_data["expires_at"]: del _pending_oauth2_states[state] + _purge_expired_states() # best-effort cleanup of other expired entries return HTMLResponse( content=_build_error_html( "State expired", @@ -311,8 +313,9 @@ async def openapi_oauth2_callback( # noqa: PLR0915 status_code=400, ) - # Consume the state (one-time use) + # Consume the state (one-time use), then do best-effort cleanup of other entries. del _pending_oauth2_states[state] + _purge_expired_states() server_id: str = state_data["server_id"] user_id: str = state_data["user_id"] @@ -555,10 +558,12 @@ async def openapi_oauth2_status( ) # 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. + # The cache is written by _get_byok_credential / _check_byok_credential when + # a tool call fetches the real token. If the token is already cached (i.e. + # the user has made at least one tool call), we can skip the DB entirely. + # We intentionally do NOT write to the cache here: writing a sentinel value + # would create dual interpretations of the same cache key across status and + # auth functions, making the cache contract subtle and fragile. try: import time as _time @@ -573,7 +578,7 @@ async def openapi_oauth2_status( if _time.monotonic() - ts < _BYOK_CRED_CACHE_TTL: return JSONResponse( { - "connected": cached_cred is not None, + "connected": bool(cached_cred), "server_id": server_id, "server_name": server_name, } @@ -596,17 +601,6 @@ async def openapi_oauth2_status( ) connected = False - # Populate cache so subsequent polls within the TTL window skip the DB. - # Use a non-None sentinel ("") for connected=True to satisfy the cache's - # None-means-no-credential invariant; invalidation via _invalidate_byok_cred_cache - # is still the authoritative signal when a new token is stored. - try: - from litellm.proxy._experimental.mcp_server.server import _write_byok_cred_cache - - _write_byok_cred_cache(user_id, server_id, "" if connected else None) - except Exception: - pass # Best-effort; never block the response - return JSONResponse( {"connected": connected, "server_id": server_id, "server_name": server_name} ) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index d94531eb919..2f2f84fef4c 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1582,12 +1582,7 @@ if MCP_AVAILABLE: if cached is not None: credential, ts = cached if time.monotonic() - ts < _BYOK_CRED_CACHE_TTL: - # Treat "" as a cache miss: the status endpoint may write an - # empty-string sentinel to record "connected=True" without a - # real token value. Fall through to the DB to fetch the actual - # credential rather than returning an empty Bearer header. - if credential is None or credential: - return credential + return credential from litellm.proxy._experimental.mcp_server.db import get_user_credential from litellm.proxy.proxy_server import prisma_client @@ -1654,11 +1649,6 @@ if MCP_AVAILABLE: ) # Check shared credential cache before hitting the DB. - # Note: the status endpoint writes "" as a sentinel for "connected per latest - # status poll". We treat "" identically to _get_byok_credential (cache miss) - # so that both functions always verify from DB when only the sentinel is present. - # This prevents a race where a deleted credential passes the auth check but - # returns None from _get_byok_credential, causing a silent 401 to the backend. cache_key = (user_id, mcp_server.server_id) cached = _byok_cred_cache.get(cache_key) if cached is not None: @@ -1680,10 +1670,7 @@ if MCP_AVAILABLE: "WWW-Authenticate": 'Bearer resource_metadata="/.well-known/oauth-protected-resource"' }, ) - # Only return early for real cached credentials; treat "" as a miss - # so we always verify against the DB when only the status sentinel exists. - if cached_cred: - return + return from litellm.proxy._experimental.mcp_server.db import get_user_credential from litellm.proxy.proxy_server import 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 79c99e64982..c03f209c1d5 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 @@ -104,13 +104,13 @@ async def test_connect_missing_client_secret_raises_400(): mock_server.client_id = "my-client-id" mock_server.client_secret = None # missing - # master_key is imported inline via `from litellm.proxy.proxy_server import master_key`; - # patch at the source so the function sees the mock value. + # Also mock prisma_client as non-None so the 400 assertion is robust to + # future validation reordering (without a DB mock the test would get 503). with patch( "litellm.proxy._experimental.mcp_server.openapi_oauth2_endpoints.global_mcp_server_manager" ) as mock_mgr, patch( - "litellm.proxy.proxy_server.master_key", - "sk-test", + "litellm.proxy.proxy_server.prisma_client", + MagicMock(), create=True, ): mock_mgr.get_mcp_server_by_id.return_value = mock_server