fix(router): reconcile heuristic v2 entitlement

This commit is contained in:
Tin 2026-09-02 18:14:55 -07:00
parent 5fa015b2e5
commit 1f8f8c3f39
14 changed files with 190 additions and 26 deletions

View file

@ -1,9 +1,18 @@
ALTER TABLE "LiteLLM_ProxyModelTable"
ADD COLUMN IF NOT EXISTS "heuristic_v2_unlimited" BOOLEAN NOT NULL DEFAULT FALSE;
ADD COLUMN IF NOT EXISTS "heuristic_v2_unlimited" BOOLEAN,
ADD COLUMN IF NOT EXISTS "heuristic_v2_license_blocked" BOOLEAN;
-- Existing rows remain NULL so pre-gating duplicates cannot break this migration.
-- Proxy startup reconciles them atomically before loading database-backed routers.
ALTER TABLE "LiteLLM_ProxyModelTable"
ALTER COLUMN "heuristic_v2_unlimited" SET DEFAULT FALSE,
ALTER COLUMN "heuristic_v2_license_blocked" SET DEFAULT FALSE;
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_ProxyModelTable_one_heuristic_v2_router"
ON "LiteLLM_ProxyModelTable" ((1))
WHERE NOT "heuristic_v2_unlimited" AND CASE
WHERE "heuristic_v2_unlimited" IS FALSE
AND "heuristic_v2_license_blocked" IS FALSE
AND CASE
WHEN jsonb_typeof(litellm_params) = 'object'
THEN (litellm_params #>> '{complexity_router_config,classifier_type}') = 'heuristic_v2'
WHEN jsonb_typeof(litellm_params) = 'string'

View file

@ -53,7 +53,8 @@ model LiteLLM_ProxyModelTable {
litellm_params Json
model_info Json?
blocked Boolean @default(false)
heuristic_v2_unlimited Boolean @default(false)
heuristic_v2_unlimited Boolean? @default(false)
heuristic_v2_license_blocked Boolean? @default(false)
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")

View file

@ -4887,6 +4887,7 @@ class PrismaCompatibleUpdateDBModel(TypedDict, total=False):
model_info: str
blocked: bool
heuristic_v2_unlimited: bool # writable-ok: the model write path stamps this internal license marker
heuristic_v2_license_blocked: bool # writable-ok: proxy reconciliation stamps this internal routing marker
updated_at: str
updated_by: str

View file

@ -3,7 +3,7 @@
import base64
import json
import os
from datetime import datetime
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Final
import httpx
@ -155,6 +155,16 @@ class LicenseCheck:
license_data: Final = self.airgapped_license_data
if license_data is None:
return False
expiration_value: Final = license_data.get("expiration_date")
if not isinstance(expiration_value, str):
return False
try:
expiration_date: Final = datetime.strptime(expiration_value, "%Y-%m-%d").replace(tzinfo=timezone.utc)
except ValueError:
return False
if expiration_date < datetime.now(timezone.utc):
self.airgapped_license_data = None
return False
allowed_features: Final = license_data.get("allowed_features")
return isinstance(allowed_features, list) and feature in allowed_features
@ -188,19 +198,22 @@ class LicenseCheck:
# Decode and parse the data
license_data: Final = json.loads(message.decode())
self.airgapped_license_data = EnterpriseLicenseData(**license_data)
# debug information provided in license data
verbose_proxy_logger.debug("License data: %s", license_data)
# Check expiration date
expiration_date: Final = datetime.strptime(license_data["expiration_date"], "%Y-%m-%d")
if expiration_date < datetime.now():
expiration_date: Final = datetime.strptime(license_data["expiration_date"], "%Y-%m-%d").replace(
tzinfo=timezone.utc
)
if expiration_date < datetime.now(timezone.utc):
self.airgapped_license_data = None
return False, "License has expired"
self.airgapped_license_data = EnterpriseLicenseData(**license_data)
return True
except Exception as e:
self.airgapped_license_data = None
verbose_proxy_logger.debug(
"litellm.proxy.auth.litellm_license.py::verify_license_without_api_request - Unable to verify License locally. - %s",
e,

View file

@ -308,6 +308,29 @@ def _heuristic_v2_unlimited_marker(complexity_router_config: Mapping[str, object
return uses_heuristic_v2(complexity_router_config) and _license_allows_unlimited_heuristic_v2()
_FIND_HEURISTIC_V2_OWNER_SQL: Final = """
SELECT model_id, litellm_params
FROM "LiteLLM_ProxyModelTable"
WHERE "heuristic_v2_license_blocked" IS NOT TRUE
AND ($1::text IS NULL OR model_id <> $1::text)
AND CASE
WHEN jsonb_typeof(litellm_params) = 'object'
THEN (litellm_params #>> '{complexity_router_config,classifier_type}') = 'heuristic_v2'
WHEN jsonb_typeof(litellm_params) = 'string'
THEN (((litellm_params #>> '{}')::jsonb) #>> '{complexity_router_config,classifier_type}') = 'heuristic_v2'
ELSE FALSE
END
LIMIT 1
"""
async def _find_persisted_heuristic_v2_owner(
prisma_client: PrismaClient, current_model_id: str | None
) -> Sequence[Mapping[str, object]]:
rows: Final = await prisma_client.db.query_raw(_FIND_HEURISTIC_V2_OWNER_SQL, current_model_id)
return cast("Sequence[Mapping[str, object]]", rows) # cast-ok: query selects two known proxy-model columns
async def _raise_if_heuristic_v2_slot_taken(
*,
prisma_client: PrismaClient,
@ -335,8 +358,9 @@ async def _raise_if_heuristic_v2_slot_taken(
return
from litellm.proxy.proxy_server import llm_router
rows: Final = await _proxy_model_table(prisma_client).find_many(
where={} # mutable-ok: Prisma requires a mutable filter mapping
rows: Final = await _find_persisted_heuristic_v2_owner(
prisma_client=prisma_client,
current_model_id=current_model_id,
)
violation: Final = _heuristic_v2_slot_violation(
persisted_rows=rows,
@ -367,7 +391,7 @@ def _heuristic_v2_admin_violation(*, effective_config: Mapping[str, object] | No
def _heuristic_v2_slot_violation(
*,
persisted_rows: Sequence[_ProxyModelRow],
persisted_rows: Sequence[_ProxyModelRow | Mapping[str, object]],
live_model_list: Sequence[Mapping[str, object]] = (),
incoming_params: GenericLiteLLMParams | None,
existing_params: GenericLiteLLMParams | None,
@ -390,12 +414,14 @@ def _heuristic_v2_slot_violation(
"creating a database-backed heuristic_v2 router."
)
for row in persisted_rows:
if current_model_id is not None and row.model_id == current_model_id:
row_model_id: object = row.get("model_id") if isinstance(row, Mapping) else row.model_id
if current_model_id is not None and row_model_id == current_model_id:
continue
row_litellm_params: object = row.get("litellm_params") if isinstance(row, Mapping) else row.litellm_params
if uses_heuristic_v2(
config
if isinstance(
config := _litellm_params_mapping(row.litellm_params).get("complexity_router_config"), Mapping
config := _litellm_params_mapping(row_litellm_params).get("complexity_router_config"), Mapping
)
else None
):
@ -903,6 +929,7 @@ async def patch_model(
update_data["heuristic_v2_unlimited"] = _heuristic_v2_unlimited_marker(
_effective_complexity_router_config(patch_data.litellm_params, db_model.litellm_params)
)
update_data["heuristic_v2_license_blocked"] = False
# Perform partial update
updated_model: Final = await _proxy_model_table(prisma_client).update(
@ -1154,6 +1181,7 @@ async def _add_model_to_db(
"created_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
"updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
"heuristic_v2_unlimited": _heuristic_v2_unlimited_marker(model_params.litellm_params.complexity_router_config),
"heuristic_v2_license_blocked": False,
}
if model_params.model_info.id is not None:
_data["model_id"] = model_params.model_info.id
@ -2240,6 +2268,7 @@ async def update_model(
if isinstance(merged_dictionary.get("complexity_router_config"), Mapping)
else None
),
"heuristic_v2_license_blocked": False,
}
model_response: Final = await _proxy_model_table(prisma_client).update(
where={"model_id": _model_id},

View file

@ -1685,6 +1685,35 @@ class _UserTeamsRow(Protocol):
_ProxyModelRow: TypeAlias = "prisma_models.LiteLLM_ProxyModelTable"
_RECONCILE_HEURISTIC_V2_LICENSE_SQL: Final = """
WITH ranked_heuristic_v2 AS (
SELECT model_id,
ROW_NUMBER() OVER (ORDER BY created_at, model_id) AS owner_rank
FROM "LiteLLM_ProxyModelTable"
WHERE CASE
WHEN jsonb_typeof(litellm_params) = 'object'
THEN (litellm_params #>> '{complexity_router_config,classifier_type}') = 'heuristic_v2'
WHEN jsonb_typeof(litellm_params) = 'string'
THEN (((litellm_params #>> '{}')::jsonb) #>> '{complexity_router_config,classifier_type}') = 'heuristic_v2'
ELSE FALSE
END
), expected_state AS (
SELECT model_id,
$1::boolean AS unlimited,
CASE WHEN $1::boolean THEN FALSE ELSE owner_rank > 1 END AS license_blocked
FROM ranked_heuristic_v2
)
UPDATE "LiteLLM_ProxyModelTable" AS model
SET "heuristic_v2_unlimited" = expected.unlimited,
"heuristic_v2_license_blocked" = expected.license_blocked
FROM expected_state AS expected
WHERE model.model_id = expected.model_id
AND (
model."heuristic_v2_unlimited" IS DISTINCT FROM expected.unlimited
OR model."heuristic_v2_license_blocked" IS DISTINCT FROM expected.license_blocked
)
"""
def _config_param_table(client: PrismaClient | None) -> TableActions[_ConfigParamRow]:
return cast( # cast-ok: this is prisma's LiteLLM_Config actions object, which parses its Json column to a mapping
@ -6055,7 +6084,9 @@ class ProxyConfig:
if _id is not None:
model.model_info["id"] = _id
model.model_info["db_model"] = True
model.model_info["blocked"] = bool(getattr(model, "blocked", False))
model.model_info["blocked"] = bool(
getattr(model, "blocked", False) or getattr(model, "heuristic_v2_license_blocked", False)
)
if premium_user is True:
# seeing "created_at", "updated_at", "created_by", "updated_by" is a LiteLLM Enterprise Feature
@ -6255,6 +6286,8 @@ class ProxyConfig:
return
models_list: Final[list] = new_models if isinstance(new_models, list) else []
if llm_router is not None:
llm_router.allow_multiple_heuristic_v2 = _license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE)
if llm_router is None and master_key is not None:
verbose_proxy_logger.debug("len new_models: %s", len(models_list))
@ -6957,6 +6990,11 @@ class ProxyConfig:
def _should_load_db_object(self, object_type: str | SupportedDBObjectType) -> bool:
return should_load_db_object(object_type=object_type)
async def _reconcile_heuristic_v2_license_state(self, prisma_client: PrismaClient) -> None:
"""Apply the current signed-license entitlement before loading DB routers."""
unlimited: Final = _license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE)
await prisma_client.db.execute_raw(_RECONCILE_HEURISTIC_V2_LICENSE_SQL, unlimited)
async def _get_models_from_db(self, prisma_client: PrismaClient) -> Sequence[_ProxyModelRow] | None:
"""
Fetch all model deployments from the DB.
@ -7038,6 +7076,7 @@ class ProxyConfig:
# Only load models from DB if "models" is in supported_db_objects (or if supported_db_objects is not set)
if self._should_load_db_object(object_type="models"):
await self._reconcile_heuristic_v2_license_state(prisma_client=prisma_client)
new_models: Final = await self._get_models_from_db(prisma_client=prisma_client)
# update llm router

View file

@ -53,7 +53,8 @@ model LiteLLM_ProxyModelTable {
litellm_params Json
model_info Json?
blocked Boolean @default(false)
heuristic_v2_unlimited Boolean @default(false)
heuristic_v2_unlimited Boolean? @default(false)
heuristic_v2_license_blocked Boolean? @default(false)
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")

View file

@ -757,7 +757,7 @@ class Router:
self.set_verbose = set_verbose
self.ignore_invalid_deployments = ignore_invalid_deployments
self.fallback_access_check: Final = fallback_access_check
self.allow_multiple_heuristic_v2: Final = allow_multiple_heuristic_v2
self.allow_multiple_heuristic_v2 = allow_multiple_heuristic_v2
self.debug_level = debug_level
self.enable_pre_call_checks = enable_pre_call_checks
self.enable_tag_filtering = enable_tag_filtering

View file

@ -53,7 +53,8 @@ model LiteLLM_ProxyModelTable {
litellm_params Json
model_info Json?
blocked Boolean @default(false)
heuristic_v2_unlimited Boolean @default(false)
heuristic_v2_unlimited Boolean? @default(false)
heuristic_v2_license_blocked Boolean? @default(false)
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")

View file

@ -30,11 +30,21 @@ def test_is_over_limit():
def test_allows_feature_requires_signed_license_claim():
license_check = LicenseCheck()
license_check.airgapped_license_data = {"allowed_features": [AUTO_ROUTER_LICENSE_FEATURE]}
license_check.airgapped_license_data = {
"allowed_features": [AUTO_ROUTER_LICENSE_FEATURE],
"expiration_date": "2999-01-01",
}
assert license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE) is True
license_check.airgapped_license_data = {"allowed_features": ["other_feature"]}
license_check.airgapped_license_data = {"allowed_features": ["other_feature"], "expiration_date": "2999-01-01"}
assert license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE) is False
license_check.airgapped_license_data = {
"allowed_features": [AUTO_ROUTER_LICENSE_FEATURE],
"expiration_date": "2000-01-01",
}
assert license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE) is False
assert license_check.airgapped_license_data is None
license_check.airgapped_license_data = None
assert license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE) is False

View file

@ -3339,6 +3339,36 @@ class TestGetModelInfoWithIdBlocked:
info = ProxyConfig().get_model_info_with_id(model=model, db_model=True)
assert getattr(info, "blocked") is False
def test_get_model_info_with_id_propagates_license_blocked_true(self):
from litellm.proxy.proxy_server import ProxyConfig
model = MagicMock(spec=["model_id", "model_info", "blocked", "heuristic_v2_license_blocked"])
model.model_id = "dep-license-blocked"
model.model_info = {}
model.blocked = False
model.heuristic_v2_license_blocked = True
info = ProxyConfig().get_model_info_with_id(model=model, db_model=True)
assert getattr(info, "blocked") is True
class TestHeuristicV2LicenseReconciliation:
@pytest.mark.asyncio
@pytest.mark.parametrize("unlimited", [False, True])
async def test_reconcile_uses_current_license_state(self, unlimited):
from litellm.proxy.proxy_server import ProxyConfig
prisma = MagicMock()
prisma.db.execute_raw = AsyncMock(return_value=2)
with patch( # test-quality-ok: isolates the proxy license singleton used by reconciliation
"litellm.proxy.proxy_server._license_check.allows_feature",
return_value=unlimited,
):
await ProxyConfig()._reconcile_heuristic_v2_license_state(prisma_client=prisma)
query, entitlement = prisma.db.execute_raw.await_args.args
assert "ROW_NUMBER()" in query
assert entitlement is unlimited
class TestPatchModelBlockedAuthGate:
"""Only proxy admins may flip `blocked` — team admins authorized for
@ -4150,7 +4180,29 @@ class TestStrategyRouterWriteValidation:
existing_params=None,
)
prisma_client.db.litellm_proxymodeltable.find_many.assert_not_called()
prisma_client.db.query_raw.assert_not_called()
@pytest.mark.asyncio
async def test_singleton_lookup_uses_bounded_database_query(self):
from litellm.proxy.management_endpoints.model_management_endpoints import (
_find_persisted_heuristic_v2_owner,
)
prisma_client = MagicMock()
prisma_client.db.query_raw = AsyncMock(
return_value=[
{
"model_id": "owner",
"litellm_params": {"complexity_router_config": {"classifier_type": "heuristic_v2"}},
}
]
)
rows = await _find_persisted_heuristic_v2_owner(prisma_client, "current")
query, excluded_model_id = prisma_client.db.query_raw.await_args.args
assert len(rows) == 1
assert "LIMIT 1" in query
assert excluded_model_id == "current"
def test_double_prefix_rejected_against_stored_params(self):
from litellm.proxy.management_endpoints.model_management_endpoints import (

View file

@ -50,11 +50,22 @@ describe("heuristicV2Selection", () => {
isProxyAdmin: true,
stateLoading: false,
hasUnlimitedLicense: false,
slotTaken: true,
slotTaken: false,
currentOwnsSlot: true,
};
expect(heuristicV2Selection(params).allowed).toBe(true);
});
it("blocks a former licensed owner when another router now owns the singleton", () => {
const params = {
isProxyAdmin: true,
stateLoading: false,
hasUnlimitedLicense: false,
slotTaken: true,
currentOwnsSlot: true,
};
expect(heuristicV2Selection(params).allowed).toBe(false);
});
});
vi.mock("@/components/networking", () => ({

View file

@ -154,7 +154,7 @@ export const heuristicV2Selection = ({
currentOwnsSlot = false,
}: HeuristicV2SelectionParams): HeuristicV2Selection => {
if (!isProxyAdmin) return { allowed: false, lockedReason: "Only proxy admins can configure Heuristic v2." };
if (currentOwnsSlot) return { allowed: true };
if (currentOwnsSlot && !slotTaken) return { allowed: true };
if (stateLoading) return { allowed: false, lockedReason: "Checking Heuristic v2 availability..." };
if (slotTaken && !hasUnlimitedLicense) {
return {

View file

@ -142,10 +142,7 @@ export type ClassifierFallback = "heuristic" | "default_model";
export const DEFAULT_CLASSIFIER_FALLBACK: ClassifierFallback = "heuristic";
export interface AdaptiveRouterWeights {
quality: number;
cost: number;
}
export type AdaptiveRouterWeights = { quality: number; cost: number };
export const DEFAULT_ADAPTIVE_WEIGHTS: AdaptiveRouterWeights = { quality: 0.3, cost: 0.7 };