fix(router): validate static cost metadata

This commit is contained in:
林SO 2026-08-05 06:00:19 +08:00
parent f9eb359558
commit 1a16389584
2 changed files with 47 additions and 6 deletions

View file

@ -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

View file

@ -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"