mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(router): resolve provider-prefixed model costs
This commit is contained in:
parent
22b36cbcf6
commit
6a8f5d580e
2 changed files with 91 additions and 3 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue