mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
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>
This commit is contained in:
parent
0d8fad7660
commit
fc268b19aa
5 changed files with 50 additions and 12 deletions
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue