mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-18 04:28:19 +00:00
fix(router): propagate custom cost_per_token from db model_info in fallback path (#25888)
This commit is contained in:
+4
-2
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user