mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(auto-router): align preview and serving default resolution (#44436)
Co-authored-by: moyai-devin-berriai[bot] <336287033+moyai-devin-berriai[bot]@users.noreply.github.com>
This commit is contained in:
parent
3286782dea
commit
2e81db03b9
5 changed files with 156 additions and 27 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue