fix(router): resolve provider-prefixed model costs

This commit is contained in:
林SO 2026-08-05 02:28:45 +08:00
parent e4fd790f1c
commit c3d69fd9fd
2 changed files with 59 additions and 4 deletions

View file

@ -1,7 +1,7 @@
#### 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
from typing import Final, TypedDict, cast
import litellm
from litellm import ModelResponse, token_counter, verbose_logger
@ -10,6 +10,30 @@ from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
class _ModelCostInfo(TypedDict, total=False):
input_cost_per_token: float | None
output_cost_per_token: float | 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)
if exact_model_cost is not None:
return exact_model_cost
try:
return cast( # cast-ok: only the two declared numeric fields are read below
_ModelCostInfo, litellm.get_model_info(model_name)
)
except Exception:
return {} # mutable-ok: each unresolved model gets an independent cost map
class LowestCostLoggingHandler(CustomLogger):
test_flag: bool = False
logged_success: int = 0
@ -245,7 +269,7 @@ class LowestCostLoggingHandler(CustomLogger):
or float("inf")
)
item_litellm_model_name = _deployment.get("litellm_params", {}).get("model")
item_litellm_model_cost_map = litellm.model_cost.get(item_litellm_model_name, {})
item_litellm_model_cost_map = _get_model_cost_info(item_litellm_model_name)
# check if user provided input_cost_per_token and output_cost_per_token in litellm_params
item_input_cost = None
@ -257,10 +281,12 @@ class LowestCostLoggingHandler(CustomLogger):
item_output_cost = _deployment.get("litellm_params", {}).get("output_cost_per_token")
if item_input_cost is None:
item_input_cost = item_litellm_model_cost_map.get("input_cost_per_token", 5.0)
model_input_cost = item_litellm_model_cost_map.get("input_cost_per_token")
item_input_cost = model_input_cost if model_input_cost is not None else 5.0
if item_output_cost is None:
item_output_cost = item_litellm_model_cost_map.get("output_cost_per_token", 5.0)
model_output_cost = item_litellm_model_cost_map.get("output_cost_per_token")
item_output_cost = model_output_cost if model_output_cost is not None else 5.0
# if litellm["model"] is not in model_cost map -> use item_cost = $10

View file

@ -0,0 +1,29 @@
import pytest
from litellm.caching.caching import DualCache
from litellm.router_strategy.lowest_cost import LowestCostLoggingHandler
@pytest.mark.asyncio
async def test_provider_prefixed_models_use_resolved_costs() -> None:
deployments = [
{
"model_name": "test-group",
"litellm_params": {"model": "openai/gpt-5.6-sol"},
"model_info": {"id": "sol"},
},
{
"model_name": "test-group",
"litellm_params": {"model": "openai/gpt-5.6-luna"},
"model_info": {"id": "luna"},
},
]
handler = LowestCostLoggingHandler(router_cache=DualCache())
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"