mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(router): license heuristic v2 singleton
This commit is contained in:
parent
56f009400f
commit
5fa015b2e5
20 changed files with 336 additions and 25 deletions
|
|
@ -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'
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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" && (
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
/>
|
||||
),
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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" />,
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
) : (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue