fix: fix double-? in auth URL, simplify state token signature, cache status endpoint

- Fix malformed authorization URL when base URL already has query params
  (e.g. Google: ?access_type=offline). Use & instead of ? in that case.
- Simplify _make_state_token() to take no parameters since none were used.
  Update call site and tests accordingly.
- Status endpoint now checks _byok_cred_cache before querying the DB,
  avoiding a raw Prisma query on every 2-second poll from the UI.
  The callback's _invalidate_byok_cred_cache call ensures the first poll
  after successful auth always falls through to DB and returns connected=True.
This commit is contained in:
Ishaan Jaffer 2026-03-07 10:45:34 -08:00
parent b5ce1ea310
commit eb4f63a655
2 changed files with 39 additions and 14 deletions

View file

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

View file

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