Merge pull request #11 from BerriAI/litellm_session_token_1_101_x

refactor(auth): bind UI/CLI session tokens to their own AES-GCM context (stable/1.101.x)
This commit is contained in:
yuneng-jiang 2026-09-29 15:39:12 -07:00 • committed by GitHub
commit 1b6ee737c9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 236 additions and 29 deletions

View file

@ -3313,13 +3313,16 @@ async def get_org_object_by_alias(
)
LITELLM_SESSION_TOKEN_PREFIX: Final = "litellm_login_"
class ExperimentalUIJWTToken:
@staticmethod
def get_experimental_ui_login_jwt_auth_token(user_info: LiteLLM_UserTable) -> str:
from datetime import timedelta
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
encrypt_bearer_token,
)
if user_info.user_role is None:
@ -3345,7 +3348,7 @@ class ExperimentalUIJWTToken:
user_role=LitellmUserRoles(user_info.user_role),
)
return encrypt_value_helper(valid_token.model_dump_json(exclude_none=True))
return encrypt_bearer_token(valid_token.model_dump_json(exclude_none=True), prefix=LITELLM_SESSION_TOKEN_PREFIX)
@staticmethod
def get_cli_jwt_auth_token(
@ -3376,7 +3379,7 @@ class ExperimentalUIJWTToken:
from datetime import timedelta
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
encrypt_bearer_token,
)
if user_info.user_role is None:
@ -3414,7 +3417,7 @@ class ExperimentalUIJWTToken:
is_session_token=True,
)
return encrypt_value_helper(valid_token.model_dump_json(exclude_none=True))
return encrypt_bearer_token(valid_token.model_dump_json(exclude_none=True), prefix=LITELLM_SESSION_TOKEN_PREFIX)
@staticmethod
def get_key_object_from_ui_hash_key(
@ -3424,10 +3427,10 @@ class ExperimentalUIJWTToken:
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
decrypt_bearer_token,
)
decrypted_token: Final = decrypt_value_helper(hashed_token, key="ui_hash_key", exception_type="debug")
decrypted_token: Final = decrypt_bearer_token(hashed_token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
if decrypted_token is None:
return None
try:

View file

@ -69,26 +69,55 @@ def _derive_key(signing_key: str) -> bytes:
return hashlib.sha256(signing_key.encode()).digest()
def _encrypt_aes_gcm(value: str, signing_key: str) -> str:
"""Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string."""
def _seal_aes_gcm(value: str, signing_key: str, aad: bytes | None) -> bytes:
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
nonce: Final = os.urandom(12)
# AESGCM.encrypt returns ciphertext || tag(16); wire format is nonce || that.
blob: Final = AESGCM(_derive_key(signing_key)).encrypt(nonce, value.encode("utf-8"), None)
return _V2_GCM_PREFIX + base64.urlsafe_b64encode(nonce + blob).decode("utf-8")
return nonce + AESGCM(_derive_key(signing_key)).encrypt(nonce, value.encode("utf-8"), aad)
def _open_aes_gcm(sealed: bytes, signing_key: str, aad: bytes | None) -> str:
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
# An empty plaintext still serializes to nonce(12) || tag(16) = 28 bytes, so a
# short/empty buffer here is a corrupt value: let AESGCM.decrypt raise and be
# swallowed by the caller (returns None/original), same as legacy.
return AESGCM(_derive_key(signing_key)).decrypt(sealed[:12], sealed[12:], aad).decode("utf-8")
def _encrypt_aes_gcm(value: str, signing_key: str) -> str:
"""Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string."""
sealed: Final = _seal_aes_gcm(value=value, signing_key=signing_key, aad=None)
return _V2_GCM_PREFIX + base64.urlsafe_b64encode(sealed).decode("utf-8")
def _decrypt_aes_gcm(value: str, signing_key: str) -> str:
"""Decrypt a versioned ``v2:gcm:`` string produced by :func:`_encrypt_aes_gcm`."""
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
sealed: Final = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :])
return _open_aes_gcm(sealed=sealed, signing_key=signing_key, aad=None)
raw: Final = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :])
# An empty plaintext still serializes to nonce(12) || tag(16) = 28 bytes, so a
# short/empty buffer here is a corrupt value: let AESGCM.decrypt raise and be
# swallowed by decrypt_value_helper (returns None/original), same as legacy.
nonce, blob = raw[:12], raw[12:]
return AESGCM(_derive_key(signing_key)).decrypt(nonce, blob, None).decode("utf-8")
def encrypt_bearer_token(value: str, prefix: str) -> str:
"""AES-256-GCM as unpadded base64url behind ``prefix``, which is also the AAD so a token can't change kind."""
salt_key: Final = _get_salt_key()
if not isinstance(salt_key, str):
raise ValueError("Set LITELLM_SALT_KEY or a master key to mint bearer tokens") # noqa: TRY004 # missing config, not a bad argument type
sealed: Final = _seal_aes_gcm(value=value, signing_key=salt_key, aad=prefix.encode("utf-8"))
return prefix + base64.urlsafe_b64encode(sealed).decode("ascii").rstrip("=")
def decrypt_bearer_token(token: str, prefix: str) -> str | None:
"""None unless ``token`` came from :func:`encrypt_bearer_token` with the same ``prefix``."""
salt_key: Final = _get_salt_key()
if not isinstance(salt_key, str) or not token.startswith(prefix):
return None
encoded: Final = token.removeprefix(prefix)
try:
sealed: Final = base64.b64decode(encoded + "=" * (-len(encoded) % 4), altchars=b"-_", validate=True)
return _open_aes_gcm(sealed=sealed, signing_key=salt_key, aad=prefix.encode("utf-8"))
except Exception: # noqa: BLE001 # base64 and AES-GCM each raise their own "not a token" type
return None
def encrypt_value_helper(value: str, new_encryption_key: str | None = None):

View file

@ -48,3 +48,6 @@
- {id: other.a2a.message_send.bridge_invokes, module: other, tier: P1, area: a2a, assertions: [bridge_invokes], source: "a2a_protocol/litellm_completion_bridge/handler.py", rationale: "A2A message/send routes through the completion bridge to a real provider and logs an asend_message spend row"}
- {id: other.a2a.version.serves_pinned_0_3, module: other, tier: P1, area: a2a, assertions: [serves_pinned_0_3], source: "agent_endpoints/a2a_endpoints.py _served_version", rationale: "An agent pinning 0.3 returns the flat 0.3 message shape (parts on the result)"}
- {id: other.a2a.version.serves_pinned_1_0, module: other, tier: P1, area: a2a, assertions: [serves_pinned_1_0], source: "agent_endpoints/a2a_endpoints.py _served_version", rationale: "An agent pinning 1.0 returns the nested 1.0 message shape (result.message with ROLE_AGENT)"}
- {id: other.auth.session_token.valid_allows, module: other, tier: P0, area: auth, assertions: [valid_allows], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An unexpired LiteLLM-minted session token authenticates with the role it carries"}
- {id: other.auth.session_token.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "auth/user_api_key_auth.py expiry check", rationale: "An expired session token is rejected with the expired-key error"}
- {id: other.auth.session_token.encrypted_value_denied, module: other, tier: P0, area: auth, assertions: [encrypted_value_denied], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An encrypted value read back from a management route is not accepted as a bearer token"}

View file

@ -0,0 +1,91 @@
"""Live e2e: UI/CLI session tokens are accepted only while valid and only when minted as session tokens.
The runner mints its own session tokens under the proxy's salt key, so the valid and expired cases run in
seconds instead of waiting out a real login's expiry.
"""
from __future__ import annotations
import base64
import hashlib
import json
import os
from datetime import datetime, timedelta, timezone
from typing import Final
import pytest
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
from e2e_config import MASTER_KEY, unique_marker
from e2e_http import UnauthorizedError, unwrap
from lifecycle import ResourceManager
from models import KeyGenerateBody, KeyLoggingCallback, KeyLoggingCallbackVars, KeyMetadata
from other_client import OtherClient
pytestmark = pytest.mark.e2e
SALT_KEY: Final = os.environ.get("LITELLM_SALT_KEY") or MASTER_KEY
SESSION_TOKEN_PREFIX: Final = "litellm_login_"
ENCRYPTED_PREFIX: Final = "litellm_enc::"
def _admin_session_token(expires_at: datetime) -> str:
claims: Final = json.dumps(
{
"token": f"ui-token-{unique_marker()}",
"user_id": f"e2e-session-{unique_marker()}",
"user_role": "proxy_admin",
"team_id": "litellm-dashboard",
"expires": expires_at.isoformat(),
}
)
nonce: Final = os.urandom(12)
sealed: Final = AESGCM(hashlib.sha256(SALT_KEY.encode()).digest()).encrypt(
nonce, claims.encode(), SESSION_TOKEN_PREFIX.encode()
)
return SESSION_TOKEN_PREFIX + base64.urlsafe_b64encode(nonce + sealed).decode().rstrip("=")
class TestSessionToken:
@pytest.mark.covers("other.auth.session_token.valid_allows")
def test_unexpired_session_token_reaches_admin_route(self, client: OtherClient) -> None:
token: Final = _admin_session_token(datetime.now(timezone.utc) + timedelta(minutes=10))
listing: Final = unwrap(client.list_users_as(token))
assert listing.total >= 0, f"an unexpired admin session token did not reach /user/list: {listing}"
@pytest.mark.covers("other.auth.session_token.expired_denied")
def test_expired_session_token_is_denied(self, client: OtherClient) -> None:
token: Final = _admin_session_token(datetime.now(timezone.utc) - timedelta(minutes=1))
result: Final = client.list_users_as(token)
assert isinstance(result, UnauthorizedError), f"an expired session token must get 401, got {result}"
assert "expired" in result.body.lower(), f"expected the expired-key error, got {result.body[:300]}"
@pytest.mark.covers("other.auth.session_token.encrypted_value_denied")
def test_encrypted_stored_value_is_not_a_bearer_token(
self, client: OtherClient, resources: ResourceManager
) -> None:
stored_value: Final = f'{{"token": "{unique_marker()}", "user_role": "proxy_admin"}}'
key: Final = client.proxy.generate_key(
KeyGenerateBody(
key_alias=f"e2e-session-{unique_marker()}",
metadata=KeyMetadata(
logging=[
KeyLoggingCallback(
callback_name="langfuse",
callback_vars=KeyLoggingCallbackVars(langfuse_secret_key=stored_value),
)
]
),
)
)
resources.defer(lambda: client.proxy.delete_key(key))
metadata: Final = client.proxy.key_info(key).metadata
assert metadata is not None and metadata.logging, f"/key/info dropped the logging metadata: {metadata}"
encrypted: Final = metadata.logging[0].callback_vars.langfuse_secret_key
assert encrypted is not None and encrypted.startswith(ENCRYPTED_PREFIX), (
f"expected /key/info to return the stored secret encrypted, got {encrypted!r}"
)
for bearer in (encrypted.removeprefix(ENCRYPTED_PREFIX), encrypted):
result = client.list_users_as(bearer)
assert isinstance(result, UnauthorizedError), f"an encrypted stored value must get 401, got {result}"

View file

@ -1,5 +1,7 @@
import asyncio
import base64
import json
import re
from types import SimpleNamespace
from typing import TYPE_CHECKING, Optional
from unittest.mock import AsyncMock, MagicMock, patch
@ -31,6 +33,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import (
LITELLM_SESSION_TOKEN_PREFIX,
ExperimentalUIJWTToken,
_cache_management_object,
_can_object_call_model,
@ -58,7 +61,9 @@ from litellm.constants import (
REGISTRY_ERROR_NEGATIVE_CACHE_TTL,
TAG_REGISTRY_MAX_SIZE,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.proxy.auth.user_api_key_auth import check_api_key_for_custom_headers_or_pass_through_endpoints
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_bearer_token, encrypt_value_helper
from litellm.proxy.common_utils.user_api_key_cache import (
END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL,
TAG_REGISTRY_OVERFLOW_SENTINEL,
@ -129,7 +134,7 @@ def test_get_experimental_ui_login_jwt_auth_token_valid(valid_sso_user_defined_v
token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(valid_sso_user_defined_values)
# Decrypt and verify token contents
decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted_token = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
# Check that decrypted_token is not None before using json.loads
assert decrypted_token is not None
token_data = json.loads(decrypted_token)
@ -155,7 +160,7 @@ def test_get_cli_jwt_auth_token_includes_team_alias(valid_sso_user_defined_value
team_alias="test-team",
)
decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted_token = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
assert decrypted_token is not None
token_data = json.loads(decrypted_token)
@ -182,7 +187,7 @@ def test_get_cli_jwt_auth_token_carries_team_grants_not_user_allowlist(
team_model_aliases={"team-fast": "gpt-4.1-mini"},
)
decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted_token = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
assert decrypted_token is not None
token_data = json.loads(decrypted_token)
@ -199,7 +204,7 @@ def test_get_cli_jwt_auth_token_keeps_user_allowlist_when_no_team(
"""A session token with no team bound still carries the user's own allowlist."""
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted_token = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
assert decrypted_token is not None
token_data = json.loads(decrypted_token)
@ -213,7 +218,7 @@ def test_get_experimental_ui_login_jwt_auth_token_uses_10_min_expiry(
):
"""Test that Experimental UI token uses fixed 10-minute expiry (does not use LITELLM_UI_SESSION_DURATION)."""
token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(valid_sso_user_defined_values)
decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted_token = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
assert decrypted_token is not None
token_data = json.loads(decrypted_token)
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
@ -231,7 +236,7 @@ def test_experimental_ui_token_ignores_litellm_ui_session_duration(
was incorrectly wired to the experimental flow."""
# Default LITELLM_UI_SESSION_DURATION is "24h" - token must still expire in ~10 min
token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(valid_sso_user_defined_values)
decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted_token = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
assert decrypted_token is not None
token_data = json.loads(decrypted_token)
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
@ -268,6 +273,51 @@ def test_get_key_object_from_ui_hash_key_valid(valid_sso_user_defined_values, mo
assert key_object.max_budget == litellm.max_ui_session_budget
@pytest.mark.parametrize("encryption_algorithm", ["xsalsa20-poly1305", "aes-256-gcm"])
def test_get_key_object_from_ui_hash_key_accepts_only_minted_session_tokens(
valid_sso_user_defined_values, monkeypatch, encryption_algorithm
):
monkeypatch.setattr(proxy_server, "general_settings", {"encryption_algorithm": encryption_algorithm})
session_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
stored_value = encrypt_value_helper(json.dumps({"user_role": LitellmUserRoles.PROXY_ADMIN.value}))
key_object = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(session_token)
assert key_object is not None
assert key_object.user_role == LitellmUserRoles.PROXY_ADMIN
reshaped = LITELLM_SESSION_TOKEN_PREFIX + stored_value.removeprefix("v2:gcm:").rstrip("=")
for candidate in (stored_value, reshaped):
assert ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(candidate) is None
def test_session_tokens_are_header_safe_and_never_look_like_virtual_keys(valid_sso_user_defined_values):
for token in (
ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(valid_sso_user_defined_values),
ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values),
):
assert re.fullmatch(r"litellm_login_[A-Za-z0-9_-]+", token), token
assert ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(token) is not None
@pytest.mark.asyncio
async def test_session_token_survives_langfuse_basic_auth_parsing(valid_sso_user_defined_values):
session_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
basic_credentials = base64.b64encode(f"{session_token}:sk-lf-secret".encode()).decode()
request = MagicMock()
request.headers = {}
api_key = await check_api_key_for_custom_headers_or_pass_through_endpoints(
request=request,
route="/api/public/ingestion",
pass_through_endpoints=[
{"path": "/api/public/ingestion", "target": "https://example.com", "custom_auth_parser": "langfuse"}
],
api_key=f"Basic {basic_credentials}",
)
assert api_key == session_token
assert ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(session_token) is not None
def test_get_key_object_from_ui_hash_key_invalid():
"""Test getting key object from invalid UI hash key"""
# Test with invalid token
@ -609,7 +659,7 @@ def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
# Decrypt and verify token contents
decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted_token = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
assert decrypted_token is not None
token_data = json.loads(decrypted_token)
@ -649,7 +699,7 @@ def test_get_cli_jwt_auth_token_custom_expiration(valid_sso_user_defined_values,
token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
# Decrypt and verify token contents
decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted_token = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
assert decrypted_token is not None
token_data = json.loads(decrypted_token)
@ -667,7 +717,7 @@ def test_get_cli_jwt_auth_token_unique_per_session(valid_sso_user_defined_values
from litellm.constants import CLI_SESSION_KEY_PREFIX
def _decode(token: str) -> dict:
decrypted = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
assert decrypted is not None
return json.loads(decrypted)
@ -687,7 +737,7 @@ def test_get_cli_jwt_auth_token_applies_fallback_budget(valid_sso_user_defined_v
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
valid_sso_user_defined_values, max_budget=litellm.max_ui_session_budget
)
decrypted = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
assert decrypted is not None
assert json.loads(decrypted).get("max_budget") == litellm.max_ui_session_budget
@ -696,7 +746,7 @@ def test_get_cli_jwt_auth_token_no_fallback_when_budget_provided(
valid_sso_user_defined_values,
):
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values, max_budget=None)
decrypted = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
decrypted = decrypt_bearer_token(token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
assert decrypted is not None
assert json.loads(decrypted).get("max_budget") is None

View file

@ -6,12 +6,16 @@ gate, and the backward-compatibility guarantees that let legacy XSalsa20-Poly130
(nacl) ciphertext and new AES values coexist and decrypt correctly.
"""
import re
import pytest
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
_V2_GCM_PREFIX,
decrypt_bearer_token,
decrypt_value_helper,
encrypt_bearer_token,
encrypt_value_helper,
)
@ -185,3 +189,30 @@ def test_decrypt_failure_debug_log_omits_raw_value(monkeypatch):
"the failing key should still be named in the breadcrumb"
)
assert result == secret
def test_bearer_token_opens_only_under_its_own_prefix():
token = encrypt_bearer_token("session", prefix="kind_a_")
relabeled = "kind_b_" + token.removeprefix("kind_a_")
assert decrypt_bearer_token(token, prefix="kind_a_") == "session"
assert decrypt_bearer_token(token, prefix="kind_b_") is None
assert decrypt_bearer_token(relabeled, prefix="kind_b_") is None
@pytest.mark.parametrize("use_aes", [False, True])
def test_stored_value_is_not_a_bearer_token_even_when_reshaped(monkeypatch, use_aes: bool):
if use_aes:
_use_aes(monkeypatch)
stored = encrypt_value_helper("stored-secret")
for candidate in (stored, "kind_a_" + stored.removeprefix(_V2_GCM_PREFIX).rstrip("=")):
assert decrypt_bearer_token(candidate, prefix="kind_a_") is None
@pytest.mark.parametrize("length", range(6))
def test_bearer_token_uses_only_header_safe_characters(length: int):
token = encrypt_bearer_token("x" * length, prefix="kind_a_")
assert re.fullmatch(r"kind_a_[A-Za-z0-9_-]+", token), token
assert decrypt_bearer_token(token, prefix="kind_a_") == "x" * length