mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(router): avoid dynamic cost metadata lookup
This commit is contained in:
parent
c3d69fd9fd
commit
f9eb359558
2 changed files with 41 additions and 6 deletions
|
|
@ -1,7 +1,9 @@
|
|||
#### What this does ####
|
||||
# picks based on response time (for streaming, this is time to first token)
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Final, TypedDict, cast
|
||||
from typing import Final, cast
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm import ModelResponse, token_counter, verbose_logger
|
||||
|
|
@ -13,6 +15,7 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
class _ModelCostInfo(TypedDict, total=False):
|
||||
input_cost_per_token: float | None
|
||||
output_cost_per_token: float | None
|
||||
litellm_provider: str | None
|
||||
|
||||
|
||||
def _get_model_cost_info(model_name: str | None) -> _ModelCostInfo:
|
||||
|
|
@ -26,13 +29,15 @@ def _get_model_cost_info(model_name: str | None) -> _ModelCostInfo:
|
|||
if exact_model_cost is not None:
|
||||
return exact_model_cost
|
||||
|
||||
try:
|
||||
return cast( # cast-ok: only the two declared numeric fields are read below
|
||||
_ModelCostInfo, litellm.get_model_info(model_name)
|
||||
)
|
||||
except Exception:
|
||||
provider_name, separator, unprefixed_model_name = model_name.partition("/")
|
||||
if separator == "":
|
||||
return {} # mutable-ok: each unresolved model gets an independent cost map
|
||||
|
||||
unprefixed_model_cost = model_cost.get(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
|
||||
|
|
|
|||
|
|
@ -1,3 +1,6 @@
|
|||
from unittest.mock import patch
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -27,3 +30,30 @@ async def test_provider_prefixed_models_use_resolved_costs() -> None:
|
|||
|
||||
assert selected is not None
|
||||
assert selected["model_info"]["id"] == "luna"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unknown_provider_model_does_not_query_dynamic_metadata() -> None:
|
||||
deployments = [
|
||||
{
|
||||
"model_name": "test-group",
|
||||
"litellm_params": {"model": "ollama/caller-controlled-model"},
|
||||
"model_info": {"id": "unknown"},
|
||||
},
|
||||
{
|
||||
"model_name": "test-group",
|
||||
"litellm_params": {"model": "openai/gpt-5.6-luna"},
|
||||
"model_info": {"id": "luna"},
|
||||
},
|
||||
]
|
||||
handler = LowestCostLoggingHandler(router_cache=DualCache())
|
||||
|
||||
with patch.object(litellm, "get_model_info") as get_model_info:
|
||||
selected = await handler.async_get_available_deployments(
|
||||
model_group="test-group",
|
||||
healthy_deployments=deployments,
|
||||
)
|
||||
|
||||
assert selected is not None
|
||||
assert selected["model_info"]["id"] == "luna"
|
||||
get_model_info.assert_not_called()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue