mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
493ff1a6d1
commit
bd0a7d6113
3 changed files with 18 additions and 37 deletions
|
|
@ -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}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue