From fc268b19aaff87dfcadee076715a205d38f60770 Mon Sep 17 00:00:00 2001 From: kerry Date: Sun, 27 Sep 2026 00:51:26 +0000 Subject: [PATCH] fix(databricks): keep the served service_tier on streamed chunks and bill it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/cost_calculator.py | 2 +- .../llms/databricks/chat/transformation.py | 20 ++++++++------ litellm/llms/databricks/cost_calculator.py | 3 ++- .../test_databricks_chat_transformation.py | 10 +++++++ .../test_databricks_cost_calculator.py | 27 +++++++++++++++++-- 5 files changed, 50 insertions(+), 12 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index a279b9f0903..bf67c32436f 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -696,7 +696,7 @@ def cost_per_token( data_residency=data_residency, ) elif custom_llm_provider == "databricks": - return databricks_cost_per_token(model=model, usage=usage_block) + return databricks_cost_per_token(model=model, usage=usage_block, service_tier=service_tier) elif custom_llm_provider == "fireworks_ai": return fireworks_ai_cost_per_token(model=model, usage=usage_block) elif custom_llm_provider == "azure": diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index 538904b34e6..d2ab2178688 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -778,14 +778,18 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator): ) choice["delta"]["thinking_blocks"] = thinking_blocks translated_choices.append(choice) - return ModelResponseStream( - id=chunk["id"], - object="chat.completion.chunk", - created=chunk["created"], - model=chunk["model"], - choices=translated_choices, - usage=chunk.get("usage"), - ) + kwargs: Final[dict[str, Any]] = { + "id": chunk["id"], + "object": "chat.completion.chunk", + "created": chunk["created"], + "model": chunk["model"], + "choices": translated_choices, + "usage": chunk.get("usage"), + } + service_tier: Final = chunk.get("service_tier") + if isinstance(service_tier, str) and service_tier: + kwargs["service_tier"] = service_tier + return ModelResponseStream(**kwargs) except KeyError as e: raise DatabricksException( message=f"KeyError: {e}, Got unexpected response from Databricks: {chunk}", diff --git a/litellm/llms/databricks/cost_calculator.py b/litellm/llms/databricks/cost_calculator.py index 64166e6fc11..2bb5b99f0ad 100644 --- a/litellm/llms/databricks/cost_calculator.py +++ b/litellm/llms/databricks/cost_calculator.py @@ -30,7 +30,7 @@ def _registry_key(model: str) -> str: ) -def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: +def cost_per_token(model: str, usage: Usage, service_tier: str | None = None) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. @@ -45,4 +45,5 @@ def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: model=_registry_key(model), usage=usage, custom_llm_provider="databricks", + service_tier=service_tier, ) diff --git a/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py b/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py index 9cd17bd3580..1cbc9eeb897 100644 --- a/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py +++ b/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py @@ -883,3 +883,13 @@ def test_completion_merges_system_messages_when_one_has_empty_content(respx_mock {"role": "system", "content": "You are terse."}, {"role": "user", "content": "Hello"}, ] + + +def test_chunk_parser_relays_the_served_service_tier(): + iterator = DatabricksChatResponseIterator(streaming_response=None, sync_stream=True) + + with_tier: Final = iterator.chunk_parser({**_streaming_chunk(), "service_tier": "priority"}) + assert with_tier.model_dump()["service_tier"] == "priority" + + without_tier: Final = iterator.chunk_parser(_streaming_chunk()) + assert getattr(without_tier, "service_tier", None) is None diff --git a/tests/unit/llms/databricks/test_databricks_cost_calculator.py b/tests/unit/llms/databricks/test_databricks_cost_calculator.py index 494b99c1d11..7120a130462 100644 --- a/tests/unit/llms/databricks/test_databricks_cost_calculator.py +++ b/tests/unit/llms/databricks/test_databricks_cost_calculator.py @@ -156,8 +156,6 @@ def test_uncached_request_bills_every_prompt_token_at_the_input_rate(local_model assert completion_cost == pytest.approx(200 * info["output_cost_per_token"]) - - @pytest.mark.parametrize("model", NEW_MODELS) def test_new_models_carry_cache_pricing(local_model_cost_map: None, model: str) -> None: info: Final = _model_info(model) @@ -232,3 +230,28 @@ def test_sonnet_5_ships_standard_rates_not_introductory(local_model_cost_map: No for field in PRICE_FIELDS: assert sonnet_5[field] == pytest.approx(sonnet_4_6[field]), field + + +def test_cost_per_token_bills_the_served_priority_tier( + local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch +) -> None: + rates: Final = { + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + "input_cost_per_token_priority": 0.01, + "output_cost_per_token_priority": 0.02, + "litellm_provider": "databricks", + "mode": "chat", + } + monkeypatch.setitem(litellm.model_cost, "databricks/dbrx-tiered-test", rates) + usage: Final = Usage(prompt_tokens=30, completion_tokens=40, total_tokens=70) + + prompt_cost, completion_cost = cost_per_token( + model="databricks/dbrx-tiered-test", usage=usage, service_tier="priority" + ) + assert prompt_cost == pytest.approx(30 * 0.01) + assert completion_cost == pytest.approx(40 * 0.02) + + prompt_cost, completion_cost = cost_per_token(model="databricks/dbrx-tiered-test", usage=usage) + assert prompt_cost == pytest.approx(30 * 0.001) + assert completion_cost == pytest.approx(40 * 0.002)