diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260902000000_one_heuristic_v2_router/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260902000000_one_heuristic_v2_router/migration.sql index 18302bf0245..71cde93cf87 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260902000000_one_heuristic_v2_router/migration.sql +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260902000000_one_heuristic_v2_router/migration.sql @@ -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' diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 4e63cb35e63..d398fb77394 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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") diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index c31c58eb66a..adba6df209f 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/auth/litellm_license.py b/litellm/proxy/auth/litellm_license.py index ae5a2cf8808..9d3ecabef23 100644 --- a/litellm/proxy/auth/litellm_license.py +++ b/litellm/proxy/auth/litellm_license.py @@ -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, diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index a5b16d25c2a..859ae314211 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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}, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c90bacda421..6f5de5c80bc 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 4e63cb35e63..d398fb77394 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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") diff --git a/litellm/router.py b/litellm/router.py index 5030fbfc44d..4967ac2a6d2 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/schema.prisma b/schema.prisma index 4e63cb35e63..d398fb77394 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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") diff --git a/tests/test_litellm/proxy/auth/test_litellm_license.py b/tests/test_litellm/proxy/auth/test_litellm_license.py index d6d00ac330d..76ec84f2ac6 100644 --- a/tests/test_litellm/proxy/auth/test_litellm_license.py +++ b/tests/test_litellm/proxy/auth/test_litellm_license.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 2c1bc78c363..8c93a99c017 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -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 ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts index 1bc79878df8..7351bc828ff 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts @@ -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", () => ({ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts index 10855f3f31a..57593a4becf 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts @@ -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 { diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index 1c42afa23a6..86e294b2e01 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -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 };