diff --git a/tests/unit/test_router_nested_model_info_pricing.py b/tests/unit/test_router_nested_model_info_pricing.py index 02412b49b0c..4657a47e5ca 100644 --- a/tests/unit/test_router_nested_model_info_pricing.py +++ b/tests/unit/test_router_nested_model_info_pricing.py @@ -3,12 +3,14 @@ import pytest import litellm from litellm import Router from litellm.router import _model_info_nested_under_litellm_params +from litellm.types.router import Deployment, LiteLLM_Params BACKEND_MODEL = "deepinfra/deepseek-ai/DeepSeek-V4-Flash-0731" PUBLIC_MODEL = "deepseek-ai/DeepSeek-V4-Flash-0731" INPUT_COST = 7.2e-08 OUTPUT_COST = 1.44e-07 CACHE_READ_COST = 1.44e-08 +DEPLOYMENT_ID = "db-deployment-id" def _pricing_model_info(): @@ -46,17 +48,23 @@ def _deployment_level_config(): @pytest.fixture(autouse=True) def clean_cost_map(): - """Keep registrations from bleeding across tests.""" + """Keep registrations from bleeding across tests, and put back the entries the cost map already had.""" from litellm.utils import _invalidate_model_cost_lowercase_map + cost_map = litellm.model_cost + keys = (BACKEND_MODEL, PUBLIC_MODEL, DEPLOYMENT_ID) + original = {key: cost_map[key] for key in keys if key in cost_map} + def _clear(): - litellm.model_cost.pop(BACKEND_MODEL, None) - litellm.model_cost.pop(PUBLIC_MODEL, None) + for key in keys: + cost_map.pop(key, None) _invalidate_model_cost_lowercase_map() _clear() yield _clear() + cost_map.update(original) + _invalidate_model_cost_lowercase_map() def _deployment_pricing(router): @@ -64,6 +72,16 @@ def _deployment_pricing(router): return deployment.get("model_info", {}) +def _nested_deployment(model_info): + """A deployment built the way the proxy loads one from the DB, nested block included.""" + config = _nested_config() + return Deployment( + model_name=config["model_name"], + litellm_params=LiteLLM_Params.model_validate(config["litellm_params"]), + model_info=model_info, + ) + + def test_promotes_nothing_when_no_nested_block(): """No nested block, and a non-dict nested value, both promote nothing.""" assert _model_info_nested_under_litellm_params({"model": BACKEND_MODEL}, {}) == {} @@ -175,3 +193,45 @@ async def test_nested_model_info_produces_non_zero_response_cost(): response_cost = response._hidden_params.get("response_cost") assert response_cost is not None assert response_cost > 0 + + +@pytest.mark.asyncio +async def test_add_deployment_applies_nested_model_info_pricing(): + """ + DB-loaded and /model/new deployments reach the router through add_deployment, + not _create_deployment, so the nested block has to be applied there as well. + """ + router = Router(model_list=[]) + assert router.add_deployment(_nested_deployment({"id": DEPLOYMENT_ID})) is not None + + model_info = _deployment_pricing(router) + assert model_info.get("input_cost_per_token") == INPUT_COST + assert model_info.get("output_cost_per_token") == OUTPUT_COST + assert model_info.get("cache_read_input_token_cost") == CACHE_READ_COST + + # Exact, not just > 0: the shipped cost map also prices this backend model. + response = await router.acompletion( + model=PUBLIC_MODEL, + messages=[{"role": "user", "content": "hi"}], + mock_response="hello", + ) + usage = response.usage + assert response._hidden_params.get("response_cost") == pytest.approx( + usage.prompt_tokens * INPUT_COST + usage.completion_tokens * OUTPUT_COST + ) + + +def test_add_deployment_keeps_deployment_level_values_and_id(): + """Only unset fields are applied, and the nested `id` never re-keys the deployment.""" + deployment = _nested_deployment({"id": DEPLOYMENT_ID, "input_cost_per_token": INPUT_COST}) + deployment.litellm_params.model_info.update(id="nested-id-should-be-ignored", input_cost_per_token=999.0) + + router = Router(model_list=[]) + router.add_deployment(deployment) + + model_info = _deployment_pricing(router) + assert model_info.get("input_cost_per_token") == INPUT_COST + assert model_info.get("output_cost_per_token") == OUTPUT_COST + assert model_info.get("id") == DEPLOYMENT_ID + assert router.has_model_id(DEPLOYMENT_ID) + assert not router.has_model_id("nested-id-should-be-ignored")