fix(router): propagate custom cost_per_token from db model_info in fallback path (#25888)

This commit is contained in:
Hyogeun Oh (오효근)
2026-04-27 08:58:41 +05:30
committed by Sameer Kankute
parent 3f5e28fcdc
commit e68d5f86cf
2 changed files with 67 additions and 2 deletions
+4 -2
View File
@@ -8087,14 +8087,16 @@ 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")
model_info = ModelMapInfo(
key=model_group,
max_tokens=None,
max_input_tokens=None,
max_output_tokens=None,
input_cost_per_token=None,
output_cost_per_token=None,
input_cost_per_token=input_cost_per_token,
output_cost_per_token=output_cost_per_token,
litellm_provider=llm_provider,
mode=mode,
supported_openai_params=supported_openai_params,
+63
View File
@@ -1078,6 +1078,69 @@ def test_cached_get_model_group_info():
assert result5 is result6
def test_model_group_info_cost_from_db_model_info():
"""
When get_deployment_model_info fails (model_info is None fallback),
input_cost_per_token and output_cost_per_token should be read from db model_info.
"""
from unittest.mock import patch
router = litellm.Router(
model_list=[
{
"model_name": "my-custom-model",
"litellm_params": {
"model": "openai/my-custom-model",
"api_key": "fake",
"api_base": "https://my-custom-endpoint.com",
},
"model_info": {
"input_cost_per_token": 0.0001,
"output_cost_per_token": 0.0002,
},
},
]
)
with patch.object(
router, "get_deployment_model_info", side_effect=Exception("not found")
):
result = router._cached_get_model_group_info("my-custom-model")
assert result is not None
assert result.input_cost_per_token == 0.0001
assert result.output_cost_per_token == 0.0002
def test_model_group_info_cost_none_when_db_model_info_has_no_cost():
"""
When get_deployment_model_info fails and db model_info has no cost fields,
input/output_cost_per_token should be None.
"""
from unittest.mock import patch
router = litellm.Router(
model_list=[
{
"model_name": "my-custom-model-no-cost",
"litellm_params": {
"model": "openai/my-custom-model-no-cost",
"api_key": "fake",
"api_base": "https://my-custom-endpoint.com",
},
"model_info": {},
},
]
)
with patch.object(
router, "get_deployment_model_info", side_effect=Exception("not found")
):
result = router._cached_get_model_group_info("my-custom-model-no-cost")
assert result is not None
assert result.input_cost_per_token is None
assert result.output_cost_per_token is None
def test_get_model_access_groups_caching():
"""
Test that get_model_access_groups caches the no-args result