fix(router): avoid dynamic cost metadata lookup

This commit is contained in:
林SO 2026-08-05 05:04:09 +08:00
parent c3d69fd9fd
commit f9eb359558
2 changed files with 41 additions and 6 deletions

View file

@ -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

View file

@ -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()