mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
Merge pull request #36222 from BerriAI/litellm_lit_5292_model_info_pricing_filter
fix(model_management): stop persisting cost map pricing as a deployment override
This commit is contained in:
commit
a969319fb5
4 changed files with 241 additions and 5 deletions
|
|
@ -123,6 +123,7 @@ from litellm.types.router import (
|
|||
ModelInfo,
|
||||
updateDeployment,
|
||||
)
|
||||
from litellm.types.utils import without_server_derived_pricing
|
||||
from litellm.utils import get_utc_datetime
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -747,11 +748,10 @@ def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> Pr
|
|||
|
||||
# update model info
|
||||
if updated_patch.model_info:
|
||||
merged_model_info.update(updated_patch.model_info.model_dump(exclude_none=True))
|
||||
merged_model_info.update(without_server_derived_pricing(updated_patch.model_info.model_dump(exclude_none=True)))
|
||||
|
||||
# Honor explicit-null clears LAST, after both merges, so a model_info blob the UI
|
||||
# passes through (which today re-sends the OLD pricing on every save) cannot
|
||||
# silently undo a litellm_params clear via .update().
|
||||
# Honor explicit-null clears LAST, after both merges, so a model_info blob a client
|
||||
# passes through cannot silently undo a litellm_params clear via .update().
|
||||
#
|
||||
# Restricted to SPECIAL_MODEL_INFO_PARAMS (input/output cost per token/character
|
||||
# and cache read/write costs) so this path cannot be used to null out privileged
|
||||
|
|
@ -2094,6 +2094,10 @@ async def add_new_model(
|
|||
enforced=bool(general_settings.get(ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING, False)),
|
||||
)
|
||||
|
||||
model_params.model_info = ModelInfo( # rebind-ok: downstream team-model handling mutates this same object
|
||||
**without_server_derived_pricing(model_params.model_info.model_dump(exclude_none=True))
|
||||
)
|
||||
|
||||
model_response: prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None = None
|
||||
# update DB
|
||||
incoming_model_info: Final = model_params.model_info.model_dump(exclude_none=True)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import json
|
||||
import re
|
||||
import time
|
||||
from collections.abc import Mapping, Sequence
|
||||
from enum import Enum
|
||||
|
|
@ -3646,6 +3647,38 @@ def shared_backend_model_info(model_info: dict[str, Any]) -> dict[str, Any]:
|
|||
return {k: v for k, v in model_info.items() if k in SHARED_BACKEND_MODEL_INFO_FIELDS}
|
||||
|
||||
|
||||
ABOVE_THRESHOLD_COST_KEY_PATTERN: Final = re.compile(r"_above_\d+k?_tokens$")
|
||||
|
||||
_PRICING_FIELD_EXEMPTIONS: Final[frozenset[str]] = frozenset({"output_vector_size"})
|
||||
|
||||
SERVER_DERIVED_PRICING_FIELDS: Final[frozenset[str]] = (
|
||||
frozenset(CustomPricingLiteLLMParams.model_fields) - _PRICING_FIELD_EXEMPTIONS
|
||||
)
|
||||
|
||||
|
||||
def is_server_derived_pricing_key(key: str) -> bool:
|
||||
"""Whether ``/model/info`` can fill ``key`` into ``model_info`` from the cost map.
|
||||
|
||||
Two sources, because ``get_model_info`` emits two: the declared pricing fields, and
|
||||
the tiered ``*_above_<N>_tokens`` rates that ride through on a pattern match and are
|
||||
declared nowhere. Both are read here from the same objects the read path uses, so the
|
||||
set cannot drift as new rates are added.
|
||||
"""
|
||||
return key in SERVER_DERIVED_PRICING_FIELDS or ABOVE_THRESHOLD_COST_KEY_PATTERN.search(key) is not None
|
||||
|
||||
|
||||
def without_server_derived_pricing(model_info: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
"""Drop the pricing ``/model/info`` derives for display, keeping everything else.
|
||||
|
||||
``/model/info`` fills a deployment's missing pricing in from the cost map so the
|
||||
Admin UI has a rate to show. Clients that echo that response back on save would
|
||||
otherwise persist the display value as a real per-deployment override, freezing the
|
||||
deployment at that day's price where no cost map refresh can reach it. A deployment's
|
||||
own pricing belongs on ``litellm_params``, which is unaffected.
|
||||
"""
|
||||
return MappingProxyType({k: v for k, v in model_info.items() if not is_server_derived_pricing_key(k)})
|
||||
|
||||
|
||||
# Server-controlled fields that bound or drive an interceptor's agentic loop
|
||||
# (depth, cycle fingerprints, ceiling, code-interpreter sandbox state). Listed
|
||||
# in all_litellm_params so they are treated as LiteLLM-level and excluded from
|
||||
|
|
|
|||
|
|
@ -216,6 +216,7 @@ from litellm.types.llms.openai import (
|
|||
OpenAIWebSearchOptions,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
ABOVE_THRESHOLD_COST_KEY_PATTERN,
|
||||
OPENAI_RESPONSE_HEADERS,
|
||||
CallTypes,
|
||||
ChatCompletionDeltaToolCall,
|
||||
|
|
@ -5643,7 +5644,7 @@ def _is_potential_model_name_in_model_cost(
|
|||
)
|
||||
|
||||
|
||||
_ABOVE_THRESHOLD_COST_KEY: Final = re.compile(r"_above_\d+k?_tokens$")
|
||||
_ABOVE_THRESHOLD_COST_KEY: Final = ABOVE_THRESHOLD_COST_KEY_PATTERN
|
||||
|
||||
|
||||
def _get_model_info_helper(
|
||||
|
|
|
|||
|
|
@ -3666,6 +3666,204 @@ class TestUpdateDBModelClearPricing:
|
|||
assert info["cache_creation_input_token_cost"] == 0.000003
|
||||
|
||||
|
||||
class TestModelInfoServerDerivedPricingFilter:
|
||||
"""LIT-5292. `/model/info` fills a deployment's missing pricing in from the cost map
|
||||
so the Admin UI has a rate to display. Clients echo that whole blob back on save, so
|
||||
without a write-path filter an unrelated edit persists the display value as a real
|
||||
per-deployment override and no cost map refresh can move the deployment again.
|
||||
|
||||
A deployment's own pricing rides `litellm_params`, which stays writable.
|
||||
"""
|
||||
|
||||
def test_echoed_cost_map_pricing_is_not_persisted(self):
|
||||
"""The ticket's repro: a deployment with no override, edited for an unrelated
|
||||
reason, must not gain one from the pricing the form was displaying."""
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_db_model,
|
||||
)
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
db_model = Deployment(
|
||||
model_name="haiku",
|
||||
litellm_params=LiteLLM_Params(model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"),
|
||||
model_info=ModelInfo(id="dep-unpriced-0"),
|
||||
)
|
||||
|
||||
result = update_db_model(
|
||||
db_model=db_model,
|
||||
updated_patch=updateDeployment(
|
||||
model_info=ModelInfo(
|
||||
id="dep-unpriced-0",
|
||||
access_groups=["prod"],
|
||||
input_cost_per_token=0.0000008,
|
||||
output_cost_per_token=0.000004,
|
||||
cache_read_input_token_cost=0.00000008,
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
info = json.loads(result["model_info"])
|
||||
params = json.loads(result["litellm_params"])
|
||||
assert info["access_groups"] == ["prod"]
|
||||
for field in ("input_cost_per_token", "output_cost_per_token", "cache_read_input_token_cost"):
|
||||
assert field not in info, f"{field} was persisted as a per-deployment override"
|
||||
assert field not in params
|
||||
|
||||
def test_tiered_above_threshold_pricing_is_dropped(self):
|
||||
"""Tiered rates ride `get_model_info` on a pattern match and are declared on no
|
||||
model, so a filter built only from the declared pricing fields would miss them."""
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_db_model,
|
||||
)
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
db_model = Deployment(
|
||||
model_name="sonnet",
|
||||
litellm_params=LiteLLM_Params(model="claude-sonnet-4-5"),
|
||||
model_info=ModelInfo(id="dep-tiered-0"),
|
||||
)
|
||||
|
||||
result = update_db_model(
|
||||
db_model=db_model,
|
||||
updated_patch=updateDeployment(
|
||||
model_info=ModelInfo(
|
||||
id="dep-tiered-0",
|
||||
input_cost_per_token_above_200k_tokens=0.000006,
|
||||
cache_creation_input_token_cost_above_1hr_above_200k_tokens=0.000012,
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
info = json.loads(result["model_info"])
|
||||
assert "input_cost_per_token_above_200k_tokens" not in info
|
||||
assert "cache_creation_input_token_cost_above_1hr_above_200k_tokens" not in info
|
||||
|
||||
def test_output_vector_size_and_client_owned_fields_survive(self):
|
||||
"""`output_vector_size` sits on the pricing model but is an embedding dimension,
|
||||
not a rate. It and the operator-owned keys stay writable."""
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_db_model,
|
||||
)
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
db_model = Deployment(
|
||||
model_name="embed",
|
||||
litellm_params=LiteLLM_Params(model="openai/text-embedding-3-large"),
|
||||
model_info=ModelInfo(id="dep-embed-0"),
|
||||
)
|
||||
|
||||
result = update_db_model(
|
||||
db_model=db_model,
|
||||
updated_patch=updateDeployment(
|
||||
model_info=ModelInfo(
|
||||
id="dep-embed-0",
|
||||
output_vector_size=3072,
|
||||
base_model="azure/text-embedding-3-large",
|
||||
tier="paid",
|
||||
team_id="team-1",
|
||||
access_groups=["research"],
|
||||
my_custom_key="my_custom_value",
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
info = json.loads(result["model_info"])
|
||||
assert info["output_vector_size"] == 3072
|
||||
assert info["base_model"] == "azure/text-embedding-3-large"
|
||||
assert info["tier"] == "paid"
|
||||
assert info["team_id"] == "team-1"
|
||||
assert info["access_groups"] == ["research"]
|
||||
assert info["my_custom_key"] == "my_custom_value"
|
||||
|
||||
def test_litellm_params_pricing_still_persists(self):
|
||||
"""The supported way to set a deployment override is untouched."""
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_db_model,
|
||||
)
|
||||
from litellm.types.router import (
|
||||
Deployment,
|
||||
LiteLLM_Params,
|
||||
ModelInfo,
|
||||
updateLiteLLMParams,
|
||||
)
|
||||
|
||||
db_model = Deployment(
|
||||
model_name="haiku",
|
||||
litellm_params=LiteLLM_Params(model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"),
|
||||
model_info=ModelInfo(id="dep-priced-0"),
|
||||
)
|
||||
|
||||
result = update_db_model(
|
||||
db_model=db_model,
|
||||
updated_patch=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(input_cost_per_token=0.00000123),
|
||||
model_info=ModelInfo(id="dep-priced-0", input_cost_per_token=0.0000008),
|
||||
),
|
||||
)
|
||||
|
||||
params = json.loads(result["litellm_params"])
|
||||
assert params["input_cost_per_token"] == 0.00000123
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_new_model_drops_echoed_pricing_and_keeps_identity(self):
|
||||
"""The create path filters too, and rebuilding the blob must not mint a fresh id
|
||||
or flip `db_model`, which would detach the row from its router deployment."""
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
add_new_model,
|
||||
)
|
||||
|
||||
model_id = "dep-create-0"
|
||||
db_row = LiteLLM_ProxyModelTable(
|
||||
model_id=model_id,
|
||||
model_name="haiku",
|
||||
litellm_params={"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"},
|
||||
model_info={"id": model_id},
|
||||
created_by="test-admin",
|
||||
updated_by="test-admin",
|
||||
)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
|
||||
mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=db_row)
|
||||
|
||||
mock_proxy_config = MagicMock()
|
||||
mock_proxy_config.add_deployment = AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None))
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_model_ids.return_value = [model_id]
|
||||
|
||||
_PS = "litellm.proxy.proxy_server"
|
||||
_ENCRYPT = "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper"
|
||||
with (
|
||||
patch(f"{_PS}.prisma_client", mock_prisma),
|
||||
patch(f"{_PS}.store_model_in_db", True),
|
||||
patch(f"{_PS}.proxy_config", mock_proxy_config),
|
||||
patch(f"{_PS}.proxy_logging_obj", MagicMock()),
|
||||
patch(f"{_PS}.general_settings", {}),
|
||||
patch(f"{_PS}.premium_user", True),
|
||||
patch(f"{_PS}.llm_router", mock_router),
|
||||
patch(_ENCRYPT, side_effect=lambda value, **kwargs: value),
|
||||
):
|
||||
await add_new_model(
|
||||
model_params=Deployment(
|
||||
model_name="haiku",
|
||||
litellm_params=LiteLLM_Params(model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"),
|
||||
model_info={
|
||||
"id": model_id,
|
||||
"access_groups": ["prod"],
|
||||
"input_cost_per_token": 0.0000008,
|
||||
},
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
written = json.loads(mock_prisma.db.litellm_proxymodeltable.create.call_args.kwargs["data"]["model_info"])
|
||||
assert "input_cost_per_token" not in written
|
||||
assert written["id"] == model_id, "filtering must not mint a fresh deployment id"
|
||||
assert written["access_groups"] == ["prod"]
|
||||
|
||||
|
||||
class TestGetModelInfoWithIdBlocked:
|
||||
"""`ProxyConfig.get_model_info_with_id` must propagate the DB-level `blocked`
|
||||
column into the in-memory `model_info` dict so the router filter can read it."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue