diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 33fb069afbd..9806a824c63 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -571,13 +571,16 @@ async def preview_auto_router_routing( llm_router=llm_router, ) - complexity_router: Final = ComplexityRouter( - model_name=resolved.router_name, - litellm_router_instance=llm_router, - complexity_router_config=resolved.complexity_router_config.model_dump(exclude_none=True), - default_model=resolved.default_model, - derive_savings_baseline=False, - ) + try: + complexity_router: Final = ComplexityRouter( + model_name=resolved.router_name, + litellm_router_instance=llm_router, + complexity_router_config=resolved.complexity_router_config.model_dump(exclude_none=True), + default_model=resolved.complexity_router_config.resolve_default_model(resolved.default_model), + derive_savings_baseline=False, + ) + except ValueError as e: + raise HTTPException(status_code=400, detail={"error": f"Could not route this prompt: {e}"}) from e request_kwargs: Final = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata( data=request_data, diff --git a/litellm/router.py b/litellm/router.py index 0b9f12c8da3..afc842a05a0 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -9210,23 +9210,13 @@ class Router: if limit_violation is not None: raise ValueError(limit_violation) - 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 - # validated model, not the raw dict, so normalization (e.g. fallback_tier - # whitespace) is applied by its one owner before the tiers lookup. - if default_model is None and complexity_router_config: - validated: Final = ComplexityRouterConfig.model_validate(complexity_router_config) - # Custom tier sets name their fallback tier; built-in sets default to MEDIUM or SIMPLE - derived: Final = ( - (validated.tiers.get(validated.fallback_tier) if validated.fallback_tier is not None else None) - or validated.tiers.get("MEDIUM") - or validated.tiers.get("SIMPLE") + default_model: Final = ( + ComplexityRouterConfig.model_validate(complexity_router_config).resolve_default_model( + deployment.litellm_params.complexity_router_default_model ) - if isinstance(derived, list): - default_model = derived[0] if derived else None - else: - default_model = derived + if complexity_router_config + else deployment.litellm_params.complexity_router_default_model + ) if default_model is None: raise ValueError( diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 88907731468..8d00b45fed0 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -1955,6 +1955,20 @@ class ComplexityRouterConfig(BaseModel): def _normalize_classification_examples_field(cls, value: str | None) -> str | None: return normalize_classification_examples(value) + def resolve_default_model(self, default_model: str | None = None) -> str | None: + if default_model is not None: + return default_model + if self.default_model is not None: + return self.default_model + derived: Final = ( + (self.tiers.get(self.fallback_tier) if self.fallback_tier is not None else None) + or self.tiers.get("MEDIUM") + or self.tiers.get("SIMPLE") + ) + if isinstance(derived, list): + return derived[0] if derived else None + return derived + @property def has_custom_tiers(self) -> bool: """True when the operator replaced the built-in tier set via tier_definitions.""" diff --git a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py index 01e41e8b03f..0aff5852c96 100644 --- a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py @@ -150,6 +150,89 @@ async def _classifier_user_payload(body: Mapping[str, object], monkeypatch: pyte return router.recorded_calls[0]["messages"][1]["content"] +@pytest.mark.parametrize( + "tiers,config_default,explicit_default,expected", + ( + ({"SIMPLE": "cheap-model", "MEDIUM": "mid-model"}, None, None, "mid-model"), + ({"SIMPLE": ["cheap-model"], "MEDIUM": ["mid-model", "strong-model"]}, None, None, "mid-model"), + ({"SIMPLE": ["cheap-model"], "MEDIUM": []}, None, None, "cheap-model"), + ({"SIMPLE": "cheap-model"}, None, None, "cheap-model"), + ({"MEDIUM": "mid-model"}, "cheap-model", None, "cheap-model"), + ({"MEDIUM": "mid-model"}, "cheap-model", "strong-model", "strong-model"), + ({"MEDIUM": "mid-model"}, None, "strong-model", "strong-model"), + ({}, "cheap-model", None, "cheap-model"), + ({}, None, "strong-model", "strong-model"), + ), +) +@pytest.mark.asyncio +async def test_preview_and_serving_share_default_model_resolution( + monkeypatch: pytest.MonkeyPatch, + tiers: Mapping[str, object], + config_default: str | None, + explicit_default: str | None, + expected: str, +): + from litellm.router_utils.auto_router_model_naming import validate_complexity_router_config_write + from litellm.types.management_endpoints.auto_router_endpoints import ComplexityRouterConfigValidationRequest + + config: Final = { + "tiers": tiers, + "default_model": config_default, + "classifier_type": "llm", + "classifier_fallback": "default_model", + "classifier_llm_config": {"model": "unconfigured-classifier"}, + } + assert validate_complexity_router_config_write(config) is None + verdict: Final = await auto_router_endpoints.validate_complexity_router_config( + ComplexityRouterConfigValidationRequest(complexity_router_config=config), ADMIN + ) + assert verdict.valid and verdict.error is None + serving: Final = _router() + serving.init_complexity_router_deployment( + Deployment( + model_name="default-parity", + litellm_params={ + "model": "auto_router/complexity_router", + "complexity_router_config": config, + "complexity_router_default_model": explicit_default, + }, + model_info={"id": "default-parity"}, + ) + ) + strategy: Final = serving.complexity_routers["default-parity"][0].strategy + assert strategy.config.default_model == expected + decision: Final = await strategy.async_pre_routing_hook( + model="default-parity", messages=[{"role": "user", "content": "hello"}], request_kwargs={} + ) + assert decision is not None and decision.model == expected + monkeypatch.setattr(proxy_server, "llm_router", serving) + preview: Final = await preview_auto_router_routing( + http_request=ROUTING_HTTP_REQUEST, + data=AutoRouterRoutingTestRequest.model_validate( + {"prompt": "hello", "complexity_router_config": config, "default_model": explicit_default} + ), + user_api_key_dict=ADMIN, + ) + assert preview.routed_model == expected + assert preview.routing_decision["cause"] == "default_model_fallback" + + +@pytest.mark.asyncio +async def test_preview_missing_unresolvable_default_is_a_config_error(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(proxy_server, "llm_router", _router()) + with pytest.raises(HTTPException) as error: + await preview_auto_router_routing( + http_request=ROUTING_HTTP_REQUEST, + data=_request( + "hello", tiers={}, classifier_type="llm", classifier_fallback="default_model", + classifier_llm_config={"model": "unconfigured-classifier"}, + ), + user_api_key_dict=ADMIN, + ) + assert error.value.status_code == 400 + assert "requires a default model" in error.value.detail["error"] + + @pytest.mark.asyncio async def test_simple_prompt_routes_to_the_simple_tier(monkeypatch: pytest.MonkeyPatch): response = await _route("what is 2+2", monkeypatch) @@ -195,12 +278,17 @@ async def test_escalation_keyword_bumps_the_classified_tier(monkeypatch: pytest. assert response.routing_decision["escalation_keyword"] == "ultrathink" +@pytest.mark.parametrize("default_model,expected,configured", ((None, "mid-model", True), ("never-configured", "never-configured", False))) @pytest.mark.asyncio -async def test_tier_model_missing_from_the_proxy_is_reported(monkeypatch: pytest.MonkeyPatch): - response = await _route("what is 2+2", monkeypatch, tiers={**TIERS, "SIMPLE": ["never-configured"]}) +async def test_tier_model_missing_from_the_proxy_is_reported( + monkeypatch: pytest.MonkeyPatch, default_model: str | None, expected: str, configured: bool +): + response = await _route( + "what is 2+2", monkeypatch, tiers={**TIERS, "SIMPLE": ["never-configured"]}, default_model=default_model + ) - assert response.routed_model == "never-configured" - assert response.routed_model_configured is False + assert response.routed_model == expected + assert response.routed_model_configured is configured @pytest.mark.asyncio diff --git a/tests/unit/router_strategy/test_complexity_router.py b/tests/unit/router_strategy/test_complexity_router.py index 333524ffffc..dfeb9f8d961 100644 --- a/tests/unit/router_strategy/test_complexity_router.py +++ b/tests/unit/router_strategy/test_complexity_router.py @@ -1671,6 +1671,40 @@ class TestRouterComplexityDeploymentMethods: router.init_complexity_router_deployment(deployment) assert "auto_router/complexity_router/test-router" in router.complexity_routers + @pytest.mark.parametrize("tier_models", ("custom-model", ["custom-model", "other-model"])) + @pytest.mark.parametrize("explicit,expected", ((None, "configured-model"), ("top-model", "top-model"))) + def test_custom_default_resolution_preserves_explicit_precedence(self, tier_models, explicit, expected): + from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig + + config: Final = ComplexityRouterConfig.model_validate({ + "tiers": {"CUSTOM": tier_models, "OTHER": "other-model"}, + "tier_definitions": [ + {"name": "CUSTOM", "description": "Custom work"}, + {"name": "OTHER", "description": "Other work"}, + ], + "fallback_tier": " CUSTOM ", + "classifier_type": "llm", + "classifier_llm_config": {"model": "classifier"}, + "default_model": "configured-model", + }) + assert config.resolve_default_model(explicit) == expected + inferred: Final = config.model_copy(update={"default_model": None}) + assert inferred.resolve_default_model() == "custom-model" + assert config.default_model == "configured-model" + + @pytest.mark.parametrize("config", (None, {})) + def test_absent_deployment_config_still_requires_explicit_default(self, config): + from litellm.types.router import Deployment + + router: Final = Router(model_list=[]) + deployment: Final = Deployment( + model_name="no-default", + litellm_params={"model": "auto_router/complexity_router", "complexity_router_config": config}, + model_info={"id": "no-default"}, + ) + with pytest.raises(ValueError, match="complexity_router_default_model is required"): + router.init_complexity_router_deployment(deployment) + @staticmethod def _forecast_row(model_name: str, model_id: str, classifier_type: str) -> dict[str, object]: settings: Final = (