fix(cost): bill cache-read tokens at the input rate when the map has no cache-read rate

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-09-18 16:22:14 +00:00
parent 0bd8b7fe02
commit 88d1eb3b35
2 changed files with 40 additions and 22 deletions

View file

@ -519,7 +519,6 @@ def _get_token_base_cost(
current_time: datetime | None = None,
*,
threshold_is_inclusive: bool = False,
missing_cache_read_uses_input: bool = False,
) -> tuple[float, float, float, float, float]:
"""
Return prompt cost, completion cost, and cache costs for a given model and usage.
@ -530,13 +529,11 @@ def _get_token_base_cost(
`threshold_is_inclusive` switches that comparison to >=, for providers such as xAI
that bill the higher tier once the prompt reaches the threshold.
`missing_cache_read_uses_input` resolves an absent cache-read rate to the resolved
input rate instead of 0.0; an explicit 0.0 rate stays a real price either way.
An absent cache-creation rate always resolves to the resolved input rate, the way the
tiered table and custom deployment pricing already do, since a provider that publishes
no write price bills cache writes as ordinary input. An absent 1h write rate resolves
to the cache-creation rate, off-peak included. An explicit 0.0 stays a real price for both.
An absent cache-creation or cache-read rate always resolves to the resolved input
rate, the way the tiered table and custom deployment pricing already do, since a
provider that publishes no cache price bills cached tokens as ordinary input. An
absent 1h write rate resolves to the cache-creation rate, off-peak included. An
explicit 0.0 stays a real price for all of them.
Returns:
Tuple[float, float, float, float] - (prompt_cost, completion_cost, cache_creation_cost, cache_read_cost)
@ -663,8 +660,7 @@ def _get_token_base_cost(
"input_cost_per_token",
prompt_base_cost,
)
if cache_read_cost is None:
cache_read_cost = input_rate_for_missing_cache_rates if missing_cache_read_uses_input else 0.0
resolved_cache_read_cost: Final = input_rate_for_missing_cache_rates if cache_read_cost is None else cache_read_cost
resolved_cache_creation_cost: Final = (
input_rate_for_missing_cache_rates if cache_creation_cost is None else cache_creation_cost
)
@ -677,7 +673,7 @@ def _get_token_base_cost(
completion_base_cost,
resolved_cache_creation_cost,
cache_creation_cost_above_1hr,
cache_read_cost,
resolved_cache_read_cost,
),
)
@ -1588,7 +1584,6 @@ def calculate_prompt_caching_savings(
service_tier=service_tier,
current_time=billed_at,
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
missing_cache_read_uses_input=True,
)
write_rate: Final = cache_creation_cost or prompt_base_cost
write_rate_1h: Final = cache_creation_cost_above_1hr or write_rate

View file

@ -48,7 +48,7 @@ def _local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.mark.parametrize("prompt_tokens", [100, 200000, 200001])
@pytest.mark.parametrize("read_rate", [None, 0.0, 0.25e-6])
@pytest.mark.parametrize("service_tier", [None, "priority"])
def test_missing_cache_read_policy_preserves_billing(prompt_tokens, read_rate, service_tier):
def test_missing_cache_read_rate_resolves_to_input_rate(prompt_tokens, read_rate, service_tier):
info = {
"input_cost_per_token": 3e-6,
"input_cost_per_token_priority": 4e-6,
@ -59,16 +59,43 @@ def test_missing_cache_read_policy_preserves_billing(prompt_tokens, read_rate, s
}
usage = Usage(prompt_tokens=prompt_tokens, prompt_tokens_details={"cached_tokens": 100})
billed = _get_token_base_cost(info, usage, service_tier=service_tier)
savings = _get_token_base_cost(info, usage, service_tier=service_tier, missing_cache_read_uses_input=True)
prompt_cost, _ = generic_cost_per_token(
"policy-fixture", usage, "openai", service_tier=service_tier, model_info=info
)
assert billed[4] == pytest.approx(read_rate or 0.0)
assert savings[:4] == billed[:4]
assert savings[4] == pytest.approx(billed[0] if read_rate is None else read_rate)
assert billed[4] == pytest.approx(read_rate if read_rate is not None else billed[0])
assert prompt_cost == pytest.approx((prompt_tokens - 100) * billed[0] + 100 * billed[4])
def test_generic_cost_per_token_bills_cache_reads_at_input_rate_when_no_cache_read_rate() -> None:
model_info: ModelInfo = {
"key": "bare-model",
"max_tokens": None,
"max_input_tokens": None,
"max_output_tokens": None,
"input_cost_per_token": 2.4e-7,
"output_cost_per_token": 9.7e-7,
"litellm_provider": "bedrock",
"mode": "chat",
"supported_openai_params": None,
}
usage = Usage(
prompt_tokens=12928,
completion_tokens=380,
total_tokens=13308,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=12288),
)
prompt_cost, completion_cost = generic_cost_per_token(
model="bare-model",
usage=usage,
custom_llm_provider="bedrock",
model_info=model_info,
)
assert prompt_cost == pytest.approx(12928 * 2.4e-7)
assert completion_cost == pytest.approx(380 * 9.7e-7)
def test_generic_cost_per_token_prefers_audio_per_second_rate() -> None:
model_info: ModelInfo = {
"key": "gemini-embedding-2",
@ -180,11 +207,7 @@ def test_missing_cache_read_uses_off_peak_input_rate():
}
when = datetime(2026, 9, 7, 12, tzinfo=timezone.utc)
billed = _get_token_base_cost(info, Usage(prompt_tokens=100), current_time=when)
savings = _get_token_base_cost(
info, Usage(prompt_tokens=100), current_time=when, missing_cache_read_uses_input=True
)
assert billed[4] == 0.0
assert savings[0] == savings[4] == 5e-6
assert billed[0] == billed[4] == 5e-6
def test_reasoning_tokens_no_price_set(_local_model_cost_map):