mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(router): limit heuristic v2 to one router
This commit is contained in:
parent
993766be0e
commit
4bc01c3b47
5 changed files with 225 additions and 0 deletions
|
|
@ -92,6 +92,7 @@ from litellm.router_strategy.complexity_router import (
|
|||
from litellm.router_utils.auto_router_model_naming import (
|
||||
STRATEGY_ROUTER_PARAM_FIELDS,
|
||||
carries_complexity_router_settings,
|
||||
uses_heuristic_v2,
|
||||
validate_complexity_router_config_placement,
|
||||
validate_complexity_router_config_write,
|
||||
validate_strategy_router_model_write,
|
||||
|
|
@ -143,6 +144,9 @@ class _ProxyModelRow(Protocol):
|
|||
|
||||
def model_dump_json(self, *, exclude_none: bool = False) -> str: ...
|
||||
|
||||
@property
|
||||
def litellm_params(self) -> object: ...
|
||||
|
||||
|
||||
class _ProxyModelTable(Protocol):
|
||||
def find_unique(self, *, where: Mapping[str, object]) -> Awaitable[BaseModel | None]: ...
|
||||
|
|
@ -265,6 +269,68 @@ def _raise_on_strategy_router_write_violation(
|
|||
)
|
||||
|
||||
|
||||
def _litellm_params_mapping(value: object) -> Mapping[str, object]:
|
||||
"""Normalize Prisma's parsed or JSON-encoded Json column at the typed boundary."""
|
||||
if isinstance(value, Mapping):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
parsed: Final = json.loads(value)
|
||||
except JSONDecodeError:
|
||||
return MappingProxyType({})
|
||||
return parsed if isinstance(parsed, Mapping) else MappingProxyType({})
|
||||
return MappingProxyType({})
|
||||
|
||||
|
||||
def _effective_complexity_router_config(
|
||||
incoming_params: GenericLiteLLMParams | None,
|
||||
existing_params: GenericLiteLLMParams | None,
|
||||
) -> Mapping[str, object] | None:
|
||||
"""Return the whole config that the write leaves on the deployment."""
|
||||
config: Final = (
|
||||
incoming_params.complexity_router_config
|
||||
if incoming_params is not None and incoming_params.complexity_router_config is not None
|
||||
else existing_params.complexity_router_config
|
||||
if existing_params is not None
|
||||
else None
|
||||
)
|
||||
return config if isinstance(config, Mapping) else None
|
||||
|
||||
|
||||
async def _raise_if_heuristic_v2_slot_taken(
|
||||
*,
|
||||
prisma_client: PrismaClient,
|
||||
incoming_params: GenericLiteLLMParams | None,
|
||||
existing_params: GenericLiteLLMParams | None,
|
||||
current_model_id: str | None = None,
|
||||
) -> None:
|
||||
"""Allow at most one persisted heuristic-v2 complexity router per proxy."""
|
||||
effective_config: Final = _effective_complexity_router_config(incoming_params, existing_params)
|
||||
if not uses_heuristic_v2(effective_config):
|
||||
return
|
||||
rows: Final = await _proxy_model_table(prisma_client).find_many(where={})
|
||||
for row in rows:
|
||||
if current_model_id is not None and row.model_id == current_model_id:
|
||||
continue
|
||||
if uses_heuristic_v2(
|
||||
config
|
||||
if isinstance(
|
||||
config := _litellm_params_mapping(row.litellm_params).get("complexity_router_config"), Mapping
|
||||
)
|
||||
else None
|
||||
):
|
||||
raise ProxyException(
|
||||
message=(
|
||||
"Only one complexity router can use classifier_type='heuristic_v2' per proxy. "
|
||||
"Use classifier_type='heuristic' for this router, or change or delete the existing "
|
||||
"heuristic_v2 router first."
|
||||
),
|
||||
type=ProxyErrorTypes.validation_error.value,
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
param="litellm_params.complexity_router_config.classifier_type",
|
||||
)
|
||||
|
||||
|
||||
ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING: Final = "enforce_rpm_tpm_on_model_add"
|
||||
_REQUIRED_RATE_LIMIT_FIELDS: Final = ("rpm", "tpm")
|
||||
|
||||
|
|
@ -714,6 +780,12 @@ async def patch_model(
|
|||
incoming_params=patch_data.litellm_params,
|
||||
existing_params=db_model.litellm_params,
|
||||
)
|
||||
await _raise_if_heuristic_v2_slot_taken(
|
||||
prisma_client=prisma_client,
|
||||
incoming_params=patch_data.litellm_params,
|
||||
existing_params=db_model.litellm_params,
|
||||
current_model_id=model_id,
|
||||
)
|
||||
|
||||
# Handle team model updates with proper alias management
|
||||
update_data: Final = await _update_team_model_in_db(
|
||||
|
|
@ -1830,6 +1902,11 @@ async def add_new_model(
|
|||
incoming_params=model_params.litellm_params,
|
||||
existing_params=None,
|
||||
)
|
||||
await _raise_if_heuristic_v2_slot_taken(
|
||||
prisma_client=prisma_client,
|
||||
incoming_params=model_params.litellm_params,
|
||||
existing_params=None,
|
||||
)
|
||||
|
||||
_raise_if_rate_limits_required_but_missing(
|
||||
litellm_params=model_params.litellm_params,
|
||||
|
|
@ -2008,6 +2085,12 @@ async def update_model(
|
|||
incoming_params=model_params.litellm_params,
|
||||
existing_params=deployment.litellm_params,
|
||||
)
|
||||
await _raise_if_heuristic_v2_slot_taken(
|
||||
prisma_client=prisma_client,
|
||||
incoming_params=model_params.litellm_params,
|
||||
existing_params=deployment.litellm_params,
|
||||
current_model_id=_model_id,
|
||||
)
|
||||
|
||||
# update DB
|
||||
if store_model_in_db is True:
|
||||
|
|
|
|||
|
|
@ -8702,6 +8702,15 @@ 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":
|
||||
for registered in self.complexity_routers.values():
|
||||
if any(tagged.strategy.config.classifier_type == "heuristic_v2" for tagged in registered):
|
||||
raise ValueError(
|
||||
"Only one complexity router can use classifier_type='heuristic_v2' per proxy. "
|
||||
"Use classifier_type='heuristic' for this router, or change or remove the existing "
|
||||
"heuristic_v2 router first."
|
||||
)
|
||||
|
||||
default_model: str | None = deployment.litellm_params.complexity_router_default_model
|
||||
|
||||
# If no default model specified, try to get from config tiers. Derived from the
|
||||
|
|
|
|||
|
|
@ -210,6 +210,11 @@ def carries_complexity_router_settings(model: str | None, present_fields: frozen
|
|||
)
|
||||
|
||||
|
||||
def uses_heuristic_v2(complexity_router_config: Mapping[str, object] | None) -> bool:
|
||||
"""Whether a complexity-router config selects the singleton v2 classifier."""
|
||||
return complexity_router_config is not None and complexity_router_config.get("classifier_type") == "heuristic_v2"
|
||||
|
||||
|
||||
def validate_complexity_router_config_placement(litellm_params: Mapping[str, object] | None) -> str | None:
|
||||
"""Reject a complexity-router setting written beside ``complexity_router_config``.
|
||||
|
||||
|
|
|
|||
|
|
@ -3953,6 +3953,80 @@ class TestStrategyRouterWriteValidation:
|
|||
model_info={"id": model_id},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_second_heuristic_v2_router_is_rejected(self):
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_raise_if_heuristic_v2_slot_taken,
|
||||
)
|
||||
|
||||
table = MagicMock()
|
||||
row = MagicMock()
|
||||
row.model_id = "first-v2"
|
||||
row.litellm_params = {"complexity_router_config": {"classifier_type": "heuristic_v2"}}
|
||||
table.find_many = AsyncMock(return_value=[row])
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints._proxy_model_table",
|
||||
return_value=table,
|
||||
):
|
||||
with pytest.raises(ProxyException, match="Only one complexity router"):
|
||||
await _raise_if_heuristic_v2_slot_taken(
|
||||
prisma_client=MagicMock(),
|
||||
incoming_params=LiteLLM_Params(
|
||||
model="auto_router/complexity_router",
|
||||
complexity_router_config={"classifier_type": "heuristic_v2"},
|
||||
),
|
||||
existing_params=None,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_existing_heuristic_v2_router_can_be_edited(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_raise_if_heuristic_v2_slot_taken,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
table = MagicMock()
|
||||
row = MagicMock()
|
||||
row.model_id = "first-v2"
|
||||
row.litellm_params = json.dumps({"complexity_router_config": {"classifier_type": "heuristic_v2"}})
|
||||
table.find_many = AsyncMock(return_value=[row])
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints._proxy_model_table",
|
||||
return_value=table,
|
||||
):
|
||||
await _raise_if_heuristic_v2_slot_taken(
|
||||
prisma_client=MagicMock(),
|
||||
incoming_params=updateLiteLLMParams(rpm=10),
|
||||
existing_params=LiteLLM_Params(
|
||||
model="auto_router/complexity_router",
|
||||
complexity_router_config={"classifier_type": "heuristic_v2"},
|
||||
),
|
||||
current_model_id="first-v2",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_heuristic_v1_does_not_consume_v2_slot(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_raise_if_heuristic_v2_slot_taken,
|
||||
)
|
||||
|
||||
table = MagicMock()
|
||||
table.find_many = AsyncMock()
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints._proxy_model_table",
|
||||
return_value=table,
|
||||
):
|
||||
await _raise_if_heuristic_v2_slot_taken(
|
||||
prisma_client=MagicMock(),
|
||||
incoming_params=LiteLLM_Params(
|
||||
model="auto_router/complexity_router",
|
||||
complexity_router_config={"classifier_type": "heuristic"},
|
||||
),
|
||||
existing_params=None,
|
||||
)
|
||||
table.find_many.assert_not_awaited()
|
||||
|
||||
def test_double_prefix_rejected_against_stored_params(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_strategy_router_write_violation,
|
||||
|
|
|
|||
|
|
@ -1085,6 +1085,60 @@ class TestRouterComplexityDeploymentMethods:
|
|||
router.init_complexity_router_deployment(deployment)
|
||||
assert "auto_router/complexity_router/test-router" in router.complexity_routers
|
||||
|
||||
def test_only_one_heuristic_v2_complexity_router_can_register(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
def deployment(name: str) -> Deployment:
|
||||
return 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},
|
||||
)
|
||||
|
||||
router.init_complexity_router_deployment(deployment("first"))
|
||||
with pytest.raises(ValueError, match="Only one complexity router"):
|
||||
router.init_complexity_router_deployment(deployment("second"))
|
||||
|
||||
def test_heuristic_v1_complexity_routers_remain_unlimited(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini"},
|
||||
}
|
||||
]
|
||||
)
|
||||
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",
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini"},
|
||||
},
|
||||
),
|
||||
model_info={"id": name},
|
||||
)
|
||||
)
|
||||
assert set(router.complexity_routers) == {"first", "second"}
|
||||
|
||||
def test_hybrid_initialization_waits_for_later_pool_deployments(self):
|
||||
router = Router(
|
||||
model_list=[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue