mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(router): validate static cost metadata
This commit is contained in:
parent
f9eb359558
commit
1a16389584
2 changed files with 47 additions and 6 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue