mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(router): match deployment pricing ids against the cost map only within the deployment's provider (#45472)
* fix(router): match deployment pricing ids against the cost map only within the deployment's provider A deployment model_info.id that equals another provider's catalog key (e.g. baseten/zai-org/glm-5.2 on an openai-compatible api_base) was merged into that provider's built-in row by register_model, keeping litellm_provider=baseten on the entry. _check_provider_match then rejected the row at request time and the deployment billed $0. register_model now takes a keyword-only custom_llm_provider used to scope _get_builtin_model_info_for_registration and _resolve_builtin_model_cost_entry, and the router passes the deployment's provider through. The stored entry stays provider-less for non-colliding ids. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(utils): pass custom_llm_provider straight to the registration lookup Drop the per-entry lookup-provider expression in register_model per review: the kwarg alone scopes _get_builtin_model_info_for_registration, and _resolve_builtin_model_cost_entry keeps its main signature and caller. _register_custom_pricing_for_request passes the provider through so router-originated per-request registrations get the same scoped lookup. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(utils): drop docstring and simplify set restore in per-request collision test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(router): type the colliding-id deployment fixture Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry <kerry@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
911aff2752
commit
fe1e8d1182
5 changed files with 166 additions and 4 deletions
|
|
@ -1319,6 +1319,7 @@ def _register_custom_pricing_for_request(
|
|||
},
|
||||
persist_across_reloads=False,
|
||||
warning_display_name=shared_key,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -10451,6 +10451,7 @@ class Router:
|
|||
model_cost={model_id: model_info},
|
||||
persist_across_reloads=False,
|
||||
warning_display_name=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
## OLD MODEL REGISTRATION ## Kept to prevent breaking changes
|
||||
|
|
|
|||
|
|
@ -3260,7 +3260,7 @@ def is_generalized_model_info(model_info: ModelInfo) -> bool:
|
|||
return key not in litellm.model_cost and match_capability_generalizations(key) is not None
|
||||
|
||||
|
||||
def _get_builtin_model_info_for_registration(model: str) -> ModelInfo | None:
|
||||
def _get_builtin_model_info_for_registration(model: str, custom_llm_provider: str | None) -> ModelInfo | None:
|
||||
"""Resolve ``model`` to its built-in cost-map entry for registration merging.
|
||||
|
||||
Returns ``None`` when the lookup raises or when it resolved via a
|
||||
|
|
@ -3269,7 +3269,7 @@ def _get_builtin_model_info_for_registration(model: str) -> ModelInfo | None:
|
|||
inheritance for prefix-mangled keys.
|
||||
"""
|
||||
try:
|
||||
info: Final = get_model_info(model=model)
|
||||
info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
except Exception:
|
||||
return None
|
||||
return None if is_generalized_model_info(info) else info
|
||||
|
|
@ -3344,6 +3344,7 @@ def register_model(
|
|||
*,
|
||||
persist_across_reloads: bool = True,
|
||||
warning_display_name: str | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
):
|
||||
"""
|
||||
Register new / Override existing models (and their pricing) to specific providers.
|
||||
|
|
@ -3368,6 +3369,10 @@ def register_model(
|
|||
``warning_display_name`` names the model in the missing-cache-pricing
|
||||
warning instead of the registered key, for callers that register under an
|
||||
opaque key (e.g. the router's hashed deployment ids).
|
||||
|
||||
``custom_llm_provider`` scopes the built-in cost-map match to entries for
|
||||
that provider, so a deployment id that happens to equal another provider's
|
||||
catalog key stays its own provider-less entry instead of merging into it.
|
||||
"""
|
||||
|
||||
loaded_model_cost = {}
|
||||
|
|
@ -3394,7 +3399,9 @@ def register_model(
|
|||
existing_model = litellm.model_cost.get(key, {})
|
||||
model_cost_key = key
|
||||
else:
|
||||
builtin_model_info = _get_builtin_model_info_for_registration(model=_key_str)
|
||||
builtin_model_info = _get_builtin_model_info_for_registration(
|
||||
model=_key_str, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
if builtin_model_info is not None:
|
||||
existing_model = cast(dict, builtin_model_info)
|
||||
model_cost_key = existing_model["key"]
|
||||
|
|
|
|||
|
|
@ -1037,3 +1037,64 @@ def test_update_model_cost():
|
|||
assert litellm.model_cost["gpt-4"]["input_cost_per_token"] == 0.00002
|
||||
except Exception as e:
|
||||
pytest.fail(f"An error occurred: {e}")
|
||||
|
||||
|
||||
def test_register_model_scopes_builtin_match_to_the_given_provider():
|
||||
"""Registering under an id that collides with another provider's catalog
|
||||
key keeps the entry provider-less under the id instead of merging into the
|
||||
baseten row."""
|
||||
colliding_id: Final = "baseten/zai-org/glm-5.2"
|
||||
builtin_key: Final = "baseten/zai-org/GLM-5.2"
|
||||
model_cost_entries: Final = _snapshot_model_cost_entries((colliding_id, builtin_key))
|
||||
builtin_row_before: Final = copy.deepcopy(litellm.model_cost[builtin_key])
|
||||
try:
|
||||
litellm.register_model(
|
||||
{
|
||||
colliding_id: {
|
||||
"input_cost_per_token": 0.00000096,
|
||||
"output_cost_per_token": 0.00000302,
|
||||
"cache_read_input_token_cost": 0.00000010,
|
||||
"mode": "chat",
|
||||
}
|
||||
},
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
entry: Final = litellm.model_cost[colliding_id]
|
||||
assert "litellm_provider" not in entry
|
||||
assert entry["input_cost_per_token"] == 0.00000096
|
||||
assert entry["output_cost_per_token"] == 0.00000302
|
||||
assert entry["cache_read_input_token_cost"] == 0.00000010
|
||||
assert litellm.model_cost[builtin_key] == builtin_row_before
|
||||
finally:
|
||||
_restore_model_cost_entries(model_cost_entries)
|
||||
|
||||
|
||||
def test_per_request_custom_pricing_scopes_a_colliding_deployment_id_to_its_provider():
|
||||
from litellm.main import _register_custom_pricing_for_request
|
||||
|
||||
deployment_id: Final = "baseten/zai-org/glm-5.2"
|
||||
builtin_key: Final = "baseten/zai-org/GLM-5.2"
|
||||
shared_key: Final = "openai/zai-org/GLM-5.2"
|
||||
model_cost_entries: Final = _snapshot_model_cost_entries((deployment_id, builtin_key, shared_key))
|
||||
builtin_row_before: Final = copy.deepcopy(litellm.model_cost[builtin_key])
|
||||
openai_models_before: Final = frozenset(litellm.open_ai_chat_completion_models)
|
||||
try:
|
||||
_register_custom_pricing_for_request(
|
||||
model="zai-org/GLM-5.2",
|
||||
custom_llm_provider="openai",
|
||||
kwargs={
|
||||
"input_cost_per_token": 0.00000096,
|
||||
"output_cost_per_token": 0.00000302,
|
||||
"metadata": {"model_info": {"id": deployment_id}},
|
||||
},
|
||||
model_info={"mode": "chat"},
|
||||
)
|
||||
|
||||
entry: Final = litellm.model_cost[deployment_id]
|
||||
assert entry["input_cost_per_token"] == 0.00000096
|
||||
assert entry["output_cost_per_token"] == 0.00000302
|
||||
assert litellm.model_cost[builtin_key] == builtin_row_before
|
||||
finally:
|
||||
_restore_model_cost_entries(model_cost_entries)
|
||||
litellm.open_ai_chat_completion_models.intersection_update(openai_models_before)
|
||||
|
|
|
|||
|
|
@ -26,7 +26,13 @@ from litellm.litellm_core_utils.ptu_pricing import ptu_config_error
|
|||
from litellm.litellm_core_utils.llm_cost_calc.utils import SERVICE_TIER_COST_KEY_SUFFIXES
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.openai_like.model_info import MODEL_INFO_REFRESH_SECONDS
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
from litellm.types.router import (
|
||||
Deployment,
|
||||
DeploymentTypedDict,
|
||||
LiteLLM_Params,
|
||||
LiteLLMParamsTypedDict,
|
||||
ModelInfo,
|
||||
)
|
||||
from litellm.utils import (
|
||||
_invalidate_model_cost_lowercase_map,
|
||||
reapply_runtime_model_cost_registrations,
|
||||
|
|
@ -3227,3 +3233,89 @@ def test_price_data_reload_refreshes_the_cached_model_group_and_deployment_info(
|
|||
|
||||
assert router.cached_model_group_info("grp").input_cost_per_token == new_price
|
||||
assert router.cached_deployment_model_info("dep-a", "openai/gpt-4o")["input_cost_per_token"] == new_price
|
||||
|
||||
|
||||
_COLLIDING_BUILTIN_KEY: Final = "baseten/zai-org/GLM-5.2"
|
||||
_COLLIDING_SHARED_OPENAI_KEY: Final = "openai/zai-org/GLM-5.2"
|
||||
_COLLIDING_MODEL_INFO: Final = {
|
||||
"input_cost_per_token": 0.00000096,
|
||||
"output_cost_per_token": 0.00000302,
|
||||
"cache_read_input_token_cost": 0.00000010,
|
||||
"mode": "chat",
|
||||
}
|
||||
|
||||
|
||||
def _colliding_id_deployment(model_id: str, custom_llm_provider: str) -> DeploymentTypedDict:
|
||||
litellm_params: Final = LiteLLMParamsTypedDict(
|
||||
model="zai-org/GLM-5.2",
|
||||
api_base="https://inference.baseten.co/v1",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
model_info: Final[dict] = {"id": model_id, **_COLLIDING_MODEL_INFO}
|
||||
return DeploymentTypedDict(
|
||||
model_name="nvidia/zai-org/glm-5.2",
|
||||
litellm_params=litellm_params,
|
||||
model_info=model_info,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_id", ("baseten/zai-org/glm-5.2",))
|
||||
def test_deployment_id_colliding_with_another_providers_catalog_key_keeps_its_own_pricing(model_id: str) -> None:
|
||||
"""A deployment id equal to another provider's catalog key must not merge
|
||||
into that row: the merge leaves litellm_provider=baseten on the entry, which
|
||||
_check_provider_match then rejects for the openai request, billing $0."""
|
||||
builtin_row_before: Final = copy.deepcopy(litellm.model_cost[_COLLIDING_BUILTIN_KEY])
|
||||
model_cost_entries: Final = {
|
||||
key: copy.deepcopy(litellm.model_cost.get(key))
|
||||
for key in (model_id, _COLLIDING_BUILTIN_KEY, _COLLIDING_SHARED_OPENAI_KEY)
|
||||
}
|
||||
try:
|
||||
router: Final = Router(model_list=[_colliding_id_deployment(model_id, "openai")])
|
||||
|
||||
response: Final = router.completion(
|
||||
model="nvidia/zai-org/glm-5.2",
|
||||
messages=[{"role": "user", "content": "colliding id pricing"}],
|
||||
mock_response=litellm.ModelResponse(
|
||||
model="zai-org/GLM-5.2",
|
||||
usage=litellm.Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500),
|
||||
),
|
||||
)
|
||||
|
||||
assert isinstance(response, litellm.ModelResponse)
|
||||
assert response._hidden_params["response_cost"] == pytest.approx(
|
||||
1000 * 0.00000096 + 500 * 0.00000302
|
||||
)
|
||||
assert litellm.model_cost[_COLLIDING_BUILTIN_KEY] == builtin_row_before
|
||||
finally:
|
||||
_restore_model_cost_entries(model_cost_entries)
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
|
||||
def test_deployment_id_matching_its_own_providers_catalog_key_still_merges() -> None:
|
||||
"""Same-provider collision keeps today's merge behavior: the deployment's
|
||||
custom prices land on the baseten row the openai-compatible baseten request
|
||||
matches against."""
|
||||
model_id: Final = "baseten/zai-org/glm-5.2"
|
||||
model_cost_entries: Final = {
|
||||
key: copy.deepcopy(litellm.model_cost.get(key))
|
||||
for key in (model_id, _COLLIDING_BUILTIN_KEY)
|
||||
}
|
||||
try:
|
||||
router: Final = Router(model_list=[_colliding_id_deployment(model_id, "baseten")])
|
||||
|
||||
response: Final = router.completion(
|
||||
model="nvidia/zai-org/glm-5.2",
|
||||
messages=[{"role": "user", "content": "same provider pricing"}],
|
||||
mock_response=litellm.ModelResponse(
|
||||
model="zai-org/GLM-5.2",
|
||||
usage=litellm.Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500),
|
||||
),
|
||||
)
|
||||
|
||||
assert isinstance(response, litellm.ModelResponse)
|
||||
assert response._hidden_params["response_cost"] == pytest.approx(
|
||||
1000 * 0.00000096 + 500 * 0.00000302
|
||||
)
|
||||
finally:
|
||||
_restore_model_cost_entries(model_cost_entries)
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue