litellm/tests/test_litellm/proxy/auth/test_litellm_license.py

128 lines
5.9 KiB
Python

import asyncio
import json
from unittest.mock import AsyncMock, MagicMock, patch
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPublicKey
from litellm.proxy.auth.litellm_license import LicenseCheck
def test_read_public_key_loads_successfully():
"""Ensure public_key.pem is valid PEM with no leading whitespace."""
license_check = LicenseCheck()
assert (
license_check.public_key is not None
), "public_key.pem could not be loaded — check for leading whitespace or malformed PEM header"
def test_is_over_limit():
license_check = LicenseCheck()
license_check.airgapped_license_data = {"max_users": 100}
assert license_check.is_over_limit(101) is True
assert license_check.is_over_limit(100) is False
assert license_check.is_over_limit(99) is False
license_check.airgapped_license_data = {}
assert license_check.is_over_limit(101) is False
assert license_check.is_over_limit(100) is False
assert license_check.is_over_limit(99) is False
license_check.airgapped_license_data = None
assert license_check.is_over_limit(101) is False
assert license_check.is_over_limit(100) is False
assert license_check.is_over_limit(99) is False
def test_auto_router_capability_limit() -> None:
"""The signed license's auto_router feature or its "*" wildcard lifts the one-router limit; an
API-verified license (no airgapped data) and an airgapped license without either keep it."""
license_check = LicenseCheck()
license_check.airgapped_license_data = {"expiration_date": "2999-01-01", "allowed_features": ["auto_router"]}
assert license_check.auto_router_capability_limit() is None
license_check.airgapped_license_data = {
"expiration_date": "2999-01-01",
"allowed_features": ["sso", "auto_router", "audit_logs"],
}
assert license_check.auto_router_capability_limit() is None
license_check.airgapped_license_data = {"expiration_date": "2999-01-01", "allowed_features": ["*"]}
assert license_check.auto_router_capability_limit() is None
license_check.airgapped_license_data = {"expiration_date": "2999-01-01", "allowed_features": ["sso", "*"]}
assert license_check.auto_router_capability_limit() is None
license_check.airgapped_license_data = {"expiration_date": "2999-01-01", "allowed_features": ["sso"]}
assert license_check.auto_router_capability_limit() == 1
license_check.airgapped_license_data = {"expiration_date": "2999-01-01", "allowed_features": "*"}
assert license_check.auto_router_capability_limit() is None
license_check.airgapped_license_data = {"expiration_date": "2999-01-01"}
assert license_check.auto_router_capability_limit() == 1
license_check.airgapped_license_data = None
assert license_check.auto_router_capability_limit() == 1
def _signed_license(
expiration_date: str, allowed_features: tuple[str, ...] = ("auto_router",)
) -> tuple[RSAPublicKey, str]:
import base64
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.asymmetric import padding, rsa
private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
message = json.dumps(
{"expiration_date": expiration_date, "user_id": "u", "allowed_features": list(allowed_features)}
).encode()
signature = private_key.sign(
message,
padding.PSS(mgf=padding.MGF1(hashes.SHA256()), salt_length=padding.PSS.MAX_LENGTH),
hashes.SHA256(),
)
return private_key.public_key(), base64.b64encode(message + b"." + signature).decode()
def test_expired_or_unreadable_license_grants_no_features() -> None:
"""The verifier stores the signed payload only after the expiry check passes and clears it when a
later verify rejects the license, so a stale payload cannot keep lifting the heuristic_v2 limit."""
license_check = LicenseCheck()
public_key, valid_key = _signed_license("2999-01-01")
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=valid_key) is True
assert license_check.auto_router_capability_limit() is None
_, expired_key = _signed_license("2000-01-01")
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=expired_key) is not True
assert license_check.airgapped_license_data is None
assert license_check.auto_router_capability_limit() == 1
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=valid_key) is True
assert license_check.verify_license_without_api_request(public_key=public_key, license_key="not-a-license") is not True
assert license_check.airgapped_license_data is None
def test_valid_signed_license_with_auto_router_lifts_the_limit() -> None:
license_check = LicenseCheck()
public_key, license_key = _signed_license("2999-01-01")
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=license_key) is True
assert license_check.auto_router_capability_limit() is None
def test_valid_signed_wildcard_license_lifts_the_limit() -> None:
"""The license generator defaults allowed_features to ["*"], meaning every feature, so a wildcard
license grants auto_router the same way a license that names it does."""
license_check = LicenseCheck()
public_key, license_key = _signed_license("2999-01-01", allowed_features=("*",))
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=license_key) is True
assert license_check.grants_feature("auto_router") is True
assert license_check.auto_router_capability_limit() is None
named_public_key, named_key = _signed_license("2999-01-01", allowed_features=("sso", "audit_logs"))
assert license_check.verify_license_without_api_request(public_key=named_public_key, license_key=named_key) is True
assert license_check.grants_feature("auto_router") is False
assert license_check.auto_router_capability_limit() == 1