From 6a8f5d580eae16aed9a3c9fe428b4c62336effd9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9E=97SO?= <142557582+Linxiushen@users.noreply.github.com> Date: Mon, 28 Sep 2026 01:14:31 +0800 Subject: [PATCH] fix(router): resolve provider-prefixed model costs --- litellm/router_strategy/lowest_cost.py | 49 +++++++++++++++++-- .../unit/router_strategy/test_lowest_cost.py | 45 +++++++++++++++++ 2 files changed, 91 insertions(+), 3 deletions(-) diff --git a/litellm/router_strategy/lowest_cost.py b/litellm/router_strategy/lowest_cost.py index 22c321c65fb..fa439a23332 100644 --- a/litellm/router_strategy/lowest_cost.py +++ b/litellm/router_strategy/lowest_cost.py @@ -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 diff --git a/tests/unit/router_strategy/test_lowest_cost.py b/tests/unit/router_strategy/test_lowest_cost.py index ab3ef099410..e203ee6c271 100644 --- a/tests/unit/router_strategy/test_lowest_cost.py +++ b/tests/unit/router_strategy/test_lowest_cost.py @@ -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):