fix(router): resolve provider-prefixed model costs

This commit is contained in:
林SO 2026-09-28 01:14:31 +08:00
parent 22b36cbcf6
commit 6a8f5d580e
2 changed files with 91 additions and 3 deletions

View file

@ -3,6 +3,9 @@
from datetime import datetime
from typing import Final
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm import ModelResponse, token_counter, verbose_logger
from litellm._logging import verbose_router_logger
@ -11,6 +14,44 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
class _ModelCostInfo(TypedDict, total=False):
input_cost_per_token: ReadOnly[float | None]
output_cost_per_token: ReadOnly[float | None]
litellm_provider: ReadOnly[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 {}
exact_model_cost: Final = _get_validated_model_cost_info(model_name)
if exact_model_cost is not None:
return exact_model_cost
provider_name, separator, unprefixed_model_name = model_name.partition("/")
if separator == "":
return {}
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
class LowestCostLoggingHandler(CustomLogger):
test_flag: bool = False
logged_success: int = 0
@ -240,7 +281,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
@ -252,10 +293,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

@ -1,5 +1,7 @@
import copy
from datetime import datetime
from typing import Final
from unittest.mock import patch
import pytest
@ -47,6 +49,49 @@ def test_log_success_event_counts_a_response_with_no_completion_tokens():
assert _recorded_minute_counters(cache) == {"tpm": 12, "rpm": 1}
@pytest.mark.parametrize(
("candidate_model", "candidate_price", "expected_id"),
(
("openai/cheap", None, "candidate"),
("other/cheap", None, "reference"),
(None, None, "reference"),
("unknown", None, "reference"),
("ollama/unknown", None, "reference"),
("custom/cheap", {"input_cost_per_token": 0.01, "output_cost_per_token": 0.02}, "candidate"),
("openai/cheap", {"input_cost_per_token": 2.0, "output_cost_per_token": 3.0}, "reference"),
("custom/broken", {"input_cost_per_token": "invalid"}, "reference"),
("custom/null", {"input_cost_per_token": None, "output_cost_per_token": None}, "reference"),
),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("reverse_order", [False, True], ids=["candidate-first", "reference-first"])
async def test_provider_prefix_uses_matching_static_prices_and_preserves_exact_entries(
candidate_model: str | None,
candidate_price: dict[str, float | str | None] | None,
expected_id: str,
reverse_order: bool,
) -> None:
deployments: Final = [
{"litellm_params": {"model": candidate_model}, "model_info": {"id": "candidate"}},
{"litellm_params": {"model": "openai/reference"}, "model_info": {"id": "reference"}},
]
prices: Final = {
"cheap": {"input_cost_per_token": 0.01, "output_cost_per_token": 0.02, "litellm_provider": "openai"},
"reference": {"input_cost_per_token": 0.2, "output_cost_per_token": 0.3, "litellm_provider": "openai"},
**({candidate_model: candidate_price} if candidate_model is not None and candidate_price is not None else {}),
}
handler: Final = LowestCostLoggingHandler(router_cache=DualCache())
with patch.dict(litellm.model_cost, prices, clear=True), patch("litellm.get_model_info") as metadata_lookup:
selected: Final = await handler.async_get_available_deployments(
model_group="test-group", healthy_deployments=deployments[::-1] if reverse_order else deployments
)
expected: Final = next(deployment for deployment in deployments if deployment["model_info"]["id"] == expected_id)
assert selected is expected
metadata_lookup.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("use_async", [False, True], ids=["sync", "async"])
async def test_log_success_event_keeps_cost_bookkeeping_out_of_the_latency_routing_entry(use_async: bool):