mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #34151 from BerriAI/litellm_lit4663_autorouter_prefix
fix(proxy): reject model writes that corrupt an auto-router pseudo-model
This commit is contained in:
commit
9b48bf6084
5 changed files with 457 additions and 16 deletions
|
|
@ -57,10 +57,15 @@ from litellm.router import Router
|
|||
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
|
||||
UpdateUsefulLinksRequest,
|
||||
)
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
STRATEGY_ROUTER_PARAM_FIELDS,
|
||||
validate_strategy_router_model_write,
|
||||
)
|
||||
from litellm.types.router import (
|
||||
SPECIAL_MODEL_INFO_PARAMS,
|
||||
Deployment,
|
||||
DeploymentTypedDict,
|
||||
GenericLiteLLMParams,
|
||||
LiteLLMParamsTypedDict,
|
||||
updateDeployment,
|
||||
)
|
||||
|
|
@ -98,6 +103,45 @@ async def get_db_model(model_id: str, prisma_client: PrismaClient) -> Optional[D
|
|||
return deployment_pydantic_obj
|
||||
|
||||
|
||||
def _strategy_router_write_violation(
|
||||
incoming_params: GenericLiteLLMParams | None,
|
||||
existing_params: GenericLiteLLMParams | None,
|
||||
) -> str | None:
|
||||
"""Reject writes that would corrupt a strategy router's pseudo-model.
|
||||
|
||||
An auto-router deployment's ``litellm_params.model`` (``auto_router/...``) is
|
||||
the discriminator the router loads it by; a write that mangles it makes the
|
||||
router drop the deployment silently under ``ignore_invalid_deployments``.
|
||||
Only writes that supply ``litellm_params.model`` are judged, against the
|
||||
merged (stored + incoming) params, so partial patches and restores of an
|
||||
already-corrupted row stay legal. Returns the violation, or None.
|
||||
"""
|
||||
if incoming_params is None or incoming_params.model is None:
|
||||
return None
|
||||
present_fields = frozenset(
|
||||
field
|
||||
for field in STRATEGY_ROUTER_PARAM_FIELDS
|
||||
for source in (incoming_params, existing_params)
|
||||
if source is not None and getattr(source, field, None) is not None
|
||||
)
|
||||
return validate_strategy_router_model_write(model=incoming_params.model, present_fields=present_fields)
|
||||
|
||||
|
||||
def _raise_on_strategy_router_write_violation(
|
||||
incoming_params: GenericLiteLLMParams | None,
|
||||
existing_params: GenericLiteLLMParams | None,
|
||||
) -> None:
|
||||
violation = _strategy_router_write_violation(incoming_params=incoming_params, existing_params=existing_params)
|
||||
if violation is None:
|
||||
return
|
||||
raise ProxyException(
|
||||
message=violation,
|
||||
type=ProxyErrorTypes.validation_error.value,
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
param="litellm_params.model",
|
||||
)
|
||||
|
||||
|
||||
def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> PrismaCompatibleUpdateDBModel:
|
||||
merged_deployment_dict = DeploymentTypedDict(
|
||||
model_name=db_model.model_name,
|
||||
|
|
@ -255,6 +299,11 @@ async def patch_model(
|
|||
param="blocked",
|
||||
)
|
||||
|
||||
_raise_on_strategy_router_write_violation(
|
||||
incoming_params=patch_data.litellm_params,
|
||||
existing_params=db_model.litellm_params,
|
||||
)
|
||||
|
||||
# Handle team model updates with proper alias management
|
||||
update_data = await _update_team_model_in_db(
|
||||
db_model=db_model,
|
||||
|
|
@ -1292,6 +1341,11 @@ async def add_new_model(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
_raise_on_strategy_router_write_violation(
|
||||
incoming_params=model_params.litellm_params,
|
||||
existing_params=None,
|
||||
)
|
||||
|
||||
model_response: Optional[LiteLLM_ProxyModelTable] = None
|
||||
# update DB
|
||||
if store_model_in_db is True:
|
||||
|
|
@ -1446,6 +1500,11 @@ async def update_model(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
_raise_on_strategy_router_write_violation(
|
||||
incoming_params=model_params.litellm_params,
|
||||
existing_params=deployment.litellm_params,
|
||||
)
|
||||
|
||||
# update DB
|
||||
if store_model_in_db is True:
|
||||
_existing_litellm_params_dict = dict(_existing_litellm_params.litellm_params)
|
||||
|
|
|
|||
|
|
@ -110,6 +110,9 @@ from litellm.router_utils.batch_utils import (
|
|||
replace_model_in_jsonl,
|
||||
should_replace_model_in_jsonl,
|
||||
)
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
classify_strategy_router_model,
|
||||
)
|
||||
from litellm.router_utils.client_initalization_utils import InitalizeCachedClient
|
||||
from litellm.router_utils.clientside_credential_handler import (
|
||||
get_dynamic_litellm_params,
|
||||
|
|
@ -7623,15 +7626,7 @@ class Router:
|
|||
but NOT "auto_router/complexity_router" or "auto_router/adaptive_router"
|
||||
(which use the complexity-router and adaptive-router strategies).
|
||||
"""
|
||||
if litellm_params.model.startswith("auto_router/complexity_router"):
|
||||
return False # This is handled by complexity_router
|
||||
if litellm_params.model.startswith("auto_router/adaptive_router"):
|
||||
return False # This is handled by adaptive_router
|
||||
if litellm_params.model.startswith("auto_router/quality_router"):
|
||||
return False # This is handled by quality_router
|
||||
if litellm_params.model.startswith("auto_router/"):
|
||||
return True
|
||||
return False
|
||||
return classify_strategy_router_model(litellm_params.model) == "semantic"
|
||||
|
||||
@staticmethod
|
||||
def _deployment_tags(deployment: Deployment) -> tuple[str, ...]:
|
||||
|
|
@ -7686,9 +7681,7 @@ class Router:
|
|||
|
||||
Returns True if the litellm_params model starts with "auto_router/complexity_router"
|
||||
"""
|
||||
if litellm_params.model.startswith("auto_router/complexity_router"):
|
||||
return True
|
||||
return False
|
||||
return classify_strategy_router_model(litellm_params.model) == "complexity"
|
||||
|
||||
def init_complexity_router_deployment(self, deployment: Deployment):
|
||||
"""
|
||||
|
|
@ -7738,7 +7731,7 @@ class Router:
|
|||
|
||||
def _is_adaptive_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
|
||||
"""True when this deployment opts in via the `auto_router/adaptive_router` model prefix."""
|
||||
return litellm_params.model.startswith("auto_router/adaptive_router")
|
||||
return classify_strategy_router_model(litellm_params.model) == "adaptive"
|
||||
|
||||
def _deployment_participates_in_adaptive_routing(self, litellm_params: LiteLLM_Params) -> bool:
|
||||
"""True when this deployment owns an `adaptive_routers` entry once finalized:
|
||||
|
|
@ -7964,9 +7957,7 @@ class Router:
|
|||
|
||||
Returns True if the litellm_params model starts with "auto_router/quality_router".
|
||||
"""
|
||||
if litellm_params.model.startswith("auto_router/quality_router"):
|
||||
return True
|
||||
return False
|
||||
return classify_strategy_router_model(litellm_params.model) == "quality"
|
||||
|
||||
def init_quality_router_deployment(self, deployment: Deployment):
|
||||
"""
|
||||
|
|
|
|||
101
litellm/router_utils/auto_router_model_naming.py
Normal file
101
litellm/router_utils/auto_router_model_naming.py
Normal file
|
|
@ -0,0 +1,101 @@
|
|||
"""Naming contract for strategy-router (auto-router) pseudo-models.
|
||||
|
||||
A deployment whose ``litellm_params.model`` starts with ``auto_router/`` does not
|
||||
name a provider model; the string is the discriminator that selects which
|
||||
pre-routing strategy owns the deployment. This module is the single source of
|
||||
truth for classifying that string (``Router._is_*_router_deployment`` delegates
|
||||
here) and for checking that a client-supplied write leaves the deployment
|
||||
coherent, so management endpoints can reject corruption with a 400 instead of
|
||||
the router silently dropping the deployment at load time under
|
||||
``ignore_invalid_deployments``.
|
||||
"""
|
||||
|
||||
from typing import Literal, Mapping
|
||||
|
||||
AUTO_ROUTER_MODEL_PREFIX = "auto_router/"
|
||||
|
||||
StrategyRouterKind = Literal["semantic", "complexity", "adaptive", "quality"]
|
||||
|
||||
STRATEGY_ROUTER_PARAM_FIELDS: frozenset[str] = frozenset(
|
||||
{
|
||||
"auto_router_config",
|
||||
"auto_router_config_path",
|
||||
"auto_router_default_model",
|
||||
"auto_router_embedding_model",
|
||||
"complexity_router_config",
|
||||
"complexity_router_default_model",
|
||||
"adaptive_router_config",
|
||||
"quality_router_config",
|
||||
"quality_router_default_model",
|
||||
}
|
||||
)
|
||||
|
||||
_REQUIRED_FIELD_GROUPS: Mapping[StrategyRouterKind, tuple[tuple[str, ...], ...]] = {
|
||||
"semantic": (
|
||||
("auto_router_config", "auto_router_config_path"),
|
||||
("auto_router_default_model",),
|
||||
("auto_router_embedding_model",),
|
||||
),
|
||||
"complexity": (("complexity_router_config", "complexity_router_default_model"),),
|
||||
"adaptive": (("adaptive_router_config",),),
|
||||
"quality": (("quality_router_config", "quality_router_default_model"),),
|
||||
}
|
||||
|
||||
|
||||
def classify_strategy_router_model(model: str) -> StrategyRouterKind | None:
|
||||
"""Classify a ``litellm_params.model`` string the way the Router does.
|
||||
|
||||
Returns None for regular provider models. Mirrors Router registration
|
||||
exactly: reserved names are matched by prefix, everything else under
|
||||
``auto_router/`` is a semantic router.
|
||||
"""
|
||||
if not model.startswith(AUTO_ROUTER_MODEL_PREFIX):
|
||||
return None
|
||||
remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :]
|
||||
if remainder.startswith("complexity_router"):
|
||||
return "complexity"
|
||||
if remainder.startswith("adaptive_router"):
|
||||
return "adaptive"
|
||||
if remainder.startswith("quality_router"):
|
||||
return "quality"
|
||||
return "semantic"
|
||||
|
||||
|
||||
def validate_strategy_router_model_write(model: str, present_fields: frozenset[str]) -> str | None:
|
||||
"""Check that writing ``model`` leaves a deployment the router can load.
|
||||
|
||||
``present_fields`` is the set of strategy-router param fields that are
|
||||
non-None on the deployment after the write (stored fields merged with the
|
||||
incoming ones). Returns a human-readable violation, or None when coherent.
|
||||
"""
|
||||
kind = classify_strategy_router_model(model)
|
||||
if kind is None:
|
||||
offending = sorted(present_fields & STRATEGY_ROUTER_PARAM_FIELDS)
|
||||
if offending:
|
||||
return (
|
||||
f"litellm_params.model='{model}' does not start with '{AUTO_ROUTER_MODEL_PREFIX}' but the "
|
||||
f"deployment carries auto-router settings ({', '.join(offending)}), so the router could not "
|
||||
f"load it. Keep the '{AUTO_ROUTER_MODEL_PREFIX}' prefix; to change the name clients call, "
|
||||
"edit the public model_name instead."
|
||||
)
|
||||
return None
|
||||
remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :]
|
||||
if remainder.startswith(AUTO_ROUTER_MODEL_PREFIX):
|
||||
return (
|
||||
f"litellm_params.model='{model}' repeats the '{AUTO_ROUTER_MODEL_PREFIX}' prefix, so the router "
|
||||
f"could not load it. Use '{remainder}'; to change the name clients call, edit the public "
|
||||
"model_name instead."
|
||||
)
|
||||
if not remainder:
|
||||
return (
|
||||
f"litellm_params.model='{model}' is missing the router name after the '{AUTO_ROUTER_MODEL_PREFIX}' prefix."
|
||||
)
|
||||
missing = tuple(
|
||||
" or ".join(group) for group in _REQUIRED_FIELD_GROUPS[kind] if not any(f in present_fields for f in group)
|
||||
)
|
||||
if missing:
|
||||
return (
|
||||
f"litellm_params.model='{model}' selects the {kind} router, which requires "
|
||||
f"{'; '.join(missing)} in litellm_params."
|
||||
)
|
||||
return None
|
||||
|
|
@ -3330,3 +3330,216 @@ class TestModelInfoAsMapping:
|
|||
assert model_info_as_mapping("{not json") is None
|
||||
assert model_info_as_mapping('["a", "b"]') is None
|
||||
assert model_info_as_mapping(42) is None
|
||||
|
||||
|
||||
class TestStrategyRouterWriteValidation:
|
||||
"""Management write paths must reject litellm_params.model values that would
|
||||
corrupt a strategy router's pseudo-model (LIT-4663). The router loads these
|
||||
deployments by the auto_router/ discriminator, so a mangled string makes it
|
||||
drop the deployment silently under ignore_invalid_deployments; the mistake
|
||||
has to fail loudly at the API boundary instead."""
|
||||
|
||||
def _stored_complexity_params(self) -> LiteLLM_Params:
|
||||
return LiteLLM_Params(
|
||||
model="auto_router/complexity_router",
|
||||
complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini"}},
|
||||
)
|
||||
|
||||
def _db_complexity_router(self, model_id: str) -> Deployment:
|
||||
return Deployment(
|
||||
model_name="my-auto-router",
|
||||
litellm_params=self._stored_complexity_params(),
|
||||
model_info={"id": model_id},
|
||||
)
|
||||
|
||||
def test_double_prefix_rejected_against_stored_params(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_strategy_router_write_violation,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
violation = _strategy_router_write_violation(
|
||||
incoming_params=updateLiteLLMParams(model="auto_router/auto_router/complexity_router"),
|
||||
existing_params=self._stored_complexity_params(),
|
||||
)
|
||||
assert violation is not None
|
||||
assert "repeats" in violation
|
||||
|
||||
def test_prefix_strip_rejected_against_stored_params(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_strategy_router_write_violation,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
violation = _strategy_router_write_violation(
|
||||
incoming_params=updateLiteLLMParams(model="complexity_router"),
|
||||
existing_params=self._stored_complexity_params(),
|
||||
)
|
||||
assert violation is not None
|
||||
assert "does not start with" in violation
|
||||
|
||||
def test_patch_without_model_is_not_judged(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_strategy_router_write_violation,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
assert (
|
||||
_strategy_router_write_violation(
|
||||
incoming_params=updateLiteLLMParams(rpm=10),
|
||||
existing_params=self._stored_complexity_params(),
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert _strategy_router_write_violation(incoming_params=None, existing_params=None) is None
|
||||
|
||||
def test_restore_of_corrupted_row_is_allowed(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_strategy_router_write_violation,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
corrupted = LiteLLM_Params(
|
||||
model="auto_router/auto_router/complexity_router",
|
||||
complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini"}},
|
||||
)
|
||||
assert (
|
||||
_strategy_router_write_violation(
|
||||
incoming_params=updateLiteLLMParams(model="auto_router/complexity_router"),
|
||||
existing_params=corrupted,
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
def test_create_semantic_router_missing_embedding_rejected(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_strategy_router_write_violation,
|
||||
)
|
||||
|
||||
violation = _strategy_router_write_violation(
|
||||
incoming_params=LiteLLM_Params(
|
||||
model="auto_router/my-router",
|
||||
auto_router_config="{}",
|
||||
auto_router_default_model="gpt-4o-mini",
|
||||
),
|
||||
existing_params=None,
|
||||
)
|
||||
assert violation is not None
|
||||
assert "auto_router_embedding_model" in violation
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_patch_model_rejects_double_prefix(self):
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
patch_model,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
model_id = "strategy-router-patch-test"
|
||||
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.get_db_model",
|
||||
new=AsyncMock(return_value=self._db_complexity_router(model_id)),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints._update_team_model_in_db",
|
||||
new=AsyncMock(),
|
||||
) as mock_update,
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await patch_model(
|
||||
model_id=model_id,
|
||||
patch_data=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(model="auto_router/auto_router/complexity_router")
|
||||
),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
assert "repeats" in str(exc_info.value.message)
|
||||
mock_update.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_new_model_rejects_prefixed_model_without_config(self):
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
add_new_model,
|
||||
)
|
||||
|
||||
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
mock_prisma = MagicMock()
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await add_new_model(
|
||||
model_params=Deployment(
|
||||
model_name="my-auto-router",
|
||||
litellm_params=LiteLLM_Params(model="auto_router/complexity_router"),
|
||||
model_info={"id": "strategy-router-create-test"},
|
||||
),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
assert "requires" in str(exc_info.value.message)
|
||||
mock_prisma.db.litellm_proxymodeltable.create.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_model_rejects_prefix_strip(self):
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_model,
|
||||
)
|
||||
from litellm.types.router import ModelInfo, updateLiteLLMParams
|
||||
|
||||
model_id = "strategy-router-update-test"
|
||||
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
existing_row = MagicMock()
|
||||
existing_row.model_dump.return_value = {
|
||||
"model_name": "my-auto-router",
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": {"tiers": {"SIMPLE": "gpt-4o-mini"}},
|
||||
},
|
||||
"model_info": {"id": model_id},
|
||||
}
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row)
|
||||
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await update_model(
|
||||
model_params=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(model="complexity_router"),
|
||||
model_info=ModelInfo(id=model_id),
|
||||
),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
assert "does not start with" in str(exc_info.value.message)
|
||||
mock_prisma.db.litellm_proxymodeltable.update.assert_not_awaited()
|
||||
|
|
|
|||
|
|
@ -0,0 +1,77 @@
|
|||
import pytest
|
||||
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
classify_strategy_router_model,
|
||||
validate_strategy_router_model_write,
|
||||
)
|
||||
|
||||
COMPLEXITY_FIELDS = frozenset({"complexity_router_config"})
|
||||
SEMANTIC_FIELDS = frozenset(
|
||||
{"auto_router_config", "auto_router_default_model", "auto_router_embedding_model"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,expected",
|
||||
[
|
||||
("anthropic/claude-sonnet-5", None),
|
||||
("complexity_router", None),
|
||||
("autorouter/complexity_router", None),
|
||||
("auto_router/my-router", "semantic"),
|
||||
("auto_router/complexity_router", "complexity"),
|
||||
("auto_router/complexity_router-eu", "complexity"),
|
||||
("auto_router/adaptive_router", "adaptive"),
|
||||
("auto_router/quality_router", "quality"),
|
||||
("auto_router/auto_router/complexity_router", "semantic"),
|
||||
("auto_router/", "semantic"),
|
||||
],
|
||||
)
|
||||
def test_classify_strategy_router_model(model, expected):
|
||||
assert classify_strategy_router_model(model) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,present_fields,expected_fragment",
|
||||
[
|
||||
("auto_router/auto_router/complexity_router", COMPLEXITY_FIELDS, "repeats"),
|
||||
("complexity_router", COMPLEXITY_FIELDS, "does not start with"),
|
||||
("anthropic/claude-sonnet-5", COMPLEXITY_FIELDS, "does not start with"),
|
||||
("auto_router/", frozenset(), "missing the router name"),
|
||||
("auto_router/complexity_router", frozenset(), "requires"),
|
||||
("auto_router/my-router", frozenset({"auto_router_config"}), "requires"),
|
||||
("auto_router/adaptive_router", frozenset(), "requires"),
|
||||
("auto_router/quality_router", frozenset(), "requires"),
|
||||
],
|
||||
)
|
||||
def test_validate_rejects_incoherent_writes(model, present_fields, expected_fragment):
|
||||
violation = validate_strategy_router_model_write(model=model, present_fields=present_fields)
|
||||
assert violation is not None
|
||||
assert expected_fragment in violation
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,present_fields",
|
||||
[
|
||||
("anthropic/claude-sonnet-5", frozenset()),
|
||||
("openai/gpt-4o-mini", frozenset({"api_key"})),
|
||||
("auto_router/complexity_router", COMPLEXITY_FIELDS),
|
||||
("auto_router/complexity_router", frozenset({"complexity_router_default_model"})),
|
||||
("auto_router/complexity_router-eu", COMPLEXITY_FIELDS),
|
||||
("auto_router/my-router", SEMANTIC_FIELDS),
|
||||
(
|
||||
"auto_router/my-router",
|
||||
frozenset(
|
||||
{
|
||||
"auto_router_config_path",
|
||||
"auto_router_default_model",
|
||||
"auto_router_embedding_model",
|
||||
}
|
||||
),
|
||||
),
|
||||
("auto_router/adaptive_router", frozenset({"adaptive_router_config"})),
|
||||
("auto_router/quality_router", frozenset({"quality_router_default_model"})),
|
||||
("auto_router/quality_router", frozenset({"quality_router_config"})),
|
||||
],
|
||||
)
|
||||
def test_validate_accepts_coherent_writes(model, present_fields):
|
||||
assert validate_strategy_router_model_write(model=model, present_fields=present_fields) is None
|
||||
Loading…
Add table
Reference in a new issue