feat(router): license heuristic v2 singleton

This commit is contained in:
Tin 2026-09-02 17:08:04 -07:00
parent 56f009400f
commit 5fa015b2e5
20 changed files with 336 additions and 25 deletions

View file

@ -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'

View file

@ -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")

View file

@ -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))

View file

@ -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

View file

@ -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

View file

@ -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},

View file

@ -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:

View file

@ -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")

View file

@ -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(

View file

@ -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")

View file

@ -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

View file

@ -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,

View file

@ -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=[

View file

@ -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(),

View file

@ -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<string> => {
const autoRouterGroups = selectAutoRouterModelGroups(deployments);
return new Set(

View file

@ -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 (
<RadioGroup
value={classifierType}
@ -192,7 +195,7 @@ const ClassifierTypeRadios: React.FC<{
</span>
</Label>
</SimpleTooltip>
<SimpleTooltip content={heuristicV2LockedReason}>
<SimpleTooltip content={lockedReason}>
<Label className="items-start font-normal leading-normal has-data-disabled:cursor-not-allowed has-data-disabled:opacity-50">
<RadioGroupItem value="heuristic_v2" className="mt-0.5" disabled={heuristicV2Locked} />
<span>
@ -247,6 +250,7 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
showValidationErrors = false,
defaultModel,
allowHeuristicV2 = false,
heuristicV2LockedReason,
}) => {
const [draft, setDraft] = React.useState<{ id: string; raw: string } | null>(null);
const hasDefaultModel = Boolean(defaultModel);
@ -399,6 +403,7 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
classifierType={classifierType}
onTypeChange={handleClassifierTypeChange}
allowHeuristicV2={allowHeuristicV2}
heuristicV2LockedReason={heuristicV2LockedReason}
/>
{classifierType === "heuristic_first" && (

View file

@ -485,6 +485,7 @@ interface ComplexityRouterConfigProps {
onEscalationKeywordsChange?: (keywords: string[]) => void;
showValidationErrors?: boolean;
allowHeuristicV2?: boolean;
heuristicV2LockedReason?: string;
}
export const TIER_DESCRIPTIONS: Record<
@ -607,6 +608,7 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
onEscalationKeywordsChange,
showValidationErrors = false,
allowHeuristicV2 = false,
heuristicV2LockedReason,
}) => {
const customTierSet = value.custom_tier_set;
const tierRows = activeTierRows(value);
@ -824,6 +826,7 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
showValidationErrors={showValidationErrors}
defaultModel={defaultModel}
allowHeuristicV2={allowHeuristicV2}
heuristicV2LockedReason={heuristicV2LockedReason}
/>
),
},

View file

@ -73,9 +73,10 @@ const addKeyword = async (user: ReturnType<typeof userEvent.setup>, field: HTMLE
await user.click(await screen.findByText(`Create "${keyword}"`));
};
const { mockFetchAvailableModels, mockFetchAllModelDeployments } = vi.hoisted(() => ({
const { mockFetchAvailableModels, mockFetchAllModelDeployments, mockUseLicenseInfo } = vi.hoisted(() => ({
mockFetchAvailableModels: vi.fn(),
mockFetchAllModelDeployments: vi.fn(),
mockUseLicenseInfo: vi.fn(),
}));
const { validateAutoRouterConfig } = vi.hoisted(() => ({
@ -97,6 +98,10 @@ vi.mock("@/app/(dashboard)/hooks/models/useModels", async (importOriginal) => {
return { ...actual, fetchAllModelDeployments: mockFetchAllModelDeployments };
});
vi.mock("@/app/(dashboard)/hooks/license/useLicenseInfo", () => ({
useLicenseInfo: mockUseLicenseInfo,
}));
vi.mock("./handle_add_auto_router_submit", () => ({
handleAddAutoRouterSubmit: vi.fn(),
}));
@ -141,6 +146,7 @@ describe("AddAutoRouterTab", () => {
testQueryClient.clear();
mockFetchAvailableModels.mockResolvedValue([]);
mockFetchAllModelDeployments.mockResolvedValue([]);
mockUseLicenseInfo.mockReturnValue({ data: { allowed_features: [] }, isLoading: false });
});
// Detailed Configuration starts collapsed so the modal opens onto just Name + Template; a caller
@ -180,6 +186,50 @@ describe("AddAutoRouterTab", () => {
expect(screen.queryByTestId("team-dropdown")).not.toBeInTheDocument();
});
it("disables heuristic v2 when another auto-router owns the singleton slot", async () => {
mockFetchAllModelDeployments.mockResolvedValue([
{
model_name: "existing-router",
litellm_params: {
model: "auto_router/complexity_router",
complexity_router_config: { classifier_type: "heuristic_v2" },
},
},
]);
renderWithProviders(<Harness />);
expandDetailedConfiguration();
fireEvent.click(screen.getByText("Advanced: Classification Method"));
await waitFor(() =>
expect(screen.getByRole("radio", { name: /Heuristic v2/ })).toHaveAttribute("aria-disabled", "true"),
);
});
it("keeps heuristic v2 available when the enterprise license includes auto_router", async () => {
mockFetchAllModelDeployments.mockResolvedValue([
{
model_name: "existing-router",
litellm_params: {
model: "auto_router/complexity_router",
complexity_router_config: { classifier_type: "heuristic_v2" },
},
},
]);
mockUseLicenseInfo.mockReturnValue({
data: { allowed_features: ["auto_router"] },
isLoading: false,
});
renderWithProviders(<Harness />);
expandDetailedConfiguration();
fireEvent.click(screen.getByText("Advanced: Classification Method"));
await waitFor(() =>
expect(screen.getByRole("radio", { name: /Heuristic v2/ })).not.toHaveAttribute("aria-disabled", "true"),
);
});
it("requires a team admin to pick a team", async () => {
renderWithProviders(
<AddAutoRouterTab handleOk={vi.fn()} accessToken="token" userRole="Internal User" createScope="team-required" />,

View file

@ -20,7 +20,12 @@ import { type ModelWriteScope } from "@/utils/modelPermissions";
import TeamDropdown from "../common_components/team_dropdown";
import { type AddAutoRouterValues, handleAddAutoRouterSubmit } from "./handle_add_auto_router_submit";
import { fetchAvailableModels } from "@/components/llm_calls/fetch_models";
import { autoRouterListKey, fetchAllModelDeployments } from "@/app/(dashboard)/hooks/models/useModels";
import {
autoRouterListKey,
fetchAllModelDeployments,
heuristicV2Selection,
usesHeuristicV2Deployment,
} from "@/app/(dashboard)/hooks/models/useModels";
import ComplexityRouterConfig, {
ComplexityRouterConfigValue,
effectiveClassifierType,
@ -62,6 +67,7 @@ import {
} from "@/lib/autorouter_presets";
import { useAutoRouterPresets } from "@/app/(dashboard)/hooks/autoRouter/useAutoRouterPresets";
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
import { useLicenseInfo } from "@/app/(dashboard)/hooks/license/useLicenseInfo";
interface AddAutoRouterTabProps {
handleOk: () => void;
@ -224,6 +230,17 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
queryFn: () => fetchAllModelDeployments(accessToken, userId ?? "", userRole),
enabled: Boolean(accessToken),
});
const licenseInfo = useLicenseInfo(accessToken);
const hasUnlimitedAutoRouterLicense = licenseInfo.data?.allowed_features.includes("auto_router") === true;
const heuristicV2SlotTaken = deployments?.some(usesHeuristicV2Deployment) === true;
const heuristicV2StateLoading = deploymentsLoading || licenseInfo.isLoading;
const heuristicV2SelectionParams = {
isProxyAdmin: createScope === "unscoped-ok",
stateLoading: heuristicV2StateLoading,
hasUnlimitedLicense: hasUnlimitedAutoRouterLicense,
slotTaken: heuristicV2SlotTaken,
};
const heuristicV2Access = heuristicV2Selection(heuristicV2SelectionParams);
const modelsLoading = groupsLoading || deploymentsLoading;
const modelInfo = React.useMemo(() => data ?? [], [data]);
const {
@ -615,7 +632,8 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
escalationKeywords={escalationKeywords}
onEscalationKeywordsChange={setEscalationKeywords}
showValidationErrors={showValidationErrors}
allowHeuristicV2={createScope === "unscoped-ok"}
allowHeuristicV2={heuristicV2Access.allowed}
heuristicV2LockedReason={heuristicV2Access.lockedReason}
/>
</div>
)}

View file

@ -1,4 +1,5 @@
import React, { useEffect, useMemo, useState } from "react";
import { useQuery } from "@tanstack/react-query";
import { z } from "zod/v4";
import { toast } from "@/lib/toast";
import { CircleHelp } from "lucide-react";
@ -9,10 +10,18 @@ import { Input } from "@/components/ui/input";
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { useZodForm } from "@/lib/forms/useZodForm";
import { isProxyAdminRole } from "@/utils/roles";
import AccessGroupTagsCombobox from "../add_model/AccessGroupTagsCombobox";
import ModelChoiceCombobox, { type ModelChoice } from "../add_model/ModelChoiceCombobox";
import { modelAvailableCall, modelPatchUpdateCall, validateAutoRouterConfig } from "../networking";
import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models";
import { useLicenseInfo } from "@/app/(dashboard)/hooks/license/useLicenseInfo";
import {
autoRouterListKey,
fetchAllModelDeployments,
heuristicV2Selection,
usesHeuristicV2Deployment,
} from "@/app/(dashboard)/hooks/models/useModels";
import RouterConfigBuilder from "../add_model/RouterConfigBuilder";
import { hydrateTierModelParams, normalizeTierModels } from "../add_model/complexity_router_tiers";
import {
@ -415,6 +424,28 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
classifier_type: "heuristic",
});
const isComplexityRouterModel = isComplexityRouter(modelData?.litellm_params);
const licenseInfo = useLicenseInfo(accessToken);
const { data: deployments, isLoading: deploymentsLoading } = useQuery({
queryKey: autoRouterListKey(modelData?.model_info?.id ?? "", userRole),
queryFn: () => fetchAllModelDeployments(accessToken, "", userRole),
enabled: isVisible && Boolean(accessToken),
});
const currentModelId = modelData?.model_info?.id;
const currentOwnsHeuristicV2 = usesHeuristicV2Deployment(modelData ?? {});
const anotherRouterOwnsHeuristicV2 =
deployments?.some(
(deployment) => deployment.model_info?.id !== currentModelId && usesHeuristicV2Deployment(deployment),
) === true;
const hasUnlimitedAutoRouterLicense = licenseInfo.data?.allowed_features.includes("auto_router") === true;
const heuristicV2StateLoading = deploymentsLoading || licenseInfo.isLoading;
const heuristicV2SelectionParams = {
isProxyAdmin: isProxyAdminRole(userRole),
stateLoading: heuristicV2StateLoading,
hasUnlimitedLicense: hasUnlimitedAutoRouterLicense,
slotTaken: anotherRouterOwnsHeuristicV2,
currentOwnsSlot: currentOwnsHeuristicV2,
};
const heuristicV2Access = heuristicV2Selection(heuristicV2SelectionParams);
const schema = useMemo(
() => (isComplexityRouterModel ? complexityRouterSchema : semanticRouterSchema),
@ -729,6 +760,8 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
onMatchThresholdChange={setMatchThreshold}
escalationKeywords={escalationKeywords}
onEscalationKeywordsChange={setEscalationKeywords}
allowHeuristicV2={heuristicV2Access.allowed}
heuristicV2LockedReason={heuristicV2Access.lockedReason}
/>
</div>
) : (