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
This commit is contained in:
Ishaan Jaffer 2026-03-07 15:24:00 -08:00
parent 493ff1a6d1
commit bd0a7d6113
3 changed files with 18 additions and 37 deletions

View file

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

View file

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

View file

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