fix(router): cast model_info cost values to float in _set_model_group_info

This commit is contained in:
Devin AI 2026-07-10 16:09:00 +00:00
parent bf02a4a47f
commit 1eae186b1c
4 changed files with 102 additions and 8 deletions

View file

@ -36,6 +36,29 @@ def safe_divide_seconds(seconds: float, denominator: float, default: Optional[fl
return float(seconds / denominator)
def safe_cast_to_float(value: Any, default: Optional[float] = None) -> Optional[float]:
"""
Cast a value to float, tolerating strings and non-numeric junk.
Cost fields round-tripped through the database (with store_model_in_db) can
come back as strings (e.g. "1e-05"), which breaks numeric comparisons. This
normalises them back to float.
Args:
value: The value to cast (float, int, str, or None)
default: Value to return when the input is None or not castable
Returns:
The value as a float, or default if it is None or cannot be cast
"""
if value is None:
return default
try:
return float(value)
except (TypeError, ValueError):
return default
def safe_divide(
numerator: Union[int, float],
denominator: Union[int, float],

View file

@ -67,6 +67,7 @@ from litellm.litellm_core_utils.request_timeout_resolver import (
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
get_metadata_variable_name_from_kwargs,
safe_cast_to_float,
)
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
@ -8719,8 +8720,8 @@ class Router:
# 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 = safe_cast_to_float(db_model_info.get("input_cost_per_token"))
output_cost_per_token = safe_cast_to_float(db_model_info.get("output_cost_per_token"))
model_info = ModelMapInfo(
key=model_group,
@ -8771,16 +8772,18 @@ class Router:
)
):
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 = safe_cast_to_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 > (safe_cast_to_float(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 = safe_cast_to_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 > (safe_cast_to_float(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

View file

@ -7,9 +7,34 @@ from litellm.litellm_core_utils.core_helpers import (
map_finish_reason,
reconstruct_model_name,
redact_nested_match_and_regex_keys,
safe_cast_to_float,
)
@pytest.mark.parametrize(
"value, expected",
[
("1e-05", 1e-05),
("0.0002", 0.0002),
(0.0002, 0.0002),
(5, 5.0),
(None, None),
("not-a-number", None),
({}, None),
],
)
def test_safe_cast_to_float(value, expected):
result = safe_cast_to_float(value)
assert result == expected
if expected is not None:
assert isinstance(result, float)
def test_safe_cast_to_float_custom_default():
assert safe_cast_to_float(None, default=0.0) == 0.0
assert safe_cast_to_float("junk", default=0.0) == 0.0
def test_reconstruct_model_name_prefers_deployment_value():
"""Ensure deployment metadata wins when reconstructing the model name."""

View file

@ -1455,6 +1455,49 @@ def test_model_group_info_cost_none_when_db_model_info_has_no_cost():
assert result.output_cost_per_token is None
def test_model_group_info_casts_string_cost_from_model_info():
"""
Regression for #32787: cost values round-tripped through the DB
(store_model_in_db) can come back as strings such as "1e-05". Comparing a
str against a float in _set_model_group_info raised
"TypeError: '>' not supported between instances of 'str' and 'float'".
Costs read from model_info must be cast to float.
"""
from unittest.mock import patch
router = litellm.Router(
model_list=[
{
"model_name": "grp",
"litellm_params": {"model": "openai/a", "api_key": "fake"},
"model_info": {"id": "d1"},
},
{
"model_name": "grp",
"litellm_params": {"model": "openai/b", "api_key": "fake"},
"model_info": {"id": "d2"},
},
]
)
def fake_model_info(model_id, model_name):
return {
"input_cost_per_token": "1e-05",
"output_cost_per_token": "2e-05",
}
with patch.object(
router, "get_deployment_model_info", side_effect=fake_model_info
):
result = router._cached_get_model_group_info("grp")
assert result is not None
assert isinstance(result.input_cost_per_token, float)
assert isinstance(result.output_cost_per_token, float)
assert result.input_cost_per_token == 1e-05
assert result.output_cost_per_token == 2e-05
def test_get_model_access_groups_caching():
"""
Test that get_model_access_groups caches the no-args result