mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
b5ce1ea310
commit
eb4f63a655
2 changed files with 39 additions and 14 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue