fix: replace pseudo-HMAC state token with secrets.token_urlsafe, fix comment, document cache TTL trade-off

- _make_state_token now uses secrets.token_urlsafe(32) instead of an
  HMAC digest that was never re-verified by the callback (dict lookup
  only), making the intent explicit and removing misleading crypto.
- Fix misleading comment on managed MCP tool call: original_tool_name
  is the unprefixed name, not the "full (potentially prefixed)" name.
- Document the known stale-token window trade-off in _BYOK_CRED_CACHE_TTL
  so future contributors understand the failure mode.
- Update state-token tests to reflect the new implementation.
This commit is contained in:
Ishaan Jaffer 2026-03-07 10:06:12 -08:00
parent 685483e0ec
commit 9c653fba9c
3 changed files with 23 additions and 21 deletions

View file

@ -10,12 +10,9 @@ Endpoints:
GET /v1/mcp/server/{server_id}/oauth2/status — check if user is connected
"""
import base64
import hashlib
import hmac
import html as _html_module
import json
import os
import secrets
import time
from typing import Dict, Optional
from urllib.parse import parse_qs, urlencode
@ -70,16 +67,16 @@ def _purge_expired_states() -> None:
def _make_state_token(server_id: str, user_id: str, timestamp: float, master_key: str) -> str:
"""HMAC-SHA256 of '{server_id}:{user_id}:{timestamp}:{nonce}' base64url-encoded.
"""Return a cryptographically random opaque state token.
A 16-byte random nonce is appended so that concurrent flows for the same
(server_id, user_id) pair — e.g. from multiple open browser tabs — each
produce a unique state token instead of colliding in _pending_oauth2_states.
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.
"""
nonce = base64.urlsafe_b64encode(os.urandom(16)).rstrip(b"=").decode()
message = f"{server_id}:{user_id}:{timestamp}:{nonce}".encode()
digest = hmac.new(master_key.encode(), message, hashlib.sha256).digest()
return base64.urlsafe_b64encode(digest).rstrip(b"=").decode()
return secrets.token_urlsafe(32)
def _build_success_html(server_name: str) -> str:

View file

@ -61,6 +61,11 @@ from litellm.utils import Rules, client, function_setup
# Keyed by (user_id, server_id); value is (credential_or_None, monotonic_timestamp).
# Storing the credential value (not just a bool) means _get_byok_credential and
# _check_byok_credential share a single DB round-trip per TTL window.
#
# Known trade-off: OAuth2 provider tokens may expire before the TTL elapses, causing
# a brief 401 window (up to _BYOK_CRED_CACHE_TTL seconds) after a token expires at the
# provider before the cache entry is evicted and a fresh DB lookup is made. Mitigating
# this (e.g. by storing and checking `expires_in`) is a future improvement.
_byok_cred_cache: Dict[Tuple[str, str], Tuple[Optional[str], float]] = {}
_BYOK_CRED_CACHE_TTL = 60 # seconds
_BYOK_CRED_CACHE_MAX_SIZE = 4096 # cap to prevent unbounded growth
@ -1859,7 +1864,7 @@ if MCP_AVAILABLE:
elif mcp_server:
response = await _handle_managed_mcp_tool(
server_name=server_name,
name=original_tool_name, # Pass the full name (potentially prefixed)
name=original_tool_name, # Pass the unprefixed tool name to the managed server
arguments=arguments,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,

View file

@ -28,19 +28,19 @@ def test_make_state_token_returns_string():
assert len(token) > 0
def test_make_state_token_is_unique_for_same_inputs():
"""Two calls with identical inputs must produce different tokens (random nonce)."""
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")
assert t1 != t2, "Tokens should differ due to random nonce"
assert t1 != t2, "Tokens should differ on each call"
def test_make_state_token_differs_by_server():
ts = time.time()
t1 = _make_state_token("server1", "user1", ts, "master-key")
t2 = _make_state_token("server2", "user1", ts, "master-key")
assert t1 != t2
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")
# secrets.token_urlsafe(32) produces at least 43 url-safe characters
assert len(token) >= 32
# ---------------------------------------------------------------------------