Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 19 additions & 8 deletions litellm/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -251,6 +251,15 @@
PreRoutingHookResponse = Any


def _cost_value_as_float(value: Union[str, int, float, None]) -> Optional[float]:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 The new helper uses the older Union[...] / Optional[...] typing form. This repo prefers PEP 604 union syntax (str | int | float | None, float | None) for new code.

Suggested change
def _cost_value_as_float(value: Union[str, int, float, None]) -> Optional[float]:
def _cost_value_as_float(value: str | int | float | None) -> float | None:

Rule Used: In this repo, prefer str | None (PEP 604 union s... (source)

Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!

if value is None:
return None
try:
return float(value)
except (TypeError, ValueError):
return None


class RoutingArgs(enum.Enum):
ttl = 60 # 1min (RPM/TPM expire key)

Expand Down Expand Up @@ -8750,8 +8759,8 @@ def _set_model_group_info(self, model_group: str, user_facing_model_group_name:
# Get mode from database model_info if available, otherwise default to "chat"
db_model_info = model.get("model_info", {})
mode = db_model_info.get("mode", "chat")
input_cost_per_token = db_model_info.get("input_cost_per_token")
output_cost_per_token = db_model_info.get("output_cost_per_token")
input_cost_per_token = _cost_value_as_float(db_model_info.get("input_cost_per_token"))
output_cost_per_token = _cost_value_as_float(db_model_info.get("output_cost_per_token"))

model_info = ModelMapInfo(
key=model_group,
Expand Down Expand Up @@ -8802,16 +8811,18 @@ def _set_model_group_info(self, model_group: str, user_facing_model_group_name:
)
):
model_group_info.max_output_tokens = model_info["max_output_tokens"]
if model_info.get("input_cost_per_token", None) is not None and (
_input_cost_per_token = _cost_value_as_float(model_info.get("input_cost_per_token"))
if _input_cost_per_token is not None and (
model_group_info.input_cost_per_token is None
or (model_info["input_cost_per_token"] or 0.0) > (model_group_info.input_cost_per_token or 0.0)
or _input_cost_per_token > (model_group_info.input_cost_per_token or 0.0)
):
model_group_info.input_cost_per_token = model_info["input_cost_per_token"]
if model_info.get("output_cost_per_token", None) is not None and (
model_group_info.input_cost_per_token = _input_cost_per_token
_output_cost_per_token = _cost_value_as_float(model_info.get("output_cost_per_token"))
if _output_cost_per_token is not None and (
model_group_info.output_cost_per_token is None
or (model_info["output_cost_per_token"] or 0.0) > (model_group_info.output_cost_per_token or 0.0)
or _output_cost_per_token > (model_group_info.output_cost_per_token or 0.0)
):
model_group_info.output_cost_per_token = model_info["output_cost_per_token"]
model_group_info.output_cost_per_token = _output_cost_per_token
if (
model_info.get("supports_parallel_function_calling", None) is not None
and model_info["supports_parallel_function_calling"] is True # type: ignore
Expand Down
126 changes: 126 additions & 0 deletions tests/test_litellm/test_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -1455,6 +1455,132 @@ def test_model_group_info_cost_none_when_db_model_info_has_no_cost():
assert result.output_cost_per_token is None


@pytest.mark.parametrize(
"value,expected",
[
("1e-05", 1e-05),
("0.00001", 1e-05),
(1e-05, 1e-05),
(5, 5.0),
(None, None),
("not-a-number", None),
],
)
def test_cost_value_as_float(value, expected):
from litellm.router import _cost_value_as_float

assert _cost_value_as_float(value) == expected


def test_model_group_info_with_stringified_cost_values():
"""
YAML 1.2 parsers emit '1e-05' (integer mantissa) as a string, so cost
values in deployment model_info can arrive as str. Aggregating the model
group must not raise TypeError('>' between str and float) and must return
float costs.
"""
router = litellm.Router(
model_list=[
{
"model_name": "my-custom-model",
"litellm_params": {
"model": "openai/my-custom-backend-1",
"api_key": "fake",
},
"model_info": {
"input_cost_per_token": "1e-05",
"output_cost_per_token": "1e-05",
},
},
{
"model_name": "my-custom-model",
"litellm_params": {
"model": "openai/my-custom-backend-2",
"api_key": "fake",
},
"model_info": {
"input_cost_per_token": "2e-05",
"output_cost_per_token": "2e-05",
},
},
]
)

def _model_info_with_str_costs(model_id: str, model_name: str):
for model in router.model_list:
if model["model_info"]["id"] == model_id:
return {
"key": model_name,
"input_cost_per_token": model["model_info"]["input_cost_per_token"],
"output_cost_per_token": model["model_info"]["output_cost_per_token"],
"litellm_provider": "openai",
"mode": "chat",
}
return None

with patch.object(
router, "get_deployment_model_info", side_effect=_model_info_with_str_costs
):
result = router._set_model_group_info(
model_group="my-custom-model",
user_facing_model_group_name="my-custom-model",
)

assert result is not None
assert result.input_cost_per_token == 2e-05
assert result.output_cost_per_token == 2e-05
assert isinstance(result.input_cost_per_token, float)
assert isinstance(result.output_cost_per_token, float)


def test_model_group_info_db_fallback_with_stringified_cost_values():
"""
Fallback path: when get_deployment_model_info returns nothing, costs are
read straight from the deployment's model_info dict, which can hold
stringified floats parsed from YAML. They must be coerced to float.
"""
router = litellm.Router(
model_list=[
{
"model_name": "my-custom-model",
"litellm_params": {
"model": "openai/my-custom-backend-1",
"api_key": "fake",
},
"model_info": {
"input_cost_per_token": "1e-05",
"output_cost_per_token": "3e-05",
},
},
{
"model_name": "my-custom-model",
"litellm_params": {
"model": "openai/my-custom-backend-2",
"api_key": "fake",
},
"model_info": {
"input_cost_per_token": "2e-05",
"output_cost_per_token": "2e-05",
},
},
]
)

with patch.object(
router, "get_deployment_model_info", side_effect=Exception("not found")
):
result = router._set_model_group_info(
model_group="my-custom-model",
user_facing_model_group_name="my-custom-model",
)

assert result is not None
assert result.input_cost_per_token == 2e-05
assert result.output_cost_per_token == 3e-05
assert isinstance(result.input_cost_per_token, float)
assert isinstance(result.output_cost_per_token, float)


def test_get_model_access_groups_caching():
"""
Test that get_model_access_groups caches the no-args result
Expand Down
Loading