litellm/tests/test_litellm/proxy/auth/test_litellm_license.py
tin-berri d0d09e5343
feat(router): meter auto-router tier and prompt customization against the auto_router license feature (#39674)
Generalizes the heuristic_v2 ceiling from #39468 into a capability table whose
records own their in-process predicate, SQL spelling and refusal wording. The
existing heuristic_v2 capability keeps its own one-router ceiling. A single
customization capability combines operator-defined tier definitions with every
operator-written part of the classifier prompt. The prompt half only applies to
classifier types that call an LLM. The shipped default prompt, classification
rubric presets, tier-label renames and tier model choices remain ungated.

Scope every enforcement point to actual complexity routers. A model-less PATCH
or legacy update now decrypts the stored model before accepting strategy-router
settings, so a regular model cannot acquire a router config or spend a license
slot. Under the existing advisory lock, the cross-pod candidate query returns
only model scalars and the count decrypts and classifies them in process; old
non-router rows carrying a capability-shaped config no longer block a real
complexity router. The signed auto_router license feature makes both ceilings
unlimited.
2026-09-05 09:51:23 -07:00

101 lines
4.4 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:
"""Only the signed license's auto_router feature lifts the one-router limit; an API-verified
license (no airgapped data) and an airgapped license without the feature 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": ["sso"]}
assert license_check.auto_router_capability_limit() == 1
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) -> 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": ["auto_router"]}
).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