mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
441f3bbe26
commit
abb41870c3
2 changed files with 96 additions and 34 deletions
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
Loading…
Add table
Reference in a new issue