fix: bill cached input tokens at their real rates across vendor dialects

Cached prompt tokens were billed at the full input rate whenever a vendor's
usage dialect or cache pricing was unknown, overcharging DeepSeek topups
~5-10x on agentic workloads (hits are 10x cheaper upstream) and silently
mispricing OpenAI cached reads and Anthropic cache writes the same way.

Two root causes, two fixes:

- Usage dialects: DeepSeek reports prompt_cache_hit_tokens /
  prompt_cache_miss_tokens, which billing never parsed. Usage normalization
  now lives in payment/usage.py as a union parser over the known,
  non-colliding dialects (OpenAI prompt_tokens_details, Anthropic additive
  cache fields, DeepSeek hit/miss), producing one canonical NormalizedUsage.
  Providers expose it as an overridable BaseUpstreamProvider.normalize_usage
  hook — the escape hatch for future vendors whose fields genuinely
  conflict — and every settlement call site passes the provider's result
  through, so calculate_cost holds no vendor knowledge of its own.

- Cache rates: the OpenRouter model feed omits input_cache_read/-write for
  most DeepSeek models (and e.g. openai/gpt-4o), so billing fell back to the
  full input rate. Missing rates are now backfilled from litellm's bundled
  cost map before the provider fee is applied; the input-rate fallback
  remains only as the documented last resort.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Jeroen Ubbink
2026-06-11 21:17:11 +02:00
co-authored by Claude Fable 5
parent a16bc1220c
commit eaf74edbba
9 changed files with 779 additions and 71 deletions
+11 -2
View File
@@ -20,6 +20,7 @@ from .payment.cost_calculation import (
MaxCostData,
calculate_cost,
)
from .payment.usage import NormalizedUsage
from .wallet import credit_balance, deserialize_token_from_string
logger = get_logger(__name__)
@@ -675,12 +676,20 @@ async def revert_pay_for_request(
async def adjust_payment_for_tokens(
key: ApiKey, response_data: dict, session: AsyncSession, deducted_max_cost: int
key: ApiKey,
response_data: dict,
session: AsyncSession,
deducted_max_cost: int,
usage: NormalizedUsage | None = None,
) -> dict:
"""
Adjusts the payment based on token usage in the response.
This is called after the initial payment and the upstream request is complete.
Returns cost data to be included in the response.
``usage`` carries the upstream provider's normalized token usage (its
``normalize_usage`` hook); when omitted, the response's usage object is
normalized with the default union parser.
"""
billing_key = await get_billing_key(key, session)
model = response_data.get("model", "unknown")
@@ -763,7 +772,7 @@ async def adjust_payment_for_tokens(
extra={"error": str(e), "fee_msats": fee_msats},
)
match await calculate_cost(response_data, deducted_max_cost, session):
match await calculate_cost(response_data, deducted_max_cost, session, usage=usage):
case MaxCostData() as cost:
logger.debug(
"Using max cost data (no token adjustment)",
+28 -59
View File
@@ -6,6 +6,15 @@ from ..core import get_logger
from ..core.db import AsyncSession
from ..core.settings import settings
from .price import sats_usd_price
from .usage import NormalizedUsage, normalize_usage, parse_token_count
__all__ = [
"CostData",
"CostDataError",
"MaxCostData",
"calculate_cost",
"parse_token_count",
]
logger = get_logger(__name__)
@@ -34,13 +43,20 @@ class CostDataError(BaseModel):
async def calculate_cost(
response_data: dict, max_cost: int, session: AsyncSession
response_data: dict,
max_cost: int,
session: AsyncSession,
usage: NormalizedUsage | None = None,
) -> CostData | MaxCostData | CostDataError:
"""Calculate the cost of an API request based on token usage.
Args:
response_data: Response data containing usage information
max_cost: Maximum cost in millisats
usage: Pre-normalized usage from the upstream provider's
``normalize_usage`` hook. When omitted, the response's usage
object is normalized with the default union parser. This function
holds no vendor-dialect knowledge of its own.
Returns:
Cost data or error information
@@ -54,8 +70,10 @@ async def calculate_cost(
},
)
# Check for usage data
if "usage" not in response_data or response_data["usage"] is None:
if usage is None:
usage = normalize_usage(response_data.get("usage"))
if usage is None:
logger.warning(
"No usage data in response — billing at MaxCostData with zero "
"tokens. Dashboard will show this request as `(0+0)`. Most "
@@ -84,16 +102,14 @@ async def calculate_cost(
cache_creation_msats=0,
)
usage_data = response_data["usage"]
usage_data = response_data.get("usage") or {}
if not isinstance(usage_data, dict):
usage_data = {}
# Extract token counts
input_tokens = _extract_token_pair(usage_data, "prompt_tokens", "input_tokens")
output_tokens = _extract_token_pair(usage_data, "completion_tokens", "output_tokens")
# Extract cache tokens (handles OpenAI vs Anthropic formats)
cache_read_tokens, cache_creation_tokens, input_tokens = _extract_cache_tokens(
usage_data, input_tokens
)
input_tokens = usage.input_tokens
output_tokens = usage.output_tokens
cache_read_tokens = usage.cache_read_tokens
cache_creation_tokens = usage.cache_write_tokens
# Try USD cost first
usd_cost = _resolve_usd_cost(usage_data, response_data)
@@ -202,22 +218,6 @@ async def calculate_cost(
# ============================================================================
def parse_token_count(value: object) -> int:
"""Parse a token count from various formats (int, float, str, bool)."""
if isinstance(value, bool):
return 0
if isinstance(value, int):
return max(0, value)
if isinstance(value, float):
return max(0, int(value))
if isinstance(value, str):
try:
return max(0, int(float(value)))
except ValueError:
return 0
return 0
def _coerce_usd(value: object) -> float:
"""Coerce a value to USD float, handling various formats safely."""
if value is None or isinstance(value, bool):
@@ -230,37 +230,6 @@ def _coerce_usd(value: object) -> float:
return 0.0
def _extract_token_pair(
usage_data: dict, standard_field: str, alt_field: str
) -> int:
"""Extract token count trying two field names in order."""
value = parse_token_count(usage_data.get(standard_field, 0))
if value > 0:
return value
return parse_token_count(usage_data.get(alt_field, 0))
def _extract_cache_tokens(usage_data: dict, input_tokens: int) -> tuple[int, int, int]:
"""Extract cache tokens, handling OpenAI vs Anthropic formats.
Returns: (cache_read_tokens, cache_creation_tokens, adjusted_input_tokens)
"""
cache_read = parse_token_count(usage_data.get("cache_read_input_tokens", 0))
cache_creation = parse_token_count(
usage_data.get("cache_creation_input_tokens", 0)
)
# OpenAI: cache is included in input_tokens, subtract it
prompt_details = usage_data.get("prompt_tokens_details")
if isinstance(prompt_details, dict) and not cache_read:
openai_cached = parse_token_count(prompt_details.get("cached_tokens", 0))
if openai_cached:
cache_read = openai_cached
input_tokens = max(0, input_tokens - cache_read)
return cache_read, cache_creation, input_tokens
def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
"""Resolve USD cost with clear priority order.
+42
View File
@@ -85,6 +85,48 @@ class Model(BaseModel):
return hash(self.id)
def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing:
"""Fill missing cache rates from litellm's bundled cost map.
The OpenRouter model feed omits ``input_cache_read``/``input_cache_write``
for many models (most DeepSeek entries, openai/gpt-4o, ...). Without a
cache rate, billing falls back to the full input rate, which overcharges
cache reads (DeepSeek hits are 10x cheaper) and undercharges Anthropic
cache writes (1.25x). litellm ships per-model USD rates keyed by the exact
OpenRouter id (deepseek/deepseek-chat) or by the bare model name
(gpt-4o, claude-sonnet-4-5), so both spellings are tried.
Rates already present (e.g. provided by OpenRouter) are authoritative and
never overwritten. Unknown models are returned unchanged.
"""
needs_read = (pricing.input_cache_read or 0.0) <= 0.0
needs_write = (pricing.input_cache_write or 0.0) <= 0.0
if not (needs_read or needs_write):
return pricing
import litellm
info: dict | None = None
for key in (model_id, model_id.split("/", 1)[-1]):
candidate = litellm.model_cost.get(key)
if isinstance(candidate, dict):
info = candidate
break
if info is None:
return pricing
updated = Pricing.parse_obj(pricing.dict())
if needs_read:
read_rate = info.get("cache_read_input_token_cost")
if isinstance(read_rate, (int, float)) and read_rate > 0:
updated.input_cache_read = float(read_rate)
if needs_write:
write_rate = info.get("cache_creation_input_token_cost")
if isinstance(write_rate, (int, float)) and write_rate > 0:
updated.input_cache_write = float(write_rate)
return updated
def _has_valid_pricing(model: dict) -> bool:
"""Check if model has valid pricing (not free, no negative values)."""
pricing = model.get("pricing", {})
+92
View File
@@ -0,0 +1,92 @@
"""Vendor-agnostic normalization of upstream usage objects.
Upstream providers report token usage in vendor dialects that differ in field
names and in whether cached tokens are included in the input count:
* OpenAI: ``prompt_tokens_details.cached_tokens``, included in ``prompt_tokens``
* Anthropic: ``cache_read_input_tokens`` / ``cache_creation_input_tokens``,
additive to (not included in) ``input_tokens``
* DeepSeek: ``prompt_cache_hit_tokens`` / ``prompt_cache_miss_tokens``, with
``prompt_tokens = hit + miss``
``normalize_usage`` maps all of them onto one canonical ``NormalizedUsage``
shape so billing code needs no vendor knowledge. The known dialects' field
names do not collide, so a single union parser is safe; vendors whose fields
would genuinely conflict override ``BaseUpstreamProvider.normalize_usage``.
"""
from pydantic.v1 import BaseModel
class NormalizedUsage(BaseModel):
"""Canonical token usage: input_tokens never includes cached tokens."""
input_tokens: int = 0
output_tokens: int = 0
cache_read_tokens: int = 0
cache_write_tokens: int = 0
def parse_token_count(value: object) -> int:
"""Parse a token count from various formats (int, float, str, bool)."""
if isinstance(value, bool):
return 0
if isinstance(value, int):
return max(0, value)
if isinstance(value, float):
return max(0, int(value))
if isinstance(value, str):
try:
return max(0, int(float(value)))
except ValueError:
return 0
return 0
def _first_token_count(usage_data: dict, *fields: str) -> int:
"""Return the first positive token count among the given fields."""
for field in fields:
value = parse_token_count(usage_data.get(field, 0))
if value > 0:
return value
return 0
def normalize_usage(usage_data: object) -> NormalizedUsage | None:
"""Map a vendor usage dict onto the canonical shape, or None if absent.
Cached tokens are subtracted from the input count exactly once, only for
dialects that include them in it (OpenAI, DeepSeek). Precedence between
cache fields: Anthropic explicit > OpenAI details > DeepSeek hit/miss.
"""
if not isinstance(usage_data, dict):
return None
input_tokens = _first_token_count(usage_data, "prompt_tokens", "input_tokens")
output_tokens = _first_token_count(
usage_data, "completion_tokens", "output_tokens"
)
cache_write = parse_token_count(usage_data.get("cache_creation_input_tokens", 0))
# Anthropic: cache reads are additive, input_tokens stays untouched
cache_read = parse_token_count(usage_data.get("cache_read_input_tokens", 0))
if not cache_read:
# OpenAI: cached tokens are included in prompt_tokens
prompt_details = usage_data.get("prompt_tokens_details")
if isinstance(prompt_details, dict):
cache_read = parse_token_count(prompt_details.get("cached_tokens", 0))
# DeepSeek: prompt_tokens = prompt_cache_hit_tokens + prompt_cache_miss_tokens
if not cache_read:
cache_read = parse_token_count(
usage_data.get("prompt_cache_hit_tokens", 0)
)
if cache_read:
input_tokens = max(0, input_tokens - cache_read)
return NormalizedUsage(
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_read_tokens=cache_read,
cache_write_tokens=cache_write,
)
+63 -7
View File
@@ -35,9 +35,11 @@ from ..payment.models import (
Pricing,
_calculate_usd_max_costs,
_update_model_sats_pricing,
backfill_cache_pricing,
list_models,
)
from ..payment.price import sats_usd_price
from ..payment.usage import NormalizedUsage, normalize_usage
from ..wallet import recieve_token, send_token
from . import messages_dispatch
from .count_tokens import count_tokens_locally
@@ -101,6 +103,17 @@ class BaseUpstreamProvider:
self._models_cache = []
self._models_by_id = {}
def normalize_usage(self, usage_data: object) -> NormalizedUsage | None:
"""Map this provider's usage dialect onto the canonical shape.
The default union parser covers the known dialects (OpenAI, Anthropic,
DeepSeek), whose field names do not collide. Override in a subclass
when a vendor's usage fields would conflict with another dialect or
need bespoke interpretation; billing code consumes only the canonical
result and holds no vendor knowledge.
"""
return normalize_usage(usage_data)
def get_litellm_provider_prefix(self) -> str:
"""Resolve the litellm provider prefix for this provider instance.
@@ -880,6 +893,9 @@ class BaseUpstreamProvider:
adjustment_input,
session,
max_cost_for_model,
usage=self.normalize_usage(
adjustment_input.get("usage")
),
)
usage_finalized = True
except Exception as e:
@@ -1018,7 +1034,11 @@ class BaseUpstreamProvider:
response_json["id"] = f"chatcmpl-{uuid.uuid4()}"
cost_data = await adjust_payment_for_tokens(
key, response_json, session, deducted_max_cost
key,
response_json,
session,
deducted_max_cost,
usage=self.normalize_usage(response_json.get("usage")),
)
await session.refresh(key)
@@ -1273,6 +1293,9 @@ class BaseUpstreamProvider:
adjustment_input,
session,
max_cost_for_model,
usage=self.normalize_usage(
adjustment_input.get("usage")
),
)
usage_finalized = True
except Exception as e:
@@ -1436,7 +1459,11 @@ class BaseUpstreamProvider:
response_json["id"] = f"chatcmpl-{uuid.uuid4()}"
cost_data = await adjust_payment_for_tokens(
key, response_json, session, deducted_max_cost
key,
response_json,
session,
deducted_max_cost,
usage=self.normalize_usage(response_json.get("usage")),
)
await session.refresh(key)
@@ -1631,7 +1658,11 @@ class BaseUpstreamProvider:
"usage": None,
}
cost_data = await adjust_payment_for_tokens(
fresh_key, fallback, new_session, max_cost_for_model
fresh_key,
fallback,
new_session,
max_cost_for_model,
usage=self.normalize_usage(fallback.get("usage")),
)
usage_finalized = True
return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
@@ -1783,6 +1814,9 @@ class BaseUpstreamProvider:
combined_data,
new_session,
max_cost_for_model,
usage=self.normalize_usage(
combined_data.get("usage")
),
)
self.inject_cost_metadata(
@@ -1850,7 +1884,11 @@ class BaseUpstreamProvider:
response_json["usage"] = {"input_tokens": input_tokens}
cost_data = await adjust_payment_for_tokens(
key, response_json, session, deducted_max_cost
key,
response_json,
session,
deducted_max_cost,
usage=self.normalize_usage(response_json.get("usage")),
)
self.inject_cost_metadata(response_json, cost_data, key)
@@ -1952,7 +1990,11 @@ class BaseUpstreamProvider:
response_json["model"] = requested_model
cost_data = await adjust_payment_for_tokens(
key, response_json, session, max_cost_for_model
key,
response_json,
session,
max_cost_for_model,
usage=self.normalize_usage(response_json.get("usage")),
)
self.inject_cost_metadata(response_json, cost_data, key)
@@ -2094,6 +2136,7 @@ class BaseUpstreamProvider:
fallback,
new_session,
max_cost_for_model,
usage=self.normalize_usage(fallback.get("usage")),
)
usage_finalized = True
return (
@@ -2167,6 +2210,9 @@ class BaseUpstreamProvider:
combined_data,
new_session,
max_cost_for_model,
usage=self.normalize_usage(
combined_data.get("usage")
),
)
self.inject_cost_metadata(
combined_data, cost_data, fresh_key
@@ -3026,7 +3072,12 @@ class BaseUpstreamProvider:
)
async with create_session() as session:
match await calculate_cost(response_data, max_cost_for_model, session):
match await calculate_cost(
response_data,
max_cost_for_model,
session,
usage=self.normalize_usage(response_data.get("usage")),
):
case MaxCostData() as cost:
logger.debug(
"Using max cost pricing",
@@ -4562,14 +4613,19 @@ class BaseUpstreamProvider:
def _apply_provider_fee_to_model(self, model: Model) -> Model:
"""Apply provider fee to model's USD pricing and calculate max costs.
Cache rates missing from the upstream pricing feed are backfilled from
litellm's cost map first, so they carry the provider fee like every
other price component.
Args:
model: Model object to update
Returns:
Model with provider fee applied to pricing and max costs calculated
"""
base_pricing = backfill_cache_pricing(model.id, model.pricing)
adjusted_pricing = Pricing.parse_obj(
{k: v * self.provider_fee for k, v in model.pricing.dict().items()}
{k: v * self.provider_fee for k, v in base_pricing.dict().items()}
)
temp_model = Model(
+241
View File
@@ -0,0 +1,241 @@
"""Tests for cache-aware pricing of cached input tokens.
Specifies two things:
1. ``backfill_cache_pricing`` — when the OpenRouter model feed omits cache
rates (it does for most DeepSeek models and e.g. openai/gpt-4o), they are
filled from litellm's bundled cost map instead of silently billing cache
reads at the full input rate. Existing OpenRouter values are never
overwritten, and provider fees apply to backfilled rates like any other.
2. ``calculate_cost`` — cached tokens are billed at the cache rates from the
model's sats_pricing; the full input rate remains only as the documented
last resort when no cache rate could be resolved anywhere.
"""
import os
from unittest.mock import AsyncMock, Mock, patch
import litellm
import pytest
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to")
from routstr.core.settings import settings
from routstr.payment.cost_calculation import CostData, calculate_cost
from routstr.payment.models import (
Architecture,
Model,
Pricing,
backfill_cache_pricing,
)
from routstr.upstream import GenericUpstreamProvider
def _make_model(model_id: str, pricing: Pricing) -> Model:
return Model(
id=model_id,
name=model_id,
created=0,
description="",
context_length=64000,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="Other",
instruct_type=None,
),
pricing=pricing,
)
# ============================================================================
# backfill_cache_pricing — litellm as fallback source for missing cache rates
# ============================================================================
def test_backfill_deepseek_cache_read_from_litellm() -> None:
"""deepseek/deepseek-chat has no input_cache_read on OpenRouter; litellm
knows the real rate (10x cheaper than input)."""
pricing = Pricing(prompt=2.8e-07, completion=4.2e-07)
result = backfill_cache_pricing("deepseek/deepseek-chat", pricing)
expected = litellm.model_cost["deepseek/deepseek-chat"][
"cache_read_input_token_cost"
]
assert result.input_cache_read == expected
assert result.input_cache_read < pricing.prompt # sanity: it's a discount
def test_backfill_strips_vendor_prefix_for_litellm_lookup() -> None:
"""OpenRouter ids are vendor-prefixed (openai/gpt-4o); litellm keys most
non-DeepSeek models without the prefix (gpt-4o)."""
pricing = Pricing(prompt=2.5e-06, completion=1e-05)
result = backfill_cache_pricing("openai/gpt-4o", pricing)
expected = litellm.model_cost["gpt-4o"]["cache_read_input_token_cost"]
assert result.input_cache_read == expected
def test_backfill_fills_cache_write_rate() -> None:
"""Anthropic cache writes cost more than input (1.25x); billing them at
the input rate undercharges. litellm carries the write rate."""
pricing = Pricing(prompt=3e-06, completion=1.5e-05)
result = backfill_cache_pricing("anthropic/claude-sonnet-4-5", pricing)
expected = litellm.model_cost["claude-sonnet-4-5"][
"cache_creation_input_token_cost"
]
assert result.input_cache_write == expected
assert result.input_cache_write > pricing.prompt # sanity: write premium
def test_backfill_never_overwrites_openrouter_rates() -> None:
"""When OpenRouter provides a cache rate, it is authoritative."""
pricing = Pricing(
prompt=2.1e-07, completion=7.9e-07, input_cache_read=1.3e-07
)
result = backfill_cache_pricing("deepseek/deepseek-chat", pricing)
assert result.input_cache_read == 1.3e-07
def test_backfill_unknown_model_unchanged() -> None:
"""Models litellm doesn't know stay untouched (last-resort fallback to
the input rate happens later, at billing time)."""
pricing = Pricing(prompt=1e-06, completion=2e-06)
result = backfill_cache_pricing("artificial-dumbness/dumb-1", pricing)
assert result.input_cache_read == 0.0
assert result.input_cache_write == 0.0
def test_provider_fee_applies_to_backfilled_cache_rates() -> None:
"""Backfill happens before the provider fee, so cache rates carry the
same markup as every other price component."""
provider = GenericUpstreamProvider(
base_url="http://upstream.example", provider_fee=2.0
)
model = _make_model(
"deepseek/deepseek-chat", Pricing(prompt=2.8e-07, completion=4.2e-07)
)
adjusted = provider._apply_provider_fee_to_model(model)
litellm_read = litellm.model_cost["deepseek/deepseek-chat"][
"cache_read_input_token_cost"
]
assert adjusted.pricing.input_cache_read == pytest.approx(litellm_read * 2.0)
assert adjusted.pricing.prompt == pytest.approx(2.8e-07 * 2.0)
# ============================================================================
# calculate_cost — cached tokens billed at cache rates
# ============================================================================
@pytest.fixture(autouse=True)
def patch_sats_usd_price() -> None: # type: ignore[misc]
with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-5):
yield
@pytest.fixture
def model_pricing(monkeypatch: pytest.MonkeyPatch) -> Mock:
"""Model-based pricing: 1 msat per input token, 2 per output token,
0.1 per cache-read token, 1.25 per cache-write token."""
monkeypatch.setattr(settings, "fixed_pricing", False)
model = Mock()
model.sats_pricing = Pricing(
prompt=0.001,
completion=0.002,
input_cache_read=0.0001,
input_cache_write=0.00125,
)
return model
@pytest.mark.asyncio
async def test_deepseek_cache_hits_billed_at_cache_rate(model_pricing: Mock) -> None:
"""The reported overcharge scenario: a 10k-token prompt with 90% cache
hits costs 2900 msats at honest rates, not the 11000 msats that billing
every prompt token at the full input rate would charge."""
response = {
"model": "deepseek-chat",
"usage": {
"prompt_tokens": 10000,
"completion_tokens": 500,
"prompt_cache_hit_tokens": 9000,
"prompt_cache_miss_tokens": 1000,
},
}
with patch("routstr.proxy.get_model_instance", return_value=model_pricing):
result = await calculate_cost(response, max_cost=100000, session=AsyncMock())
assert isinstance(result, CostData)
# 1000 input @ 1 msat + 9000 cache reads @ 0.1 msat + 500 output @ 2 msat
assert result.input_msats == 1000
assert result.cache_read_msats == 900
assert result.output_msats == 1000
assert result.total_msats == 2900
@pytest.mark.asyncio
async def test_anthropic_cache_write_billed_at_write_rate(model_pricing: Mock) -> None:
"""Cache writes carry their premium rate (1.25x input here), instead of
being silently billed at the plain input rate."""
response = {
"model": "claude-sonnet-4-5",
"usage": {
"input_tokens": 300,
"output_tokens": 100,
"cache_read_input_tokens": 500,
"cache_creation_input_tokens": 2000,
},
}
with patch("routstr.proxy.get_model_instance", return_value=model_pricing):
result = await calculate_cost(response, max_cost=100000, session=AsyncMock())
assert isinstance(result, CostData)
# 300 @ 1 + 500 @ 0.1 + 2000 @ 1.25 + 100 @ 2
assert result.cache_read_msats == 50
assert result.cache_creation_msats == 2500
assert result.total_msats == 3050
@pytest.mark.asyncio
async def test_missing_cache_rate_falls_back_to_input_rate(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Documented last resort: when no cache rate could be resolved anywhere
(OpenRouter and litellm both silent), cache reads bill at the input rate —
never cheaper, never free."""
monkeypatch.setattr(settings, "fixed_pricing", False)
model = Mock()
model.sats_pricing = Pricing(prompt=0.001, completion=0.002)
response = {
"model": "dumb-1",
"usage": {
"prompt_tokens": 10000,
"completion_tokens": 500,
"prompt_cache_hit_tokens": 9000,
"prompt_cache_miss_tokens": 1000,
},
}
with patch("routstr.proxy.get_model_instance", return_value=model):
result = await calculate_cost(response, max_cost=100000, session=AsyncMock())
assert isinstance(result, CostData)
# 1000 @ 1 + 9000 @ 1 (fallback) + 500 @ 2
assert result.total_msats == 11000
+94 -1
View File
@@ -1,6 +1,7 @@
"""Tests for cache token handling in cost calculation.
Covers OpenAI vs Anthropic caching formats, edge cases, and billing accuracy.
Covers OpenAI, Anthropic and DeepSeek caching formats, dialect precedence,
edge cases, and billing accuracy.
"""
import os
@@ -316,6 +317,98 @@ async def test_zero_cache_tokens(mock_session: AsyncMock, mock_fixed_pricing: No
assert result.input_tokens == 100
# ============================================================================
# DeepSeek Cache Format
# DeepSeek emits neither OpenAI's prompt_tokens_details nor Anthropic's
# cache_read_input_tokens — only prompt_cache_hit_tokens and
# prompt_cache_miss_tokens, with the documented guarantee
# prompt_tokens = hit + miss. Hits are ~10x cheaper upstream, so billing
# them as regular input is a large overcharge.
# ============================================================================
@pytest.mark.asyncio
async def test_deepseek_cache_hit_tokens_extracted(mock_session: AsyncMock) -> None:
"""DeepSeek cache hits are extracted and removed from regular input.
Payload shape verbatim from the DeepSeek API reference (usage object).
"""
response = {
"model": "deepseek-chat",
"usage": {
"prompt_tokens": 10000, # = hit + miss
"completion_tokens": 500,
"total_tokens": 10500,
"prompt_cache_hit_tokens": 9000,
"prompt_cache_miss_tokens": 1000,
},
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 1000 # only the cache misses
assert result.cache_read_input_tokens == 9000
assert result.output_tokens == 500
@pytest.mark.asyncio
async def test_deepseek_all_tokens_cached(mock_session: AsyncMock) -> None:
"""A fully cached DeepSeek prompt bills zero regular input tokens."""
response = {
"model": "deepseek-chat",
"usage": {
"prompt_tokens": 5000,
"completion_tokens": 100,
"prompt_cache_hit_tokens": 5000,
"prompt_cache_miss_tokens": 0,
},
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 0
assert result.cache_read_input_tokens == 5000
@pytest.mark.asyncio
async def test_dialect_precedence_never_double_subtracts(mock_session: AsyncMock) -> None:
"""If a vendor emits both OpenAI-style and DeepSeek-style cache fields for
the same cached tokens, they are counted once, not subtracted twice."""
response = {
"model": "deepseek-chat",
"usage": {
"prompt_tokens": 10000,
"completion_tokens": 500,
"prompt_tokens_details": {"cached_tokens": 9000},
"prompt_cache_hit_tokens": 9000,
"prompt_cache_miss_tokens": 1000,
},
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 1000 # 10000 - 9000, applied exactly once
assert result.cache_read_input_tokens == 9000
@pytest.mark.asyncio
async def test_deepseek_malformed_hit_tokens_coerce_to_zero(mock_session: AsyncMock) -> None:
"""Malformed DeepSeek cache fields degrade to billing all input at full
rate instead of crashing or going negative."""
response = {
"model": "deepseek-chat",
"usage": {
"prompt_tokens": 1000,
"completion_tokens": 50,
"prompt_cache_hit_tokens": "garbage",
"prompt_cache_miss_tokens": -5,
},
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 1000
assert result.cache_read_input_tokens == 0
# ============================================================================
# Test 13: Missing Usage Block
# ============================================================================
+2 -2
View File
@@ -500,7 +500,7 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
captured_cost_call: dict[str, Any] = {}
async def fake_adjust(
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None
) -> dict:
captured_cost_call["combined_data"] = combined_data
captured_cost_call["max_cost"] = max_cost
@@ -589,7 +589,7 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None:
captured: dict[str, Any] = {}
async def fake_adjust(
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None
) -> dict:
captured["combined_data"] = combined_data
return fake_cost
+206
View File
@@ -0,0 +1,206 @@
"""Tests for vendor-agnostic usage normalization.
Specifies the seam that keeps vendor usage dialects out of generic billing
code: a canonical ``NormalizedUsage`` shape produced by
``routstr.payment.usage.normalize_usage`` (union parser for the known,
non-colliding dialects) and exposed as an overridable
``BaseUpstreamProvider.normalize_usage`` hook for vendors whose fields would
genuinely conflict. ``calculate_cost`` accepts a pre-normalized usage and then
needs no vendor knowledge of its own.
"""
import os
from unittest.mock import AsyncMock, patch
import pytest
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to")
from routstr.core.settings import settings
from routstr.payment.cost_calculation import CostData, calculate_cost
from routstr.payment.usage import NormalizedUsage, normalize_usage
from routstr.upstream import BaseUpstreamProvider, GenericUpstreamProvider
@pytest.fixture(autouse=True)
def mock_fixed_pricing(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(settings, "fixed_pricing", True)
monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 0.001)
monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 0.001)
@pytest.fixture(autouse=True)
def patch_sats_usd_price() -> None: # type: ignore[misc]
with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-5):
yield
# ============================================================================
# The union parser: one canonical shape for all known dialects
# ============================================================================
@pytest.mark.parametrize(
"usage,expected",
[
# OpenAI: cached_tokens included in prompt_tokens → subtracted
(
{
"prompt_tokens": 2000,
"completion_tokens": 100,
"prompt_tokens_details": {"cached_tokens": 800},
},
NormalizedUsage(
input_tokens=1200,
output_tokens=100,
cache_read_tokens=800,
cache_write_tokens=0,
),
),
# DeepSeek: hit/miss fields, prompt_tokens = hit + miss → hit subtracted
(
{
"prompt_tokens": 10000,
"completion_tokens": 500,
"prompt_cache_hit_tokens": 9000,
"prompt_cache_miss_tokens": 1000,
},
NormalizedUsage(
input_tokens=1000,
output_tokens=500,
cache_read_tokens=9000,
cache_write_tokens=0,
),
),
# Anthropic: cache fields additive, input_tokens NOT reduced
(
{
"input_tokens": 300,
"output_tokens": 100,
"cache_read_input_tokens": 500,
"cache_creation_input_tokens": 2000,
},
NormalizedUsage(
input_tokens=300,
output_tokens=100,
cache_read_tokens=500,
cache_write_tokens=2000,
),
),
# Plain OpenAI without caching
(
{"prompt_tokens": 100, "completion_tokens": 50},
NormalizedUsage(input_tokens=100, output_tokens=50),
),
],
)
def test_normalize_usage_dialects(usage: dict, expected: NormalizedUsage) -> None:
"""Each known vendor dialect maps onto the same canonical shape."""
assert normalize_usage(usage) == expected
def test_normalize_usage_absent_usage() -> None:
"""Missing/invalid usage yields None so callers can bill at max cost."""
assert normalize_usage(None) is None
assert normalize_usage("not a dict") is None # type: ignore[arg-type]
def test_normalize_usage_never_negative() -> None:
"""Buggy upstreams reporting more cached than prompt tokens clamp to 0."""
result = normalize_usage(
{
"prompt_tokens": 100,
"completion_tokens": 50,
"prompt_cache_hit_tokens": 150,
}
)
assert result is not None
assert result.input_tokens == 0
assert result.cache_read_tokens == 150
# ============================================================================
# The provider hook: default delegates to the union parser, subclasses
# override for vendors whose fields genuinely conflict
# ============================================================================
class ArtificialDumbnessProvider(BaseUpstreamProvider):
"""Fictional vendor whose usage dialect collides with nothing we know."""
provider_type = "artificial-dumbness"
def normalize_usage(self, usage_data: object) -> NormalizedUsage | None:
if not isinstance(usage_data, dict):
return None
return NormalizedUsage(
input_tokens=usage_data.get("dumb_tokens_in", 0),
output_tokens=usage_data.get("dumb_tokens_out", 0),
cache_read_tokens=usage_data.get("dumb_tokens_reused", 0),
cache_write_tokens=0,
)
def test_base_provider_hook_delegates_to_union_parser() -> None:
"""Without an override, the provider hook equals the module parser, so
DeepSeek-dialect upstreams work through GenericUpstreamProvider with no
subclass."""
provider = GenericUpstreamProvider(base_url="http://upstream.example")
usage = {
"prompt_tokens": 10000,
"completion_tokens": 500,
"prompt_cache_hit_tokens": 9000,
"prompt_cache_miss_tokens": 1000,
}
assert provider.normalize_usage(usage) == normalize_usage(usage)
@pytest.mark.asyncio
async def test_provider_override_is_honored_by_calculate_cost() -> None:
"""The escape hatch: a vendor-specific override feeds calculate_cost
through the `usage` parameter, and generic billing needs no knowledge of
the vendor's field names."""
provider = ArtificialDumbnessProvider(base_url="http://ad.example", api_key="k")
response = {
"model": "dumb-1",
"usage": {
"dumb_tokens_in": 700,
"dumb_tokens_out": 60,
"dumb_tokens_reused": 4300,
},
}
result = await calculate_cost(
response,
max_cost=100000,
session=AsyncMock(),
usage=provider.normalize_usage(response["usage"]),
)
assert isinstance(result, CostData)
assert result.input_tokens == 700
assert result.output_tokens == 60
assert result.cache_read_input_tokens == 4300
@pytest.mark.asyncio
async def test_explicit_usage_param_wins_over_response_extraction() -> None:
"""When a normalized usage is passed, calculate_cost must not re-derive
token counts from the raw response."""
response = {
"model": "dumb-1",
"usage": {"prompt_tokens": 999999, "completion_tokens": 999999},
}
result = await calculate_cost(
response,
max_cost=100000,
session=AsyncMock(),
usage=NormalizedUsage(input_tokens=10, output_tokens=5),
)
assert isinstance(result, CostData)
assert result.input_tokens == 10
assert result.output_tokens == 5