fix(databricks): keep the served service_tier on streamed chunks and bill it
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-09-27 00:51:26 +00:00
parent 0d8fad7660
commit fc268b19aa
5 changed files with 50 additions and 12 deletions

View file

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

View file

@ -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}",

View file

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

View file

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

View file

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