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 1d38f7ba04d..18302bf0245 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,7 +1,9 @@ --- Atomically reserve the proxy-wide heuristic_v2 classifier slot +ALTER TABLE "LiteLLM_ProxyModelTable" + ADD COLUMN IF NOT EXISTS "heuristic_v2_unlimited" BOOLEAN NOT NULL DEFAULT FALSE; + CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_ProxyModelTable_one_heuristic_v2_router" ON "LiteLLM_ProxyModelTable" ((1)) - WHERE CASE + WHERE NOT "heuristic_v2_unlimited" 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 7604ceadf7a..4e63cb35e63 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -53,6 +53,7 @@ model LiteLLM_ProxyModelTable { litellm_params Json model_info Json? blocked Boolean @default(false) + heuristic_v2_unlimited 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/constants.py b/litellm/constants.py index c7b74e176db..663033dfb16 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -39,6 +39,7 @@ ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset( "router_general_settings", "ignore_invalid_deployments", "fallback_access_check", + "allow_multiple_heuristic_v2", } ) DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512)) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 2da7ceb2d50..c31c58eb66a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4886,6 +4886,7 @@ class PrismaCompatibleUpdateDBModel(TypedDict, total=False): litellm_params: str model_info: str blocked: bool + heuristic_v2_unlimited: bool # writable-ok: the model write path stamps this internal license marker updated_at: str updated_by: str diff --git a/litellm/proxy/auth/litellm_license.py b/litellm/proxy/auth/litellm_license.py index 677f1a0fdda..ae5a2cf8808 100644 --- a/litellm/proxy/auth/litellm_license.py +++ b/litellm/proxy/auth/litellm_license.py @@ -15,6 +15,8 @@ from litellm.llms.custom_httpx.http_handler import HTTPHandler if TYPE_CHECKING: from litellm.proxy._types import EnterpriseLicenseData +AUTO_ROUTER_LICENSE_FEATURE: Final = "auto_router" + class LicenseCheck: """ @@ -149,6 +151,13 @@ class LicenseCheck: return False return team_count > _max_teams_in_license + def allows_feature(self, feature: str) -> bool: + license_data: Final = self.airgapped_license_data + if license_data is None: + return False + allowed_features: Final = license_data.get("allowed_features") + return isinstance(allowed_features, list) and feature in allowed_features + def verify_license_without_api_request(self, public_key, license_key): try: from cryptography.hazmat.primitives import hashes diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index c2c0d6ea146..a5b16d25c2a 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -49,6 +49,7 @@ from litellm.proxy._types import ( TeamModelDeleteRequest, UserAPIKeyAuth, ) +from litellm.proxy.auth.litellm_license import AUTO_ROUTER_LICENSE_FEATURE from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.config_sync_pubsub import ( coordination_redis_cache, @@ -297,6 +298,16 @@ def _effective_complexity_router_config( return config if isinstance(config, Mapping) else None +def _license_allows_unlimited_heuristic_v2() -> bool: + from litellm.proxy.proxy_server import _license_check + + return _license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE) + + +def _heuristic_v2_unlimited_marker(complexity_router_config: Mapping[str, object] | None) -> bool: + return uses_heuristic_v2(complexity_router_config) and _license_allows_unlimited_heuristic_v2() + + async def _raise_if_heuristic_v2_slot_taken( *, prisma_client: PrismaClient, @@ -320,9 +331,13 @@ async def _raise_if_heuristic_v2_slot_taken( code=status.HTTP_403_FORBIDDEN, param="litellm_params.complexity_router_config.classifier_type", ) + if _license_allows_unlimited_heuristic_v2(): + return from litellm.proxy.proxy_server import llm_router - rows: Final = await _proxy_model_table(prisma_client).find_many(where={}) + rows: Final = await _proxy_model_table(prisma_client).find_many( + where={} # mutable-ok: Prisma requires a mutable filter mapping + ) violation: Final = _heuristic_v2_slot_violation( persisted_rows=rows, live_model_list=llm_router.model_list if llm_router is not None else (), @@ -403,8 +418,8 @@ def _is_heuristic_v2_slot_unique_violation(error: Exception) -> bool: def _heuristic_v2_slot_proxy_exception() -> ProxyException: return ProxyException( message=( - "Only one complexity router can use classifier_type='heuristic_v2' per proxy. " - "Change or delete the existing heuristic_v2 router first." + "Heuristic v2 is limited to one auto-router. Change or delete the existing " + "heuristic_v2 router first, or reach out to tin@berri.ai to learn more." ), type=ProxyErrorTypes.validation_error.value, code=status.HTTP_400_BAD_REQUEST, @@ -885,6 +900,9 @@ async def patch_model( # Add metadata about update update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name update_data["updated_at"] = cast(str, get_utc_datetime()) + update_data["heuristic_v2_unlimited"] = _heuristic_v2_unlimited_marker( + _effective_complexity_router_config(patch_data.litellm_params, db_model.litellm_params) + ) # Perform partial update updated_model: Final = await _proxy_model_table(prisma_client).update( @@ -1135,6 +1153,7 @@ async def _add_model_to_db( "model_info": model_params.model_info.model_dump_json(exclude_none=True), "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), } if model_params.model_info.id is not None: _data["model_id"] = model_params.model_info.id @@ -2213,9 +2232,14 @@ async def update_model( else: pass - _data: Final[dict[str, str]] = { + _data: Final[dict[str, str | bool]] = { # mutable-ok: Prisma update payload is built for this write "litellm_params": json.dumps(merged_dictionary), "updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, + "heuristic_v2_unlimited": _heuristic_v2_unlimited_marker( + merged_dictionary.get("complexity_router_config") + if isinstance(merged_dictionary.get("complexity_router_config"), Mapping) + else None + ), } 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 27132c90e05..c90bacda421 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -301,7 +301,7 @@ from litellm.proxy.auth.auth_utils import ( ) from litellm.proxy.auth.fallback_model_access import router_fallback_access_check from litellm.proxy.auth.handle_jwt import JWTHandler -from litellm.proxy.auth.litellm_license import LicenseCheck +from litellm.proxy.auth.litellm_license import AUTO_ROUTER_LICENSE_FEATURE, LicenseCheck from litellm.proxy.auth.model_checks import ( expand_wildcard_deployments_for_model_info, get_all_fallbacks, @@ -5810,6 +5810,7 @@ class ProxyConfig: ), ignore_invalid_deployments=True, # don't raise an error if a deployment is invalid fallback_access_check=router_fallback_access_check, + allow_multiple_heuristic_v2=_license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE), ) if redis_usage_cache is not None and router.cache.redis_cache is None: @@ -6270,6 +6271,7 @@ class ProxyConfig: search_tools=search_tools, ignore_invalid_deployments=True, fallback_access_check=router_fallback_access_check, + allow_multiple_heuristic_v2=_license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE), ) verbose_proxy_logger.debug("updated llm_router: %s", llm_router) else: diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 7604ceadf7a..4e63cb35e63 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -53,6 +53,7 @@ model LiteLLM_ProxyModelTable { litellm_params Json model_info Json? blocked Boolean @default(false) + heuristic_v2_unlimited 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 2fc6d13c47b..5030fbfc44d 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -679,6 +679,7 @@ class Router: background_health_check_model_groups: Sequence[str] | None = None, enable_weighted_failover: bool = False, fallback_access_check: FallbackAccessCheck | None = None, + allow_multiple_heuristic_v2: bool = False, ) -> None: """ Initialize the Router class with the given parameters for caching, reliability, and routing strategy. @@ -756,6 +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.debug_level = debug_level self.enable_pre_call_checks = enable_pre_call_checks self.enable_tag_filtering = enable_tag_filtering @@ -8702,7 +8704,11 @@ class Router: complexity_router_config: Final[dict | None] = deployment.litellm_params.complexity_router_config - if complexity_router_config and complexity_router_config.get("classifier_type") == "heuristic_v2": + if ( + not self.allow_multiple_heuristic_v2 + and complexity_router_config + and complexity_router_config.get("classifier_type") == "heuristic_v2" + ): for registered in self.complexity_routers.values(): if any(tagged.strategy.config.classifier_type == "heuristic_v2" for tagged in registered): raise ValueError( diff --git a/schema.prisma b/schema.prisma index 7604ceadf7a..4e63cb35e63 100644 --- a/schema.prisma +++ b/schema.prisma @@ -53,6 +53,7 @@ model LiteLLM_ProxyModelTable { litellm_params Json model_info Json? blocked Boolean @default(false) + heuristic_v2_unlimited 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 8da365cb587..d6d00ac330d 100644 --- a/tests/test_litellm/proxy/auth/test_litellm_license.py +++ b/tests/test_litellm/proxy/auth/test_litellm_license.py @@ -1,17 +1,12 @@ -import asyncio -import json -from unittest.mock import AsyncMock, MagicMock, patch - - -from litellm.proxy.auth.litellm_license import LicenseCheck +from litellm.proxy.auth.litellm_license import AUTO_ROUTER_LICENSE_FEATURE, 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" + 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(): @@ -30,3 +25,16 @@ def test_is_over_limit(): 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_allows_feature_requires_signed_license_claim(): + license_check = LicenseCheck() + + license_check.airgapped_license_data = {"allowed_features": [AUTO_ROUTER_LICENSE_FEATURE]} + assert license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE) is True + + license_check.airgapped_license_data = {"allowed_features": ["other_feature"]} + assert license_check.allows_feature(AUTO_ROUTER_LICENSE_FEATURE) is False + + 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 5b9172a02f9..2c1bc78c363 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 @@ -1602,7 +1602,7 @@ class TestTeamModelUpdate: ) prisma_client = MockPrismaClient(team_exists=True, user_admin=False) - with patch( + with patch( # test-quality-ok: isolates the signed-license entitlement while testing marker behavior "litellm.proxy.proxy_server.premium_user", True, ): @@ -4111,6 +4111,47 @@ class TestStrategyRouterWriteValidation: ) assert not _is_heuristic_v2_slot_unique_violation(Exception("unrelated database error")) + def test_auto_router_license_marks_heuristic_v2_as_unlimited(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _heuristic_v2_unlimited_marker, + ) + + with patch( # test-quality-ok: isolates the unlicensed default while testing marker behavior + "litellm.proxy.management_endpoints.model_management_endpoints._license_allows_unlimited_heuristic_v2", + return_value=True, + ): + assert _heuristic_v2_unlimited_marker({"classifier_type": "heuristic_v2"}) is True + assert _heuristic_v2_unlimited_marker({"classifier_type": "heuristic"}) is False + + with patch( + "litellm.proxy.management_endpoints.model_management_endpoints._license_allows_unlimited_heuristic_v2", + return_value=False, + ): + assert _heuristic_v2_unlimited_marker({"classifier_type": "heuristic_v2"}) is False + + @pytest.mark.asyncio + async def test_auto_router_license_skips_singleton_lookup(self): # test-quality-ok: verifies entitlement bypasses the singleton database path + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _raise_if_heuristic_v2_slot_taken, + ) + + prisma_client = MagicMock() + with patch( # test-quality-ok: isolates the entitlement branch from global proxy license state + "litellm.proxy.management_endpoints.model_management_endpoints._license_allows_unlimited_heuristic_v2", + return_value=True, + ): + await _raise_if_heuristic_v2_slot_taken( + prisma_client=prisma_client, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + incoming_params=LiteLLM_Params( + model="auto_router/complexity_router", + complexity_router_config={"classifier_type": "heuristic_v2"}, + ), + existing_params=None, + ) + + prisma_client.db.litellm_proxymodeltable.find_many.assert_not_called() + def test_double_prefix_rejected_against_stored_params(self): from litellm.proxy.management_endpoints.model_management_endpoints import ( _strategy_router_write_violation, diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 863d2174415..7976fcb97fb 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -1113,6 +1113,35 @@ class TestRouterComplexityDeploymentMethods: with pytest.raises(ValueError, match="Only one complexity router"): router.init_complexity_router_deployment(deployment("second")) + def test_auto_router_license_allows_multiple_heuristic_v2_routers(self): + router = Router( + model_list=[ + { + "model_name": "gpt-4o-mini", + "litellm_params": {"model": "openai/gpt-4o-mini"}, + } + ], + allow_multiple_heuristic_v2=True, + ) + + for name in ("first", "second"): + router.init_complexity_router_deployment( + Deployment( + model_name=name, + litellm_params=LiteLLM_Params( + model=f"auto_router/complexity_router/{name}", + complexity_router_default_model="gpt-4o-mini", + complexity_router_config={ + "classifier_type": "heuristic_v2", + "tiers": {"SIMPLE": "gpt-4o-mini"}, + }, + ), + model_info={"id": name}, + ) + ) + + assert set(router.complexity_routers) == {"first", "second"} + def test_heuristic_v1_complexity_routers_remain_unlimited(self): router = Router( model_list=[ 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 7231c126a63..1bc79878df8 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 @@ -4,6 +4,7 @@ import React, { ReactNode } from "react"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { isAutoRouterDeployment, + heuristicV2Selection, selectAutoRouterModelGroups, selectPlainModelGroups, useAllProxyModels, @@ -20,6 +21,42 @@ import { type ProxyModel, } from "./useModels"; +describe("heuristicV2Selection", () => { + it("blocks a second router without the auto_router license feature", () => { + const params = { + isProxyAdmin: true, + stateLoading: false, + hasUnlimitedLicense: false, + slotTaken: true, + }; + expect(heuristicV2Selection(params)).toEqual({ + allowed: false, + lockedReason: "Heuristic v2 is limited to one auto-router. Reach out to tin@berri.ai to learn more.", + }); + }); + + it("allows unlimited routers with the auto_router license feature", () => { + const params = { + isProxyAdmin: true, + stateLoading: false, + hasUnlimitedLicense: true, + slotTaken: true, + }; + expect(heuristicV2Selection(params).allowed).toBe(true); + }); + + it("lets the current singleton owner keep heuristic v2 while editing", () => { + const params = { + isProxyAdmin: true, + stateLoading: false, + hasUnlimitedLicense: false, + slotTaken: true, + currentOwnsSlot: true, + }; + expect(heuristicV2Selection(params).allowed).toBe(true); + }); +}); + vi.mock("@/components/networking", () => ({ modelInfoCall: vi.fn(), modelHubCall: vi.fn(), 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 a9f7c54698a..10855f3f31a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts @@ -126,6 +126,45 @@ export const selectAutoRouterModelGroups = (deployments: AutoRouterCandidateDepl export const selectAutoRouterDeployments = (deployments: AutoRouterDeployment[]): AutoRouterDeployment[] => deployments.filter(isAutoRouterDeployment); +export const usesHeuristicV2Deployment = (deployment: AutoRouterDeployment): boolean => { + const config = deployment.litellm_params?.complexity_router_config; + return typeof config === "object" && config !== null && "classifier_type" in config + ? config.classifier_type === "heuristic_v2" + : false; +}; + +interface HeuristicV2SelectionParams { + isProxyAdmin: boolean; + stateLoading: boolean; + hasUnlimitedLicense: boolean; + slotTaken: boolean; + currentOwnsSlot?: boolean; +} + +export interface HeuristicV2Selection { + allowed: boolean; + lockedReason?: string; +} + +export const heuristicV2Selection = ({ + isProxyAdmin, + stateLoading, + hasUnlimitedLicense, + slotTaken, + currentOwnsSlot = false, +}: HeuristicV2SelectionParams): HeuristicV2Selection => { + if (!isProxyAdmin) return { allowed: false, lockedReason: "Only proxy admins can configure Heuristic v2." }; + if (currentOwnsSlot) return { allowed: true }; + if (stateLoading) return { allowed: false, lockedReason: "Checking Heuristic v2 availability..." }; + if (slotTaken && !hasUnlimitedLicense) { + return { + allowed: false, + lockedReason: "Heuristic v2 is limited to one auto-router. Reach out to tin@berri.ai to learn more.", + }; + } + return { allowed: true }; +}; + export const selectPlainModelGroups = (deployments: AutoRouterCandidateDeployment[]): ReadonlySet => { const autoRouterGroups = selectAutoRouterModelGroups(deployments); return new Set( diff --git a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx index 6e93e64d760..ebc02110add 100644 --- a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx @@ -161,6 +161,7 @@ interface ClassificationMethodConfigProps { defaultModel?: string; /** Mirrors the backend's proxy-admin-only heuristic_v2 write gate. */ allowHeuristicV2?: boolean; + heuristicV2LockedReason?: string; } const ClassifierTypeRadios: React.FC<{ @@ -168,12 +169,14 @@ const ClassifierTypeRadios: React.FC<{ classifierType: ClassifierType; onTypeChange: (classifierType: ClassifierType) => void; allowHeuristicV2: boolean; -}> = ({ value, classifierType, onTypeChange, allowHeuristicV2 }) => { + heuristicV2LockedReason?: string; +}> = ({ value, classifierType, onTypeChange, allowHeuristicV2, heuristicV2LockedReason }) => { const scorerLocked = Boolean(value.custom_tier_set); const scorerLockedReason = restrictedBy(value, "heuristicClassifier")?.reason; const heuristicV2Locked = scorerLocked || !allowHeuristicV2; - const heuristicV2LockedReason = - scorerLockedReason ?? (!allowHeuristicV2 ? "Only proxy admins can configure Heuristic v2." : undefined); + const lockedReason = + scorerLockedReason ?? + (!allowHeuristicV2 ? heuristicV2LockedReason ?? "Only proxy admins can configure Heuristic v2." : undefined); return ( - +