fix(databricks): bill cached tokens at cache rates

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-08-21 23:34:09 +00:00
parent 441f3bbe26
commit abb41870c3
2 changed files with 96 additions and 34 deletions

View file

@ -3,10 +3,33 @@ Helper util for handling databricks-specific cost calculation
- e.g.: handling 'dbrx-instruct-*'
"""
from types import MappingProxyType
from typing import Final
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
from litellm.types.utils import Usage
from litellm.utils import get_model_info
_LEGACY_ENDPOINT_NAMES: Final = MappingProxyType(
{
"dbrx-instruct": "databricks-dbrx-instruct",
"meta-llama-3.1-70b-instruct": "databricks-meta-llama-3-1-70b-instruct",
"meta-llama-3.1-405b-instruct": "databricks-meta-llama-3-1-405b-instruct",
"mixtral-8x7b-instruct-v0.1": "databricks-mixtral-8x7b-instruct",
"bge-large-en": "databricks-bge-large-en",
"gte-large-en": "databricks-gte-large-en",
"llama-2-70b-chat": "databricks-llama-2-70b-chat",
}
)
def _base_model(model: str) -> str:
"""The registry key for ``model``, mapping the endpoint names that predate the
``databricks-`` prefixed keys onto their current entries."""
name: Final = model.removeprefix("databricks/")
return next(
(key for prefix, key in _LEGACY_ENDPOINT_NAMES.items() if name.startswith(prefix)),
name,
)
def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
@ -20,36 +43,8 @@ def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
base_model = model
if model.startswith("databricks/dbrx-instruct") or model.startswith("dbrx-instruct"):
base_model = "databricks-dbrx-instruct"
elif model.startswith("databricks/meta-llama-3.1-70b-instruct") or model.startswith("meta-llama-3.1-70b-instruct"):
base_model = "databricks-meta-llama-3-1-70b-instruct"
elif model.startswith("databricks/meta-llama-3.1-405b-instruct") or model.startswith(
"meta-llama-3.1-405b-instruct"
):
base_model = "databricks-meta-llama-3-1-405b-instruct"
elif (
model.startswith("databricks/mixtral-8x7b-instruct-v0.1")
or model.startswith("mixtral-8x7b-instruct-v0.1")
or model.startswith("databricks/mixtral-8x7b-instruct-v0.1")
or model.startswith("mixtral-8x7b-instruct-v0.1")
):
base_model = "databricks-mixtral-8x7b-instruct"
elif model.startswith("databricks/bge-large-en") or model.startswith("bge-large-en"):
base_model = "databricks-bge-large-en"
elif model.startswith("databricks/gte-large-en") or model.startswith("gte-large-en"):
base_model = "databricks-gte-large-en"
elif model.startswith("databricks/llama-2-70b-chat") or model.startswith("llama-2-70b-chat"):
base_model = "databricks-llama-2-70b-chat"
## GET MODEL INFO
model_info: Final = get_model_info(model=base_model, custom_llm_provider="databricks")
## CALCULATE INPUT COST
prompt_cost: Final[float] = usage["prompt_tokens"] * model_info["input_cost_per_token"]
## CALCULATE OUTPUT COST
completion_cost: Final = usage["completion_tokens"] * model_info["output_cost_per_token"]
return prompt_cost, completion_cost
return generic_cost_per_token(
model=_base_model(model),
usage=usage,
custom_llm_provider="databricks",
)

View file

@ -0,0 +1,67 @@
from typing import Final
import pytest
import litellm
from litellm.llms.databricks.cost_calculator import cost_per_token
from litellm.types.utils import Usage
@pytest.fixture
def local_model_cost_map(monkeypatch):
"""Force get_model_info to resolve against the in-repo cost map instead of the
remote one fetched at import time, which still carries the pre-merge pricing."""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
litellm.get_model_info.cache_clear()
yield
litellm.get_model_info.cache_clear()
def _model_info(model: str) -> dict:
return litellm.get_model_info(model=model, custom_llm_provider="databricks")
def test_cached_tokens_are_billed_at_cache_rates(local_model_cost_map):
"""Cache reads and cache writes bill at the model's cache rates, not the input rate"""
model: Final = "databricks/databricks-claude-opus-5"
info: Final = _model_info(model)
usage: Final = Usage(
prompt_tokens=11000,
completion_tokens=500,
total_tokens=11500,
cache_creation_input_tokens=2000,
cache_read_input_tokens=8000,
)
prompt_cost, completion_cost = cost_per_token(model=model, usage=usage)
assert prompt_cost == pytest.approx(
1000 * info["input_cost_per_token"]
+ 2000 * info["cache_creation_input_token_cost"]
+ 8000 * info["cache_read_input_token_cost"]
)
assert completion_cost == pytest.approx(500 * info["output_cost_per_token"])
assert prompt_cost < 11000 * info["input_cost_per_token"]
def test_uncached_request_bills_every_prompt_token_at_the_input_rate(local_model_cost_map):
model: Final = "databricks/databricks-claude-sonnet-5"
info: Final = _model_info(model)
usage: Final = Usage(prompt_tokens=1000, completion_tokens=200, total_tokens=1200)
prompt_cost, completion_cost = cost_per_token(model=model, usage=usage)
assert prompt_cost == pytest.approx(1000 * info["input_cost_per_token"])
assert completion_cost == pytest.approx(200 * info["output_cost_per_token"])
def test_legacy_endpoint_names_still_resolve(local_model_cost_map):
"""Endpoint names that predate the `databricks-` prefixed registry keys keep their pricing"""
info: Final = _model_info("databricks/databricks-mixtral-8x7b-instruct")
usage: Final = Usage(prompt_tokens=100, completion_tokens=100, total_tokens=200)
prompt_cost, completion_cost = cost_per_token(model="databricks/mixtral-8x7b-instruct-v0.1", usage=usage)
assert prompt_cost == pytest.approx(100 * info["input_cost_per_token"])
assert completion_cost == pytest.approx(100 * info["output_cost_per_token"])