mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
refactor(auth): bind UI/CLI session tokens to their own AES-GCM context
Backport of BerriAI/litellm-private#5 (12981f93d3) onto stable/1.102.x, applied as the PR's net diff so main-only intermediate refactors stay out. The dashboard and lite CLI SSO specs under tests/e2e/ui/oidc are left out because this line has no OIDC e2e harness to run them.
This commit is contained in:
parent
5f8561bdc2
commit
1d5fdc87fe
6 changed files with 236 additions and 29 deletions
|
|
@ -3327,13 +3327,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:
|
||||
|
|
@ -3359,7 +3362,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(
|
||||
|
|
@ -3390,7 +3393,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:
|
||||
|
|
@ -3428,7 +3431,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(
|
||||
|
|
@ -3438,10 +3441,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:
|
||||
|
|
|
|||
|
|
@ -72,26 +72,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")
|
||||
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):
|
||||
|
|
|
|||
|
|
@ -54,3 +54,6 @@
|
|||
|
||||
- {id: other.auth.jwt.wrong_issuer_denied, module: other, tier: P0, area: auth, assertions: [wrong_issuer_denied], source: "auth/handle_jwt.py", rationale: "A signed token with the correct audience and an unexpected issuer is rejected"}
|
||||
- {id: other.auth.jwt.wrong_audience_denied, module: other, tier: P0, area: auth, assertions: [wrong_audience_denied], source: "auth/handle_jwt.py", rationale: "A signed token from the trusted issuer intended for another app is rejected"}
|
||||
- {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"}
|
||||
|
|
|
|||
91
tests/e2e/other/test_session_token_e2e.py
Normal file
91
tests/e2e/other/test_session_token_e2e.py
Normal 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}"
|
||||
|
|
@ -1,5 +1,7 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import re
|
||||
from types import SimpleNamespace
|
||||
from typing import TYPE_CHECKING, Final, Literal, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -32,6 +34,7 @@ from litellm.proxy._types import (
|
|||
WebhookEvent,
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
LITELLM_SESSION_TOKEN_PREFIX,
|
||||
ExperimentalUIJWTToken,
|
||||
_cache_management_object,
|
||||
_can_object_call_model,
|
||||
|
|
@ -62,7 +65,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,
|
||||
|
|
@ -133,7 +138,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)
|
||||
|
|
@ -159,7 +164,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)
|
||||
|
||||
|
|
@ -186,7 +191,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)
|
||||
|
||||
|
|
@ -203,7 +208,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)
|
||||
|
||||
|
|
@ -217,7 +222,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"))
|
||||
|
|
@ -235,7 +240,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"))
|
||||
|
|
@ -272,6 +277,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
|
||||
|
|
@ -650,7 +700,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)
|
||||
|
||||
|
|
@ -690,7 +740,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)
|
||||
|
||||
|
|
@ -708,7 +758,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)
|
||||
|
||||
|
|
@ -728,7 +778,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
|
||||
|
||||
|
|
@ -737,7 +787,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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue