Compare commits

...
Author SHA1 Message Date
9qeklajc 6dac01c851 more robust cost calculation 2026-06-19 00:04:14 +02:00
9qeklajc 1dceffffa7 resolve review comments 2026-06-13 23:35:14 +02:00
9qeklajc 4b74a9cf81 explicit cache 2026-06-13 22:35:55 +02:00
9qeklajc 3be2da6728 update to all providers 2026-06-13 21:06:56 +02:00
9qeklajc 17d690e77e Merge branch 'fix/cached-token-overcharge' into fix-reservation 2026-06-13 21:06:42 +02:00
Jeroen UbbinkandClaude Fable 5 cbc424e8e7 build: type-check the entire repo in make targets, matching CI
CI runs 'uv run mypy .' while the Makefile only checked routstr/, so test
files could pass locally and fail the pipeline. lint, type-check and
ci-lint now check everything; --ignore-missing-imports is dropped since
the CI invocation passes without it.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-06-11 21:22:28 +02:00
Jeroen UbbinkandClaude Fable 5 068fb3572f refactor: drop unused session parameter from calculate_cost
The session was needed when model pricing lived in the DB (73d3613) and has
been dead since pricing moved to the in-memory model map (0da08fb), yet every
caller was still obliged to supply one. get_x_cashu_cost even opened a DB
session per x-cashu request solely to feed it.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-06-11 21:17:11 +02:00
Jeroen UbbinkandClaude Fable 5 eaf74edbba 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>
2026-06-11 21:17:11 +02:00
14 changed files with 1631 additions and 222 deletions
+3 -3
View File
@@ -98,7 +98,7 @@ docker-down:
lint:
@echo "🔍 Running linting checks..."
$(RUFF) check .
$(MYPY) routstr/ --ignore-missing-imports
$(MYPY) .
format:
@echo "✨ Formatting code..."
@@ -107,7 +107,7 @@ format:
type-check:
@echo "🔎 Running type checks..."
$(MYPY) routstr/ --ignore-missing-imports
$(MYPY) .
# Development setup
dev-setup:
@@ -234,7 +234,7 @@ ci-test:
ci-lint:
@echo "🤖 Running CI linting..."
$(RUFF) check . --exit-non-zero-on-fix
$(MYPY) routstr/ --ignore-missing-imports --no-error-summary
$(MYPY) . --no-error-summary
# Debug helpers
test-debug:
+8 -2
View File
@@ -708,12 +708,18 @@ 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,
) -> 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.
The response's usage object is normalized with the default union parser in
``calculate_cost``.
"""
billing_key = await get_billing_key(key, session)
model = response_data.get("model", "unknown")
@@ -796,7 +802,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):
case MaxCostData() as cost:
logger.debug(
"Using max cost data (no token adjustment)",
+71 -68
View File
@@ -3,9 +3,17 @@ import math
from pydantic.v1 import BaseModel
from ..core import get_logger
from ..core.db import AsyncSession
from ..core.settings import settings
from .price import sats_usd_price
from .usage import normalize_usage, parse_token_count
__all__ = [
"CostData",
"CostDataError",
"MaxCostData",
"calculate_cost",
"parse_token_count",
]
logger = get_logger(__name__)
@@ -34,7 +42,8 @@ class CostDataError(BaseModel):
async def calculate_cost(
response_data: dict, max_cost: int, session: AsyncSession
response_data: dict,
max_cost: int,
) -> CostData | MaxCostData | CostDataError:
"""Calculate the cost of an API request based on token usage.
@@ -44,6 +53,9 @@ async def calculate_cost(
Returns:
Cost data or error information
The response's usage object is normalized with the default union parser;
this function holds no vendor-dialect knowledge of its own.
"""
logger.debug(
"Starting cost calculation",
@@ -54,8 +66,9 @@ async def calculate_cost(
},
)
# Check for usage data
if "usage" not in response_data or response_data["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 +97,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 +213,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 +225,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.
@@ -352,6 +316,44 @@ def _resolve_provider_fee(model_id: str) -> float:
return float(providers[0].provider_fee)
def _rate_weighted_input_fraction(
response_data: dict,
input_tokens: int,
cache_read_tokens: int,
cache_creation_tokens: int,
output_tokens: int,
) -> float:
"""Fraction of a lump-sum cost attributable to the input side.
The upstream gave a single total USD figure with no input/output breakdown,
so the split has to be reconstructed. A raw token-count split (input + cache
vs output) overstates the input share whenever cheap cache-read tokens
dominate — a cache read costs ~10x less than a regular input token yet counts
the same. Weight each bucket by its per-token rate instead. Fall back to a
token-count split when rates are unavailable (fixed pricing / unknown model).
"""
input_weight = float(input_tokens + cache_read_tokens + cache_creation_tokens)
output_weight = float(output_tokens)
try:
rates = _get_pricing_rates(response_data)
except ValueError:
rates = None
if rates is not None:
input_rate, output_rate, cache_read_rate, cache_creation_rate = rates
input_weight = (
input_tokens * input_rate
+ cache_read_tokens * cache_read_rate
+ cache_creation_tokens * cache_creation_rate
)
output_weight = output_tokens * output_rate
total_weight = input_weight + output_weight
if total_weight <= 0:
return 0.0
return input_weight / total_weight
def _calculate_from_usd_cost(
usd_cost: float,
input_usd: float,
@@ -375,15 +377,16 @@ def _calculate_from_usd_cost(
input_msats = int((input_usd * sats_per_usd) * 1000)
output_msats = int((output_usd * sats_per_usd) * 1000)
else:
effective_input_tokens = (
input_tokens + cache_read_tokens + cache_creation_tokens
)
total_tokens = effective_input_tokens + output_tokens
input_msats = (
int(cost_in_msats * effective_input_tokens / total_tokens)
if total_tokens > 0
else 0
# No explicit split: weight by per-token rates so cheap cache reads do
# not inflate the input share (a raw token-count split would).
input_fraction = _rate_weighted_input_fraction(
response_data,
input_tokens,
cache_read_tokens,
cache_creation_tokens,
output_tokens,
)
input_msats = int(cost_in_msats * input_fraction)
output_msats = cost_in_msats - input_msats
logger.info(
+136
View File
@@ -85,6 +85,142 @@ class Model(BaseModel):
return hash(self.id)
def _litellm_entry(*keys: str) -> dict | None:
"""First litellm cost-map entry matching any of the given exact keys."""
import litellm
for key in keys:
candidate = litellm.model_cost.get(key)
if isinstance(candidate, dict):
return candidate
return None
def _family_tokens(lowered_id: str) -> list[str]:
"""Candidate family tokens derived from an id, most-specific first.
The provider prefix (``deepseek`` in ``deepseek/x``) and the leading segment
of the model name (``deepseek`` in ``deepseek-v4-flash``). Tokens shorter
than 4 chars are dropped so short noise like ``gpt`` cannot match unrelated
litellm keys.
"""
provider, sep, rest = lowered_id.partition("/")
name = rest if sep else provider
tokens: list[str] = []
if sep and len(provider) >= 4:
tokens.append(provider)
leading = name.split("-", 1)[0]
if len(leading) >= 4 and leading not in tokens:
tokens.append(leading)
return tokens
def _family_reference(lowered_id: str) -> dict | None:
"""First litellm entry sharing a family token AND carrying a cache-read rate.
Generic, data-driven fallback for vanity ids that proxies invent
(``deepseek-v4-flash``) which match no litellm key directly. Any provider
whose litellm snapshots price cache reads (deepseek, anthropic, ...) is
matched without a hand-maintained family list. The family token must be the
key's root segment (``deepseek/...`` / ``deepseek-...``), not appear anywhere
in it — that excludes reseller-prefixed snapshots (``deepinfra/deepseek/...``,
``novita/deepseek/...``) whose cache markup differs from the native provider.
Keys are scanned in sorted order so the chosen reference is deterministic.
"""
import litellm
for token in _family_tokens(lowered_id):
for key in sorted(litellm.model_cost):
lowered_key = key.lower()
if not (
lowered_key == token
or lowered_key.startswith(token + "/")
or lowered_key.startswith(token + "-")
):
continue
entry = litellm.model_cost.get(key)
if not isinstance(entry, dict):
continue
ref_input = entry.get("input_cost_per_token")
ref_read = entry.get("cache_read_input_token_cost")
if (
isinstance(ref_input, (int, float))
and ref_input > 0
and isinstance(ref_read, (int, float))
and ref_read > 0
):
return entry
return None
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. Without a cache rate, billing falls back to the full input
rate, overcharging cache reads (DeepSeek hits are ~10x cheaper). Two
strategies are tried in order:
1. **Exact match** — litellm ships absolute per-token USD rates keyed by the
upstream id (``deepseek/deepseek-chat``) or its bare model name
(``gpt-4o``); both spellings are tried and copied directly.
2. **Family ratio** — vanity ids proxies invent (``deepseek-v4-flash``) match
no litellm key. A reference entry for the same provider/family is found by
generic scan (no hand-maintained list) and its ``cache_read/input``
*ratio* is scaled by THIS model's own input price — correct even when the
vanity model is priced differently from the reference.
Rates already present (e.g. provided by OpenRouter) are authoritative and
never overwritten. Models matching no family 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
updated = Pricing.parse_obj(pricing.dict())
# 1. Exact match: absolute per-token USD rates apply directly.
info = _litellm_entry(model_id, model_id.split("/", 1)[-1])
if info is not None:
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)
needs_read = False
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)
needs_write = False
if not (needs_read or needs_write):
return updated
# 2. Family ratio for vanity ids (or fields the exact entry left unpriced).
if pricing.prompt <= 0.0:
return updated
reference = _family_reference(model_id.lower())
if reference is None:
logger.debug(
"No litellm cache-rate reference for model family",
extra={"model_id": model_id},
)
return updated
ref_input = reference.get("input_cost_per_token")
if not isinstance(ref_input, (int, float)) or ref_input <= 0:
return updated
if needs_read:
ref_read = reference.get("cache_read_input_token_cost")
if isinstance(ref_read, (int, float)) and ref_read > 0:
updated.input_cache_read = pricing.prompt * (ref_read / ref_input)
if needs_write:
ref_write = reference.get("cache_creation_input_token_cost")
if isinstance(ref_write, (int, float)) and ref_write > 0:
updated.input_cache_write = pricing.prompt * (ref_write / ref_input)
return updated
def _has_valid_pricing(model: dict) -> bool:
"""Check if model has valid pricing (not free, no negative values)."""
pricing = model.get("pricing", {})
+131
View File
@@ -0,0 +1,131 @@
"""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 / Azure / xAI / Groq / Moonshot / Qwen / Gemini-compat: cache reads in
``prompt_tokens_details.cached_tokens``, included in ``prompt_tokens``.
* OpenRouter: same as OpenAI plus cache *writes* in
``prompt_tokens_details.cache_write_tokens``, also included in
``prompt_tokens``.
* litellm-normalized: same nesting, but names the write field
``prompt_tokens_details.cache_creation_tokens`` (and additionally mirrors the
Anthropic top-level fields), with ``prompt_tokens`` as the grand total.
* Anthropic native: ``cache_read_input_tokens`` / ``cache_creation_input_tokens``
top-level, additive to (not included in) ``input_tokens``.
* DeepSeek: ``prompt_cache_hit_tokens`` / ``prompt_cache_miss_tokens``, with
``prompt_tokens = hit + miss``.
What decides whether cached tokens must be subtracted out of the input count is
**which prompt field the vendor uses**, not which cache field appears:
* ``prompt_tokens`` present -> cached + cache-write tokens are *included* in it
(OpenAI family, DeepSeek, OpenRouter, litellm); subtract both so
``input_tokens`` holds only the regular-rate portion.
* only ``input_tokens`` (Anthropic native) -> cached tokens are *additive*;
leave ``input_tokens`` untouched.
``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; a vendor whose fields
would genuinely conflict needs a dedicated branch here.
"""
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 _extract_cache_tokens(usage_data: dict) -> tuple[int, int]:
"""Pull (cache_read, cache_write) across all known dialects.
Precedence (highest first), independent for reads and writes:
* Anthropic top-level: ``cache_read_input_tokens`` /
``cache_creation_input_tokens``.
* Nested ``prompt_tokens_details``: ``cached_tokens`` for reads;
``cache_creation_tokens`` (litellm) or ``cache_write_tokens``
(OpenRouter) for writes.
* DeepSeek: ``prompt_cache_hit_tokens`` for reads (no write concept).
"""
cache_read = parse_token_count(usage_data.get("cache_read_input_tokens", 0))
cache_write = parse_token_count(usage_data.get("cache_creation_input_tokens", 0))
prompt_details = usage_data.get("prompt_tokens_details")
if isinstance(prompt_details, dict):
if not cache_read:
cache_read = parse_token_count(prompt_details.get("cached_tokens", 0))
if not cache_write:
cache_write = _first_token_count(
prompt_details, "cache_creation_tokens", "cache_write_tokens"
)
if not cache_read:
# DeepSeek: prompt_tokens = prompt_cache_hit_tokens + prompt_cache_miss_tokens
cache_read = parse_token_count(usage_data.get("prompt_cache_hit_tokens", 0))
return cache_read, cache_write
def normalize_usage(usage_data: object) -> NormalizedUsage | None:
"""Map a vendor usage dict onto the canonical shape, or None if absent.
Cached reads and writes are subtracted from the input count exactly once,
only for dialects that report a ``prompt_tokens`` grand total that already
includes them (OpenAI family, DeepSeek, OpenRouter, litellm). Anthropic
native reports them additively under ``input_tokens`` and is left untouched.
"""
if not isinstance(usage_data, dict):
return None
output_tokens = _first_token_count(
usage_data, "completion_tokens", "output_tokens"
)
cache_read, cache_write = _extract_cache_tokens(usage_data)
# ``prompt_tokens`` is the inclusive grand total; ``input_tokens`` (Anthropic
# native) excludes cached tokens. The field chosen decides whether to subtract.
if "prompt_tokens" in usage_data:
input_tokens = parse_token_count(usage_data.get("prompt_tokens", 0))
input_tokens = max(0, input_tokens - cache_read - cache_write)
else:
input_tokens = parse_token_count(usage_data.get("input_tokens", 0))
return NormalizedUsage(
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_read_tokens=cache_read,
cache_write_tokens=cache_write,
)
+134 -105
View File
@@ -35,11 +35,17 @@ 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 normalize_usage
from ..wallet import recieve_token, send_token
from . import messages_dispatch
from .cache_breakpoints import (
inject_anthropic_cache_breakpoints,
is_explicit_cache_model,
)
from .count_tokens import count_tokens_locally
from .litellm_routing import detect_litellm_prefix
@@ -161,6 +167,11 @@ class BaseUpstreamProvider:
``cache_creation_input_tokens`` fields are left in place for clients
that want the breakdown.
Only Anthropic-shaped responses that lack ``prompt_tokens`` need their
additive cache fields folded into ``input_tokens``. OpenAI/litellm/
DeepSeek-shaped responses report ``prompt_tokens`` as an inclusive
total already, so it must never be inflated by top-level cache fields.
For Anthropic-shaped responses (``input_tokens`` present), the cache
fields are forced to ``0`` when the upstream omitted them, so the
client always sees a consistent shape.
@@ -180,20 +191,13 @@ class BaseUpstreamProvider:
except (TypeError, ValueError):
return
extra = cache_read + cache_creation
if extra <= 0:
if extra <= 0 or "prompt_tokens" in usage:
return
if "input_tokens" in usage:
try:
usage["input_tokens"] = int(usage.get("input_tokens") or 0) + extra
except (TypeError, ValueError):
pass
if "prompt_tokens" in usage:
try:
usage["prompt_tokens"] = (
int(usage.get("prompt_tokens") or 0) + extra
)
except (TypeError, ValueError):
pass
def _apply_provider_field(self, response_json: object) -> None:
"""Stamp the routstr ``provider`` field onto an upstream response payload.
@@ -432,6 +436,19 @@ class BaseUpstreamProvider:
return body
def _upstream_accepts_cache_control(self) -> bool:
"""True when this upstream accepts explicit ``cache_control`` markers.
Only OpenRouter (documents Anthropic + Alibaba explicit caching) and the
native Anthropic API accept the markers. Stamping them toward an
automatic-cache or non-supporting upstream risks a 400, so injection is
confined to these. Base URL is also checked so an OpenRouter endpoint
configured through the generic provider is still recognised.
"""
if self.provider_type in ("openrouter", "anthropic"):
return True
return "openrouter.ai" in (self.base_url or "")
def prepare_request_body(
self, body: bytes | None, model_obj: Model
) -> bytes | None:
@@ -501,6 +518,27 @@ class BaseUpstreamProvider:
data["stream_options"] = merged
changed = True
# Explicit-cache models (Anthropic Claude, Alibaba Qwen / deepseek-v3.2)
# cache nothing without ``cache_control`` markers in the body. Clients
# that don't recognise a routstr URL as one of these never send them, so
# caching silently never engages over routstr even though it works
# against OpenRouter directly. Stamp the standard breakpoints so caching
# works by default, deferring to any client-set markers. Gated to
# upstreams that accept the markers (OpenRouter / Anthropic) so they
# never leak to an automatic-cache provider that would reject them.
if (
"messages" in data
and isinstance(data.get("messages"), list)
and self._upstream_accepts_cache_control()
and is_explicit_cache_model(
model_obj.id,
model_obj.forwarded_model_id,
model_obj.canonical_slug,
)
):
if inject_anthropic_cache_breakpoints(data):
changed = True
if changed:
return json.dumps(data).encode()
return body
@@ -1018,7 +1056,10 @@ 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,
)
await session.refresh(key)
@@ -1436,7 +1477,10 @@ 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,
)
await session.refresh(key)
@@ -1616,6 +1660,26 @@ class BaseUpstreamProvider:
total_cost, _coerce_usd(usage_or_root.get(field))
)
def _absorb_usage(usage: object) -> None:
nonlocal input_tokens, output_tokens
nonlocal cache_read_input_tokens, cache_creation_input_tokens
normalized = normalize_usage(usage)
if normalized is None:
return
input_tokens += normalized.input_tokens
output_tokens += normalized.output_tokens
# Anthropic message streams can restate the same cumulative
# cache snapshot across events; keep the existing max() behavior
# while allowing OpenAI/litellm/DeepSeek fields to be normalized.
cache_read_input_tokens = max(
cache_read_input_tokens,
normalized.cache_read_tokens,
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
normalized.cache_write_tokens,
)
async def finalize_without_usage() -> bytes | None:
nonlocal usage_finalized
if usage_finalized:
@@ -1631,7 +1695,10 @@ 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_finalized = True
return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
@@ -1678,62 +1745,11 @@ class BaseUpstreamProvider:
changed = True
if usage := msg.get("usage"):
input_tokens += usage.get("input_tokens", 0)
output_tokens += usage.get(
"output_tokens", 0
)
# Anthropic's `message_start.usage`
# carries the cumulative cache
# snapshot for the prompt — pick
# the max() so subsequent
# `message_delta.usage` events
# (which only restate the same
# numbers) don't double-count.
cache_read_input_tokens = max(
cache_read_input_tokens,
int(
usage.get(
"cache_read_input_tokens", 0
)
or 0
),
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
int(
usage.get(
"cache_creation_input_tokens",
0,
)
or 0
),
)
_absorb_usage(usage)
_absorb_usd(usage)
if usage := data.get("usage"):
input_tokens += usage.get("input_tokens", 0)
output_tokens += usage.get(
"output_tokens", 0
)
cache_read_input_tokens = max(
cache_read_input_tokens,
int(
usage.get(
"cache_read_input_tokens", 0
)
or 0
),
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
int(
usage.get(
"cache_creation_input_tokens",
0,
)
or 0
),
)
_absorb_usage(usage)
_absorb_usd(usage)
# Some upstreams attach cost fields at
# the event root rather than nested
@@ -1850,7 +1866,10 @@ 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,
)
self.inject_cost_metadata(response_json, cost_data, key)
@@ -1952,7 +1971,10 @@ 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,
)
self.inject_cost_metadata(response_json, cost_data, key)
@@ -3025,44 +3047,46 @@ class BaseUpstreamProvider:
extra={"model": model, "has_usage": "usage" in response_data},
)
async with create_session() as session:
match await calculate_cost(response_data, max_cost_for_model, session):
case MaxCostData() as cost:
logger.debug(
"Using max cost pricing",
extra={"model": model, "max_cost_msats": cost.total_msats},
)
return cost
case CostData() as cost:
logger.debug(
"Using token-based pricing",
extra={
"model": model,
"total_cost_msats": cost.total_msats,
"input_msats": cost.input_msats,
"output_msats": cost.output_msats,
},
)
return cost
case CostDataError() as error:
logger.error(
"Cost calculation error",
extra={
"model": model,
"error_message": error.message,
"error_code": error.code,
},
)
raise HTTPException(
status_code=400,
detail={
"error": {
"message": error.message,
"type": "invalid_request_error",
"code": error.code,
}
},
)
match await calculate_cost(
response_data,
max_cost_for_model,
):
case MaxCostData() as cost:
logger.debug(
"Using max cost pricing",
extra={"model": model, "max_cost_msats": cost.total_msats},
)
return cost
case CostData() as cost:
logger.debug(
"Using token-based pricing",
extra={
"model": model,
"total_cost_msats": cost.total_msats,
"input_msats": cost.input_msats,
"output_msats": cost.output_msats,
},
)
return cost
case CostDataError() as error:
logger.error(
"Cost calculation error",
extra={
"model": model,
"error_message": error.message,
"error_code": error.code,
},
)
raise HTTPException(
status_code=400,
detail={
"error": {
"message": error.message,
"type": "invalid_request_error",
"code": error.code,
}
},
)
return None
async def send_refund(
@@ -4562,14 +4586,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(
+157
View File
@@ -0,0 +1,157 @@
"""Inject explicit prompt-cache breakpoints into OpenAI-shaped requests.
Some upstreams cache *explicitly*: the request must carry
``cache_control: {"type": "ephemeral"}`` markers on the content blocks that
should be cached. Two model families use this identical wire format:
* **Anthropic Claude** — direct or via OpenRouter's ``anthropic/*`` models.
* **Alibaba's explicit-cache models on OpenRouter** — ``qwen/qwen3-max``,
``qwen/qwen-plus``, ``qwen/qwen3.6-plus``, ``qwen/qwen3-coder-plus``,
``qwen/qwen3-coder-flash`` and ``deepseek/deepseek-v3.2`` — which OpenRouter
documents as using "the same syntax as Anthropic explicit caching".
Every other provider routstr proxies (OpenAI, Azure, xAI/Grok, Groq, Moonshot,
default DeepSeek, Gemini implicit, Fireworks) caches *automatically* and needs
no markers — they are left untouched.
A client that doesn't know it is talking to one of these models *through*
routstr (e.g. an OpenAI-compatible coding agent pointed at a routstr URL) never
emits the markers — it only adds them when it recognises the provider as
OpenRouter. So caching silently never engages over routstr even though the same
client caches fine talking to OpenRouter directly.
This module restores caching by stamping the standard breakpoints onto the
forwarded body — the system prompt, the last tool, and the last conversation
message (the format allows up to four; we use three, matching the common
agent convention) — but only when the client supplied none of its own, so
explicit client control always wins. The caller is responsible for only
applying this toward an upstream that accepts the markers (OpenRouter /
Anthropic), so they never leak to a provider that would reject them.
"""
from __future__ import annotations
from typing import Any
# The single ephemeral marker stamped onto each chosen breakpoint. A 5-minute
# TTL (the default for ``ephemeral``) — deliberately not the 1h tier, which
# carries a higher cache-write premium and should stay opt-in.
EPHEMERAL_CACHE_CONTROL: dict[str, str] = {"type": "ephemeral"}
# Alibaba's explicit-cache models on OpenRouter. Matched as substrings of the
# model id (any spelling routstr carries). Snapshot endpoints that OpenRouter
# documents as *not* supporting explicit caching (e.g. ``qwen3.5-plus-02-15``)
# are different families and deliberately absent from this list.
_ALIBABA_EXPLICIT_CACHE_SLUGS: tuple[str, ...] = (
"qwen3-max",
"qwen-plus",
"qwen3.6-plus",
"qwen3-coder-plus",
"qwen3-coder-flash",
"deepseek-v3.2",
)
def is_explicit_cache_model(model_id: str | None, *fallbacks: str | None) -> bool:
"""True when the target model uses the explicit ``cache_control`` dialect.
Covers the Claude family (broadly — every Claude model supports it) and
Alibaba's documented explicit-cache models, across the id spellings routstr
carries: the OpenRouter id (``anthropic/claude-...``, ``qwen/qwen3-max``),
the bare upstream id, and any forwarded/canonical alias.
"""
for candidate in (model_id, *fallbacks):
if not candidate:
continue
lowered = candidate.lower()
if "claude" in lowered or "anthropic/" in lowered:
return True
if any(slug in lowered for slug in _ALIBABA_EXPLICIT_CACHE_SLUGS):
return True
return False
def _has_cache_control(obj: Any) -> bool:
"""Recursively detect any client-supplied ``cache_control`` marker."""
if isinstance(obj, dict):
if "cache_control" in obj:
return True
return any(_has_cache_control(v) for v in obj.values())
if isinstance(obj, list):
return any(_has_cache_control(v) for v in obj)
return False
def body_has_cache_control(data: dict) -> bool:
"""True when the request already carries cache_control on messages/tools."""
return _has_cache_control(data.get("messages")) or _has_cache_control(
data.get("tools")
)
def _stamp_text_content(message: dict) -> bool:
"""Add the ephemeral marker to a message's last text block.
A string content is promoted to the array form Anthropic requires for
cache markers; an existing array gets the marker on its last text part.
Returns True when a marker was placed.
"""
content = message.get("content")
if isinstance(content, str):
if not content:
return False
message["content"] = [
{
"type": "text",
"text": content,
"cache_control": dict(EPHEMERAL_CACHE_CONTROL),
}
]
return True
if isinstance(content, list):
for part in reversed(content):
if isinstance(part, dict) and part.get("type") == "text":
part["cache_control"] = dict(EPHEMERAL_CACHE_CONTROL)
return True
return False
def _stamp_system_prompt(messages: list) -> None:
for message in messages:
if isinstance(message, dict) and message.get("role") in (
"system",
"developer",
):
_stamp_text_content(message)
return
def _stamp_last_tool(tools: Any) -> None:
if isinstance(tools, list) and tools and isinstance(tools[-1], dict):
tools[-1]["cache_control"] = dict(EPHEMERAL_CACHE_CONTROL)
def _stamp_last_conversation_message(messages: list) -> None:
for message in reversed(messages):
if isinstance(message, dict) and message.get("role") in ("user", "assistant"):
if _stamp_text_content(message):
return
def inject_anthropic_cache_breakpoints(data: dict) -> bool:
"""Stamp ephemeral cache breakpoints onto an OpenAI-shaped chat body.
Mutates ``data`` in place (the established convention in
``prepare_request_body``) and returns True when anything changed. No-ops
when the body isn't chat-shaped or the client already set cache_control.
"""
messages = data.get("messages")
if not isinstance(messages, list) or not messages:
return False
if body_has_cache_control(data):
return False
_stamp_system_prompt(messages)
_stamp_last_tool(data.get("tools"))
_stamp_last_conversation_message(messages)
return True
+7 -4
View File
@@ -30,6 +30,7 @@ import litellm
from ..core import get_logger
from ..core.exceptions import UpstreamError
from ..payment.models import Model
from ..payment.usage import normalize_usage
logger = get_logger(__name__)
@@ -308,10 +309,12 @@ def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent:
def _accumulate(usage: dict) -> None:
nonlocal in_tokens, out_tokens, cache_read_tokens, cache_create_tokens
nonlocal total_cost, input_cost, output_cost
in_tokens += int(usage.get("input_tokens") or 0)
out_tokens += int(usage.get("output_tokens") or 0)
cache_read_tokens += int(usage.get("cache_read_input_tokens") or 0)
cache_create_tokens += int(usage.get("cache_creation_input_tokens") or 0)
normalized = normalize_usage(usage)
if normalized is not None:
in_tokens += normalized.input_tokens
out_tokens += normalized.output_tokens
cache_read_tokens += normalized.cache_read_tokens
cache_create_tokens += normalized.cache_write_tokens
total_cost += _coerce_float(usage.get("total_cost"))
input_cost += _coerce_float(usage.get("input_cost"))
output_cost += _coerce_float(usage.get("output_cost"))
+189
View File
@@ -0,0 +1,189 @@
"""Tests for Anthropic cache-breakpoint injection on forwarded requests.
Anthropic prompt caching is explicit; a client that doesn't recognise a routstr
URL as Anthropic-backed never sends ``cache_control`` markers, so caching never
engages over routstr. ``prepare_request_body`` must stamp the standard
breakpoints for Anthropic-family models while always deferring to client-set
markers and never touching automatic-cache providers.
"""
import json
import os
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.upstream import GenericUpstreamProvider
from routstr.upstream.cache_breakpoints import (
body_has_cache_control,
inject_anthropic_cache_breakpoints,
is_explicit_cache_model,
)
def _chat_body() -> dict:
return {
"model": "anthropic/claude-sonnet-4.5",
"stream": True,
"messages": [
{"role": "system", "content": "You are concise."},
{"role": "user", "content": "Hello"},
],
"tools": [
{"type": "function", "function": {"name": "a"}},
{"type": "function", "function": {"name": "b"}},
],
}
# ---------------------------------------------------------------------------
# Model detection
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
"model_id,expected",
[
("anthropic/claude-sonnet-4.5", True),
("claude-haiku-4-5-20251001", True),
# Alibaba explicit-cache models share Anthropic's wire format
("qwen/qwen3-max", True),
("qwen/qwen3-coder-plus", True),
("deepseek/deepseek-v3.2", True),
# Automatic-cache providers need no markers
("openai/gpt-4o", False),
("google/gemini-2.5-flash", False),
("deepseek/deepseek-chat", False),
("qwen/qwen3.5-plus-02-15", False), # snapshot, no explicit caching
(None, False),
],
)
def test_is_explicit_cache_model(model_id: str | None, expected: bool) -> None:
assert is_explicit_cache_model(model_id) is expected
def test_is_explicit_cache_model_uses_fallbacks() -> None:
# routstr id is opaque but a forwarded/canonical alias reveals the family.
assert is_explicit_cache_model("model-xyz", None, "anthropic/claude-opus-4.1")
# ---------------------------------------------------------------------------
# Breakpoint placement
# ---------------------------------------------------------------------------
def test_injects_three_breakpoints() -> None:
data = _chat_body()
assert inject_anthropic_cache_breakpoints(data) is True
# system prompt promoted to array form with a marker
system = data["messages"][0]["content"]
assert system == [
{
"type": "text",
"text": "You are concise.",
"cache_control": {"type": "ephemeral"},
}
]
# last tool marked
assert data["tools"][-1]["cache_control"] == {"type": "ephemeral"}
assert "cache_control" not in data["tools"][0]
# last user message marked
user = data["messages"][1]["content"]
assert user[-1]["cache_control"] == {"type": "ephemeral"}
def test_defers_to_client_supplied_cache_control() -> None:
data = _chat_body()
data["messages"][1]["content"] = [
{"type": "text", "text": "Hello", "cache_control": {"type": "ephemeral"}}
]
assert body_has_cache_control(data) is True
# No additional stamping when the client already controls caching.
assert inject_anthropic_cache_breakpoints(data) is False
assert "cache_control" not in data["tools"][-1]
def test_marks_last_text_part_of_array_content() -> None:
data = _chat_body()
data["messages"][1]["content"] = [
{"type": "text", "text": "first"},
{"type": "image_url", "image_url": {"url": "x"}},
{"type": "text", "text": "last"},
]
inject_anthropic_cache_breakpoints(data)
parts = data["messages"][1]["content"]
assert parts[2]["cache_control"] == {"type": "ephemeral"}
assert "cache_control" not in parts[0]
def test_noop_without_messages() -> None:
assert inject_anthropic_cache_breakpoints({"prompt": "x"}) is False
# ---------------------------------------------------------------------------
# prepare_request_body integration
# ---------------------------------------------------------------------------
def _model(model_id: str): # type: ignore[no-untyped-def]
from routstr.payment.models import Architecture, Model, Pricing
return Model(
id=model_id,
name=model_id,
created=0,
description="",
context_length=200000,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="Claude",
instruct_type=None,
),
pricing=Pricing(prompt=0.0, completion=0.0),
)
def _openrouter_provider() -> "GenericUpstreamProvider":
# OpenRouter endpoint via the generic provider — recognised by base URL.
return GenericUpstreamProvider(base_url="https://openrouter.ai/api/v1")
@pytest.mark.parametrize(
"model_id", ["anthropic/claude-sonnet-4.5", "qwen/qwen3-max", "deepseek/deepseek-v3.2"]
)
def test_prepare_request_body_injects_for_explicit_models(model_id: str) -> None:
provider = _openrouter_provider()
body = json.dumps(_chat_body()).encode()
out = provider.prepare_request_body(body, _model(model_id))
assert out is not None
data = json.loads(out)
assert body_has_cache_control(data) is True
assert data["tools"][-1]["cache_control"] == {"type": "ephemeral"}
def test_prepare_request_body_skips_for_automatic_provider_model() -> None:
provider = _openrouter_provider()
body = json.dumps(_chat_body()).encode()
out = provider.prepare_request_body(body, _model("openai/gpt-4o"))
assert out is not None
data = json.loads(out)
assert body_has_cache_control(data) is False
def test_prepare_request_body_skips_when_upstream_rejects_markers() -> None:
# Claude id but a non-OpenRouter/Anthropic upstream → must NOT inject,
# since the markers could be rejected by an upstream that doesn't accept them.
from routstr.upstream import GenericUpstreamProvider
provider = GenericUpstreamProvider(base_url="https://some-gateway.example/v1")
body = json.dumps(_chat_body()).encode()
out = provider.prepare_request_body(body, _model("anthropic/claude-sonnet-4.5"))
assert out is not None
data = json.loads(out)
assert body_has_cache_control(data) is False
+281
View File
@@ -0,0 +1,281 @@
"""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 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_backfill_vanity_deepseek_id_uses_generic_family_ratio() -> None:
"""The reported bug: a proxy exposes ``deepseek-v4-flash`` — a vanity id
litellm has never heard of. Without a family fallback its cache reads bill at
the full input rate (~10x overcharge on hits). The DeepSeek family's read
discount is found by generic litellm scan (no hard-coded list) and its
cache/input *ratio* applied to the model's own input price."""
pricing = Pricing(prompt=5e-07, completion=1.5e-06)
result = backfill_cache_pricing("deepseek-v4-flash", pricing)
ref = litellm.model_cost["deepseek/deepseek-chat"]
ratio = ref["cache_read_input_token_cost"] / ref["input_cost_per_token"]
assert result.input_cache_read == pytest.approx(pricing.prompt * ratio)
assert result.input_cache_read < pricing.prompt # sanity: it's a discount
def test_backfill_vanity_id_scales_to_own_input_price() -> None:
"""Family fallback uses the *ratio*, not the reference's absolute rate, so a
vanity model priced differently from the reference is billed against its own
input price the 10x input gap is preserved in the derived read rate."""
cheap = backfill_cache_pricing(
"deepseek-mini", Pricing(prompt=1e-07, completion=2e-07)
)
pricey = backfill_cache_pricing(
"deepseek-max", Pricing(prompt=1e-06, completion=2e-06)
)
assert pricey.input_cache_read == pytest.approx(cheap.input_cache_read * 10)
def test_backfill_vanity_id_no_family_match_unchanged() -> None:
"""A vanity id whose family token matches no litellm key stays untouched."""
result = backfill_cache_pricing(
"mystery-flash-9", Pricing(prompt=1e-06, completion=2e-06)
)
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)
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)
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)
assert isinstance(result, CostData)
# 1000 @ 1 + 9000 @ 1 (fallback) + 500 @ 2
assert result.total_msats == 11000
+161 -38
View File
@@ -1,10 +1,11 @@
"""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
from unittest.mock import AsyncMock, patch
from unittest.mock import patch
import pytest
@@ -16,12 +17,6 @@ from routstr.core.settings import settings
from routstr.payment.cost_calculation import CostData, MaxCostData, calculate_cost
@pytest.fixture
def mock_session() -> AsyncMock:
"""Mock AsyncSession for cost calculation tests."""
return AsyncMock()
@pytest.fixture(autouse=True)
def mock_fixed_pricing(monkeypatch: pytest.MonkeyPatch) -> None:
"""Mock settings and price to use fixed pricing."""
@@ -41,7 +36,7 @@ def patch_sats_usd_price() -> None: # type: ignore[misc]
# Test 1: OpenAI Cache Format
# ============================================================================
@pytest.mark.asyncio
async def test_openai_cache_subtraction(mock_session: AsyncMock) -> None:
async def test_openai_cache_subtraction() -> None:
"""OpenAI includes cached_tokens in prompt_tokens, subtract them."""
response = {
"model": "gpt-4",
@@ -53,7 +48,7 @@ async def test_openai_cache_subtraction(mock_session: AsyncMock) -> None:
}
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.input_tokens == 1000 # 2000 - 1000
@@ -65,7 +60,7 @@ async def test_openai_cache_subtraction(mock_session: AsyncMock) -> None:
# Test 2: Anthropic Cache Format
# ============================================================================
@pytest.mark.asyncio
async def test_anthropic_cache_additive(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_anthropic_cache_additive(mock_fixed_pricing: None) -> None:
"""Anthropic cache tokens are separate (additive) from input_tokens."""
response = {
"model": "claude-3-5-sonnet",
@@ -76,7 +71,7 @@ async def test_anthropic_cache_additive(mock_session: AsyncMock, mock_fixed_pric
"cache_read_input_tokens": 0,
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.input_tokens == 500
@@ -89,7 +84,7 @@ async def test_anthropic_cache_additive(mock_session: AsyncMock, mock_fixed_pric
# Test 3: Invalid Cache (Edge Case)
# ============================================================================
@pytest.mark.asyncio
async def test_cache_read_exceeds_prompt_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_cache_read_exceeds_prompt_tokens(mock_fixed_pricing: None) -> None:
"""Handle buggy upstream reporting cached > prompt_tokens."""
response = {
"model": "gpt-4",
@@ -101,7 +96,7 @@ async def test_cache_read_exceeds_prompt_tokens(mock_session: AsyncMock, mock_fi
}
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
# Should not go negative
assert isinstance(result, CostData)
@@ -114,7 +109,7 @@ async def test_cache_read_exceeds_prompt_tokens(mock_session: AsyncMock, mock_fi
# Test 4: Malformed Token Values
# ============================================================================
@pytest.mark.asyncio
async def test_malformed_cache_tokens_coerce_to_zero(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_malformed_cache_tokens_coerce_to_zero(mock_fixed_pricing: None) -> None:
"""Handle non-numeric cache token values."""
response = {
"model": "gpt-4",
@@ -127,7 +122,7 @@ async def test_malformed_cache_tokens_coerce_to_zero(mock_session: AsyncMock, mo
}
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
# Both should coerce to 0
assert isinstance(result, CostData)
@@ -139,7 +134,7 @@ async def test_malformed_cache_tokens_coerce_to_zero(mock_session: AsyncMock, mo
# Test 5: Anthropic Cache Not Subtracted
# ============================================================================
@pytest.mark.asyncio
async def test_anthropic_cache_not_subtracted(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_anthropic_cache_not_subtracted(mock_fixed_pricing: None) -> None:
"""Anthropic cache fields should NOT be subtracted from input_tokens."""
response = {
"model": "claude-3-5-sonnet",
@@ -149,7 +144,7 @@ async def test_anthropic_cache_not_subtracted(mock_session: AsyncMock, mock_fixe
"cache_read_input_tokens": 200, # ← Additive, don't subtract
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
# Anthropic: input_tokens stays as-is
assert isinstance(result, CostData)
@@ -161,7 +156,7 @@ async def test_anthropic_cache_not_subtracted(mock_session: AsyncMock, mock_fixe
# Test 6: Only Cache Read, No Regular Input
# ============================================================================
@pytest.mark.asyncio
async def test_only_cache_read_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_only_cache_read_tokens(mock_fixed_pricing: None) -> None:
"""Handle response with only cache read tokens."""
response = {
"model": "gpt-4",
@@ -173,7 +168,7 @@ async def test_only_cache_read_tokens(mock_session: AsyncMock, mock_fixed_pricin
}
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.input_tokens == 0 # max(0, 0 - 1000)
@@ -185,7 +180,7 @@ async def test_only_cache_read_tokens(mock_session: AsyncMock, mock_fixed_pricin
# Test 7: Only Cache Creation
# ============================================================================
@pytest.mark.asyncio
async def test_only_cache_creation_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_only_cache_creation_tokens(mock_fixed_pricing: None) -> None:
"""Handle response with only cache creation tokens (Anthropic)."""
response = {
"model": "claude-3-5-sonnet",
@@ -196,7 +191,7 @@ async def test_only_cache_creation_tokens(mock_session: AsyncMock, mock_fixed_pr
"cache_read_input_tokens": 0,
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.input_tokens == 500
@@ -209,7 +204,7 @@ async def test_only_cache_creation_tokens(mock_session: AsyncMock, mock_fixed_pr
# Test 8: Both Cache Read and Creation
# ============================================================================
@pytest.mark.asyncio
async def test_both_cache_read_and_creation(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_both_cache_read_and_creation(mock_fixed_pricing: None) -> None:
"""Handle response with both cache read and creation."""
response = {
"model": "claude-3-5-sonnet",
@@ -220,7 +215,7 @@ async def test_both_cache_read_and_creation(mock_session: AsyncMock, mock_fixed_
"cache_read_input_tokens": 500,
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.input_tokens == 300
@@ -233,7 +228,7 @@ async def test_both_cache_read_and_creation(mock_session: AsyncMock, mock_fixed_
# Test 9: Token Field Fallback
# ============================================================================
@pytest.mark.asyncio
async def test_token_field_fallback_order(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_token_field_fallback_order(mock_fixed_pricing: None) -> None:
"""Verify fallback order for token extraction."""
# When prompt_tokens is not present, fall back to input_tokens
response = {
@@ -243,7 +238,7 @@ async def test_token_field_fallback_order(mock_session: AsyncMock, mock_fixed_pr
"completion_tokens": 50,
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.input_tokens == 250
@@ -254,20 +249,21 @@ async def test_token_field_fallback_order(mock_session: AsyncMock, mock_fixed_pr
# Test 10: Float Token Values
# ============================================================================
@pytest.mark.asyncio
async def test_float_token_values_coerced_to_int(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_float_token_values_coerced_to_int(mock_fixed_pricing: None) -> None:
"""Handle float token values by converting to int."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 100.7, # Float
"completion_tokens": 50.3, # Float
"cache_read_input_tokens": 25.9, # Float
"prompt_tokens_details": {"cached_tokens": 25.9}, # Float
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.input_tokens == 100 # Floored
# cached_tokens are part of prompt_tokens (OpenAI dialect) → subtracted: 100 - 25
assert result.input_tokens == 75 # Floored
assert result.output_tokens == 50 # Floored
assert result.cache_read_input_tokens == 25 # Floored
@@ -276,7 +272,7 @@ async def test_float_token_values_coerced_to_int(mock_session: AsyncMock, mock_f
# Test 11: Boolean Cache Tokens
# ============================================================================
@pytest.mark.asyncio
async def test_boolean_cache_tokens_coerced_to_zero(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_boolean_cache_tokens_coerced_to_zero(mock_fixed_pricing: None) -> None:
"""Handle boolean cache token values by coercing to zero."""
response = {
"model": "gpt-4",
@@ -286,7 +282,7 @@ async def test_boolean_cache_tokens_coerced_to_zero(mock_session: AsyncMock, moc
"cache_read_input_tokens": True, # Boolean
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.cache_read_input_tokens == 0 # Boolean coerced to 0
@@ -297,7 +293,7 @@ async def test_boolean_cache_tokens_coerced_to_zero(mock_session: AsyncMock, moc
# Test 12: Zero Cache Tokens
# ============================================================================
@pytest.mark.asyncio
async def test_zero_cache_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_zero_cache_tokens(mock_fixed_pricing: None) -> None:
"""Handle explicit zero cache tokens."""
response = {
"model": "gpt-4",
@@ -309,21 +305,113 @@ async def test_zero_cache_tokens(mock_session: AsyncMock, mock_fixed_pricing: No
}
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.cache_read_input_tokens == 0
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() -> 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)
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() -> 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)
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() -> 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)
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() -> 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)
assert isinstance(result, CostData)
assert result.input_tokens == 1000
assert result.cache_read_input_tokens == 0
# ============================================================================
# Test 13: Missing Usage Block
# ============================================================================
@pytest.mark.asyncio
async def test_missing_usage_block(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_missing_usage_block(mock_fixed_pricing: None) -> None:
"""When usage is missing, return MaxCostData with zero tokens."""
response = {"model": "gpt-4", "choices": [{"message": {"content": "test"}}]}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, MaxCostData)
assert result.input_tokens == 0
@@ -335,11 +423,46 @@ async def test_missing_usage_block(mock_session: AsyncMock, mock_fixed_pricing:
# Test 14: Null Usage Block
# ============================================================================
@pytest.mark.asyncio
async def test_null_usage_block(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
async def test_null_usage_block(mock_fixed_pricing: None) -> None:
"""When usage is null, return MaxCostData with zero tokens."""
response = {"model": "gpt-4", "usage": None}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, MaxCostData)
assert result.input_tokens == 0
assert result.cache_read_input_tokens == 0
# ============================================================================
# USD-path input/output split — weight by rate, not raw token count
# ============================================================================
def test_usd_split_weights_by_rate_not_token_count() -> None:
"""A lump-sum USD cost with no explicit input/output breakdown is split by
per-token *rate*. Cheap cache-read tokens (~10x discount) must not inflate
the input share the way a raw token-count split does."""
from routstr.payment.cost_calculation import _rate_weighted_input_fraction
# rates: (input, output, cache_read, cache_write); relative scale only.
with patch(
"routstr.payment.cost_calculation._get_pricing_rates",
return_value=(1.0, 4.0, 0.1, 0.1),
):
frac = _rate_weighted_input_fraction({"model": "x"}, 100, 900, 0, 100)
# input weight = 100*1 + 900*0.1 = 190; output = 100*4 = 400; total = 590.
assert frac == pytest.approx(190 / 590)
# The old raw token-count split gave (100+900)/1100 ≈ 0.909 — wildly skewed.
assert frac < 0.5
def test_usd_split_falls_back_to_token_count_without_rates() -> None:
"""When rates are unavailable (fixed pricing / unknown model) the split
degrades gracefully to the raw token-count proportion."""
from routstr.payment.cost_calculation import _rate_weighted_input_fraction
with patch(
"routstr.payment.cost_calculation._get_pricing_rates", return_value=None
):
frac = _rate_weighted_input_fraction({"model": "x"}, 100, 900, 0, 100)
assert frac == pytest.approx(1000 / 1100)
+106 -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
@@ -649,6 +649,110 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None:
assert combined["model"] == "openai/gpt-4o-mini"
@pytest.mark.asyncio
@pytest.mark.parametrize(
("usage", "expected"),
[
(
{
"prompt_tokens": 10_000,
"completion_tokens": 500,
"prompt_cache_hit_tokens": 9_000,
"prompt_cache_miss_tokens": 1_000,
},
{
"input_tokens": 1_000,
"output_tokens": 500,
"cache_read_input_tokens": 9_000,
"cache_creation_input_tokens": 0,
},
),
(
{
"prompt_tokens": 10_000,
"completion_tokens": 100,
"prompt_tokens_details": {
"cached_tokens": 5_000,
"cache_creation_tokens": 2_000,
},
},
{
"input_tokens": 3_000,
"output_tokens": 100,
"cache_read_input_tokens": 5_000,
"cache_creation_input_tokens": 2_000,
},
),
],
)
async def test_streaming_litellm_messages_normalizes_cache_usage_dialects(
usage: dict[str, Any], expected: dict[str, int]
) -> None:
provider = _make_provider()
key = _make_key()
model = _make_model()
session = _make_session()
body = _anthropic_request_body(stream=True)
async def fake_chunks() -> AsyncIterator[dict]:
yield {
"type": "message_start",
"message": {
"id": "msg_1",
"type": "message",
"role": "assistant",
"model": "openai/gpt-4o-mini",
"content": [],
},
}
yield {"type": "message_delta", "delta": {}, "usage": usage}
yield {"type": "message_stop"}
captured: dict[str, Any] = {}
async def fake_adjust(
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None
) -> dict:
captured["combined_data"] = json.loads(json.dumps(combined_data))
return {"total_msats": 999, "total_usd": 0.0001}
fake_session = MagicMock()
fake_session.get = AsyncMock(return_value=key)
class FakeSessionCtx:
async def __aenter__(self) -> Any:
return fake_session
async def __aexit__(self, *args: Any) -> None:
return None
with (
patch(
"litellm.anthropic.messages.acreate",
new=AsyncMock(return_value=fake_chunks()),
),
patch(
"routstr.upstream.base.adjust_payment_for_tokens",
new=AsyncMock(side_effect=fake_adjust),
),
patch("routstr.upstream.base.create_session", new=lambda: FakeSessionCtx()),
):
result = await provider._forward_messages_via_litellm(
request_body=body,
key=key,
session=session,
max_cost_for_model=10_000,
model_obj=model,
)
assert isinstance(result, StreamingResponse)
async for _ in result.body_iterator:
pass
combined = captured["combined_data"]
for field, value in expected.items():
assert combined["usage"][field] == value
# ---------------------------------------------------------------------------
# x-cashu non-streaming
# ---------------------------------------------------------------------------
+107
View File
@@ -20,6 +20,7 @@ comment ever reaches the client. That invariant is exactly what the buggy
import json
from collections.abc import AsyncGenerator
from typing import cast
from unittest.mock import AsyncMock, MagicMock
import pytest
@@ -88,6 +89,15 @@ def _data_payloads(out: list[bytes]) -> list[bytes]:
return payloads
def _last_adjustment_input() -> dict:
mock = cast(AsyncMock, base.adjust_payment_for_tokens)
await_args = mock.await_args
assert await_args is not None
adjustment_input = await_args.args[1]
assert isinstance(adjustment_input, dict)
return adjustment_input
def _assert_clean(out: list[bytes]) -> list[dict]:
"""Core invariant: every data line is [DONE] or valid JSON; no comments leak."""
blob = b"".join(out)
@@ -107,6 +117,103 @@ def _assert_clean(out: list[bytes]) -> list[dict]:
return objs
def test_fold_cache_does_not_inflate_inclusive_prompt_tokens() -> None:
"""OpenAI/litellm/DeepSeek prompt_tokens already includes cache tokens."""
usage = {
"prompt_tokens": 10000,
"completion_tokens": 100,
"cache_read_input_tokens": 5000,
"cache_creation_input_tokens": 2000,
}
BaseUpstreamProvider._fold_cache_into_input_tokens(usage)
assert usage["prompt_tokens"] == 10000
assert usage["cache_read_input_tokens"] == 5000
assert usage["cache_creation_input_tokens"] == 2000
def test_fold_cache_adds_only_anthropic_additive_input_tokens() -> None:
"""Anthropic native input_tokens excludes cache fields and remains folded."""
usage = {
"input_tokens": 300,
"output_tokens": 100,
"cache_read_input_tokens": 500,
"cache_creation_input_tokens": 2000,
}
BaseUpstreamProvider._fold_cache_into_input_tokens(usage)
assert usage["input_tokens"] == 2800
@pytest.mark.asyncio
async def test_deepseek_streaming_usage_preserved_for_cost_adjustment() -> None:
"""Streaming chat usage keeps DeepSeek cache fields for normalize_usage."""
usage = {
"prompt_tokens": 10000,
"completion_tokens": 500,
"prompt_cache_hit_tokens": 9000,
"prompt_cache_miss_tokens": 1000,
}
chunks = [
b'data: {"id":"x","model":"deepseek-chat","choices":[{"delta":{"content":"ok"}}]}\n\n',
b"data: "
+ json.dumps(
{
"id": "x",
"model": "deepseek-chat",
"choices": [],
"usage": usage,
}
).encode()
+ b"\n\n",
b"data: [DONE]\n\n",
]
out = await _drive(chunks)
_assert_clean(out)
adjustment_input = _last_adjustment_input()
for key, value in usage.items():
assert adjustment_input["usage"][key] == value
@pytest.mark.asyncio
async def test_openai_litellm_streaming_usage_preserved_for_cost_adjustment() -> None:
"""Streaming chat usage keeps nested and top-level cache fields."""
usage = {
"prompt_tokens": 10000,
"completion_tokens": 100,
"cache_read_input_tokens": 5000,
"cache_creation_input_tokens": 2000,
"prompt_tokens_details": {
"cached_tokens": 5000,
"cache_creation_tokens": 2000,
},
}
chunks = [
b"data: "
+ json.dumps(
{
"id": "x",
"model": "claude-litellm",
"choices": [],
"usage": usage,
}
).encode()
+ b"\n\n",
b"data: [DONE]\n\n",
]
out = await _drive(chunks)
_assert_clean(out)
adjustment_input = _last_adjustment_input()
for key, value in usage.items():
assert adjustment_input["usage"][key] == value
@pytest.mark.asyncio
async def test_openai_style_plain_stream() -> None:
"""OpenAI / Groq / Fireworks / xAI / Perplexity: plain data + [DONE]."""
+140
View File
@@ -0,0 +1,140 @@
"""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). ``calculate_cost`` normalizes the response's usage
object with this parser and needs no vendor knowledge of its own.
"""
import os
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to")
import pytest
from routstr.payment.usage import NormalizedUsage, normalize_usage
# ============================================================================
# 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),
),
# OpenRouter: cache writes nested as prompt_tokens_details.cache_write_tokens,
# both reads and writes included in prompt_tokens → both subtracted
(
{
"prompt_tokens": 10000,
"completion_tokens": 60,
"prompt_tokens_details": {
"cached_tokens": 5000,
"cache_write_tokens": 2000,
},
},
NormalizedUsage(
input_tokens=3000,
output_tokens=60,
cache_read_tokens=5000,
cache_write_tokens=2000,
),
),
# litellm-normalized Anthropic: prompt_tokens is the grand total and the
# write field is named cache_creation_tokens; top-level fields mirror it.
# prompt_tokens present → both subtracted (NOT additive like native).
(
{
"prompt_tokens": 10000,
"completion_tokens": 100,
"cache_read_input_tokens": 5000,
"cache_creation_input_tokens": 2000,
"prompt_tokens_details": {
"cached_tokens": 5000,
"cache_creation_tokens": 2000,
},
},
NormalizedUsage(
input_tokens=3000,
output_tokens=100,
cache_read_tokens=5000,
cache_write_tokens=2000,
),
),
],
)
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