From 1a163895849951056c3a145900400ca1b2e22cd4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9E=97SO?= <142557582+Linxiushen@users.noreply.github.com> Date: Wed, 5 Aug 2026 06:00:19 +0800 Subject: [PATCH] fix(router): validate static cost metadata --- litellm/router_strategy/lowest_cost.py | 24 +++++++++++---- .../router_strategy/test_lowest_cost.py | 29 +++++++++++++++++++ 2 files changed, 47 insertions(+), 6 deletions(-) diff --git a/litellm/router_strategy/lowest_cost.py b/litellm/router_strategy/lowest_cost.py index 25069c5d803..e72685ddc89 100644 --- a/litellm/router_strategy/lowest_cost.py +++ b/litellm/router_strategy/lowest_cost.py @@ -1,8 +1,9 @@ #### What this does #### # picks based on response time (for streaming, this is time to first token) from datetime import datetime, timedelta -from typing import Final, cast +from typing import Final +from pydantic import TypeAdapter, ValidationError from typing_extensions import TypedDict import litellm @@ -18,14 +19,25 @@ class _ModelCostInfo(TypedDict, total=False): litellm_provider: str | None +_MODEL_COST_INFO_ADAPTER: Final = TypeAdapter(_ModelCostInfo) + + +def _get_validated_model_cost_info(model_name: str) -> _ModelCostInfo | None: + raw_model_cost_info: Final[object] = litellm.model_cost.get(model_name) + if raw_model_cost_info is None: + return None + + try: + return _MODEL_COST_INFO_ADAPTER.validate_python(raw_model_cost_info) + except ValidationError: + return None + + def _get_model_cost_info(model_name: str | None) -> _ModelCostInfo: if model_name is None: return {} # mutable-ok: each unresolved model gets an independent cost map - model_cost = cast( # cast-ok: model_cost values come from the typed model cost JSON - dict[str, _ModelCostInfo], litellm.model_cost - ) - exact_model_cost = model_cost.get(model_name) + exact_model_cost: Final = _get_validated_model_cost_info(model_name) if exact_model_cost is not None: return exact_model_cost @@ -33,7 +45,7 @@ def _get_model_cost_info(model_name: str | None) -> _ModelCostInfo: if separator == "": return {} # mutable-ok: each unresolved model gets an independent cost map - unprefixed_model_cost = model_cost.get(unprefixed_model_name) + unprefixed_model_cost: Final = _get_validated_model_cost_info(unprefixed_model_name) if unprefixed_model_cost is None or unprefixed_model_cost.get("litellm_provider") != provider_name: return {} return unprefixed_model_cost diff --git a/tests/test_litellm/router_strategy/test_lowest_cost.py b/tests/test_litellm/router_strategy/test_lowest_cost.py index 78bbda388ca..b426e5bbde1 100644 --- a/tests/test_litellm/router_strategy/test_lowest_cost.py +++ b/tests/test_litellm/router_strategy/test_lowest_cost.py @@ -57,3 +57,32 @@ async def test_unknown_provider_model_does_not_query_dynamic_metadata() -> None: assert selected is not None assert selected["model_info"]["id"] == "luna" get_model_info.assert_not_called() + + +@pytest.mark.asyncio +async def test_invalid_static_cost_entry_is_ignored() -> None: + deployments = [ + { + "model_name": "test-group", + "litellm_params": {"model": "test-provider/corrupt-cost"}, + "model_info": {"id": "invalid"}, + }, + { + "model_name": "test-group", + "litellm_params": {"model": "openai/gpt-5.6-luna"}, + "model_info": {"id": "luna"}, + }, + ] + handler = LowestCostLoggingHandler(router_cache=DualCache()) + + with patch.dict( + litellm.model_cost, + {"test-provider/corrupt-cost": {"input_cost_per_token": "invalid"}}, + ): + selected = await handler.async_get_available_deployments( + model_group="test-group", + healthy_deployments=deployments, + ) + + assert selected is not None + assert selected["model_info"]["id"] == "luna"