mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-08-05 09:34:36 +00:00
Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
14785e4cde | ||
|
|
1dceffffa7 | ||
|
|
4b74a9cf81 | ||
|
|
3be2da6728 | ||
|
|
17d690e77e | ||
|
|
cbc424e8e7 | ||
|
|
068fb3572f | ||
|
|
eaf74edbba |
@@ -98,7 +98,7 @@ docker-down:
|
|||||||
lint:
|
lint:
|
||||||
@echo "🔍 Running linting checks..."
|
@echo "🔍 Running linting checks..."
|
||||||
$(RUFF) check .
|
$(RUFF) check .
|
||||||
$(MYPY) routstr/ --ignore-missing-imports
|
$(MYPY) .
|
||||||
|
|
||||||
format:
|
format:
|
||||||
@echo "✨ Formatting code..."
|
@echo "✨ Formatting code..."
|
||||||
@@ -107,7 +107,7 @@ format:
|
|||||||
|
|
||||||
type-check:
|
type-check:
|
||||||
@echo "🔎 Running type checks..."
|
@echo "🔎 Running type checks..."
|
||||||
$(MYPY) routstr/ --ignore-missing-imports
|
$(MYPY) .
|
||||||
|
|
||||||
# Development setup
|
# Development setup
|
||||||
dev-setup:
|
dev-setup:
|
||||||
@@ -234,7 +234,7 @@ ci-test:
|
|||||||
ci-lint:
|
ci-lint:
|
||||||
@echo "🤖 Running CI linting..."
|
@echo "🤖 Running CI linting..."
|
||||||
$(RUFF) check . --exit-non-zero-on-fix
|
$(RUFF) check . --exit-non-zero-on-fix
|
||||||
$(MYPY) routstr/ --ignore-missing-imports --no-error-summary
|
$(MYPY) . --no-error-summary
|
||||||
|
|
||||||
# Debug helpers
|
# Debug helpers
|
||||||
test-debug:
|
test-debug:
|
||||||
|
|||||||
+8
-2
@@ -708,12 +708,18 @@ async def revert_pay_for_request(
|
|||||||
|
|
||||||
|
|
||||||
async def adjust_payment_for_tokens(
|
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:
|
) -> dict:
|
||||||
"""
|
"""
|
||||||
Adjusts the payment based on token usage in the response.
|
Adjusts the payment based on token usage in the response.
|
||||||
This is called after the initial payment and the upstream request is complete.
|
This is called after the initial payment and the upstream request is complete.
|
||||||
Returns cost data to be included in the response.
|
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)
|
billing_key = await get_billing_key(key, session)
|
||||||
model = response_data.get("model", "unknown")
|
model = response_data.get("model", "unknown")
|
||||||
@@ -796,7 +802,7 @@ async def adjust_payment_for_tokens(
|
|||||||
extra={"error": str(e), "fee_msats": fee_msats},
|
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:
|
case MaxCostData() as cost:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Using max cost data (no token adjustment)",
|
"Using max cost data (no token adjustment)",
|
||||||
|
|||||||
@@ -3,9 +3,17 @@ import math
|
|||||||
from pydantic.v1 import BaseModel
|
from pydantic.v1 import BaseModel
|
||||||
|
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
from ..core.db import AsyncSession
|
|
||||||
from ..core.settings import settings
|
from ..core.settings import settings
|
||||||
from .price import sats_usd_price
|
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__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
@@ -34,7 +42,8 @@ class CostDataError(BaseModel):
|
|||||||
|
|
||||||
|
|
||||||
async def calculate_cost(
|
async def calculate_cost(
|
||||||
response_data: dict, max_cost: int, session: AsyncSession
|
response_data: dict,
|
||||||
|
max_cost: int,
|
||||||
) -> CostData | MaxCostData | CostDataError:
|
) -> CostData | MaxCostData | CostDataError:
|
||||||
"""Calculate the cost of an API request based on token usage.
|
"""Calculate the cost of an API request based on token usage.
|
||||||
|
|
||||||
@@ -44,6 +53,9 @@ async def calculate_cost(
|
|||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Cost data or error information
|
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(
|
logger.debug(
|
||||||
"Starting cost calculation",
|
"Starting cost calculation",
|
||||||
@@ -54,8 +66,9 @@ async def calculate_cost(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
# Check for usage data
|
usage = normalize_usage(response_data.get("usage"))
|
||||||
if "usage" not in response_data or response_data["usage"] is None:
|
|
||||||
|
if usage is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"No usage data in response — billing at MaxCostData with zero "
|
"No usage data in response — billing at MaxCostData with zero "
|
||||||
"tokens. Dashboard will show this request as `(0+0)`. Most "
|
"tokens. Dashboard will show this request as `(0+0)`. Most "
|
||||||
@@ -84,16 +97,14 @@ async def calculate_cost(
|
|||||||
cache_creation_msats=0,
|
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 = usage.input_tokens
|
||||||
input_tokens = _extract_token_pair(usage_data, "prompt_tokens", "input_tokens")
|
output_tokens = usage.output_tokens
|
||||||
output_tokens = _extract_token_pair(usage_data, "completion_tokens", "output_tokens")
|
cache_read_tokens = usage.cache_read_tokens
|
||||||
|
cache_creation_tokens = usage.cache_write_tokens
|
||||||
# Extract cache tokens (handles OpenAI vs Anthropic formats)
|
|
||||||
cache_read_tokens, cache_creation_tokens, input_tokens = _extract_cache_tokens(
|
|
||||||
usage_data, input_tokens
|
|
||||||
)
|
|
||||||
|
|
||||||
# Try USD cost first
|
# Try USD cost first
|
||||||
usd_cost = _resolve_usd_cost(usage_data, response_data)
|
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:
|
def _coerce_usd(value: object) -> float:
|
||||||
"""Coerce a value to USD float, handling various formats safely."""
|
"""Coerce a value to USD float, handling various formats safely."""
|
||||||
if value is None or isinstance(value, bool):
|
if value is None or isinstance(value, bool):
|
||||||
@@ -230,37 +225,6 @@ def _coerce_usd(value: object) -> float:
|
|||||||
return 0.0
|
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:
|
def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
|
||||||
"""Resolve USD cost with clear priority order.
|
"""Resolve USD cost with clear priority order.
|
||||||
|
|
||||||
@@ -454,9 +418,11 @@ def _calculate_from_tokens(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
visible_input_msats = int(calc_input_msats + calc_cache_read_msats)
|
||||||
|
|
||||||
return CostData(
|
return CostData(
|
||||||
base_msats=0,
|
base_msats=0,
|
||||||
input_msats=int(calc_input_msats),
|
input_msats=visible_input_msats,
|
||||||
output_msats=int(calc_output_msats),
|
output_msats=int(calc_output_msats),
|
||||||
total_msats=token_based_cost,
|
total_msats=token_based_cost,
|
||||||
total_usd=total_usd,
|
total_usd=total_usd,
|
||||||
|
|||||||
@@ -85,6 +85,48 @@ class Model(BaseModel):
|
|||||||
return hash(self.id)
|
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:
|
def _has_valid_pricing(model: dict) -> bool:
|
||||||
"""Check if model has valid pricing (not free, no negative values)."""
|
"""Check if model has valid pricing (not free, no negative values)."""
|
||||||
pricing = model.get("pricing", {})
|
pricing = model.get("pricing", {})
|
||||||
|
|||||||
@@ -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,
|
||||||
|
)
|
||||||
+144
-70
@@ -35,11 +35,17 @@ from ..payment.models import (
|
|||||||
Pricing,
|
Pricing,
|
||||||
_calculate_usd_max_costs,
|
_calculate_usd_max_costs,
|
||||||
_update_model_sats_pricing,
|
_update_model_sats_pricing,
|
||||||
|
backfill_cache_pricing,
|
||||||
list_models,
|
list_models,
|
||||||
)
|
)
|
||||||
from ..payment.price import sats_usd_price
|
from ..payment.price import sats_usd_price
|
||||||
|
from ..payment.usage import normalize_usage
|
||||||
from ..wallet import recieve_token, send_token
|
from ..wallet import recieve_token, send_token
|
||||||
from . import messages_dispatch
|
from . import messages_dispatch
|
||||||
|
from .cache_breakpoints import (
|
||||||
|
inject_anthropic_cache_breakpoints,
|
||||||
|
is_explicit_cache_model,
|
||||||
|
)
|
||||||
from .count_tokens import count_tokens_locally
|
from .count_tokens import count_tokens_locally
|
||||||
from .litellm_routing import detect_litellm_prefix
|
from .litellm_routing import detect_litellm_prefix
|
||||||
|
|
||||||
@@ -152,48 +158,56 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _fold_cache_into_input_tokens(usage: object) -> None:
|
def _fold_cache_into_input_tokens(usage: object) -> None:
|
||||||
"""Fold cache token counts into ``input_tokens`` / ``prompt_tokens``.
|
"""Fold additive cache token counts into Anthropic ``input_tokens``.
|
||||||
|
|
||||||
Cost calculation has already used the per-bucket counts to bill the
|
Cost calculation has already used the per-bucket counts to bill the
|
||||||
request correctly; what the client sees in the visible token total
|
request correctly; what the client sees in Anthropic-shaped visible
|
||||||
should be a single rolled-up prompt count *including* the cache
|
token totals should be a single rolled-up input count *including* the
|
||||||
portion. The standalone ``cache_read_input_tokens`` /
|
cache portion. OpenAI-compatible ``prompt_tokens`` is already inclusive
|
||||||
|
(DeepSeek, OpenAI, OpenRouter, litellm), so adding cache fields there
|
||||||
|
would double-count.
|
||||||
|
|
||||||
|
The standalone ``cache_read_input_tokens`` /
|
||||||
``cache_creation_input_tokens`` fields are left in place for clients
|
``cache_creation_input_tokens`` fields are left in place for clients
|
||||||
that want the breakdown.
|
that want the breakdown.
|
||||||
|
|
||||||
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.
|
|
||||||
"""
|
"""
|
||||||
if not isinstance(usage, dict):
|
if not isinstance(usage, dict):
|
||||||
return
|
return
|
||||||
|
|
||||||
# Normalise missing cache fields to 0 on Anthropic-shaped usage so
|
# ``prompt_tokens`` is an inclusive OpenAI-compatible grand total.
|
||||||
# downstream consumers can rely on them being present.
|
# Fold only native Anthropic-style usage where ``input_tokens`` excludes
|
||||||
if "input_tokens" in usage:
|
# cache reads/writes.
|
||||||
usage.setdefault("cache_read_input_tokens", 0)
|
if "input_tokens" not in usage or "prompt_tokens" in usage:
|
||||||
usage.setdefault("cache_creation_input_tokens", 0)
|
return
|
||||||
|
|
||||||
|
usage.setdefault("cache_read_input_tokens", 0)
|
||||||
|
usage.setdefault("cache_creation_input_tokens", 0)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
cache_read = int(usage.get("cache_read_input_tokens") or 0)
|
cache_read = int(usage.get("cache_read_input_tokens") or 0)
|
||||||
cache_creation = int(usage.get("cache_creation_input_tokens") or 0)
|
cache_creation = int(usage.get("cache_creation_input_tokens") or 0)
|
||||||
|
input_tokens = int(usage.get("input_tokens") or 0)
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return
|
return
|
||||||
extra = cache_read + cache_creation
|
extra = cache_read + cache_creation
|
||||||
if extra <= 0:
|
if extra > 0:
|
||||||
|
usage["input_tokens"] = input_tokens + extra
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _add_normalized_usage_fields(response_json: object) -> None:
|
||||||
|
"""Preserve raw usage while adding canonical fields for billing/display."""
|
||||||
|
if not isinstance(response_json, dict):
|
||||||
return
|
return
|
||||||
if "input_tokens" in usage:
|
usage = response_json.get("usage")
|
||||||
try:
|
if not isinstance(usage, dict):
|
||||||
usage["input_tokens"] = int(usage.get("input_tokens") or 0) + extra
|
return
|
||||||
except (TypeError, ValueError):
|
normalized = normalize_usage(usage)
|
||||||
pass
|
if normalized is None:
|
||||||
if "prompt_tokens" in usage:
|
return
|
||||||
try:
|
usage.setdefault("input_tokens", normalized.input_tokens)
|
||||||
usage["prompt_tokens"] = (
|
usage.setdefault("output_tokens", normalized.output_tokens)
|
||||||
int(usage.get("prompt_tokens") or 0) + extra
|
usage.setdefault("cache_read_input_tokens", normalized.cache_read_tokens)
|
||||||
)
|
usage.setdefault("cache_creation_input_tokens", normalized.cache_write_tokens)
|
||||||
except (TypeError, ValueError):
|
|
||||||
pass
|
|
||||||
|
|
||||||
def _apply_provider_field(self, response_json: object) -> None:
|
def _apply_provider_field(self, response_json: object) -> None:
|
||||||
"""Stamp the routstr ``provider`` field onto an upstream response payload.
|
"""Stamp the routstr ``provider`` field onto an upstream response payload.
|
||||||
@@ -432,6 +446,19 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
return body
|
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(
|
def prepare_request_body(
|
||||||
self, body: bytes | None, model_obj: Model
|
self, body: bytes | None, model_obj: Model
|
||||||
) -> bytes | None:
|
) -> bytes | None:
|
||||||
@@ -501,6 +528,27 @@ class BaseUpstreamProvider:
|
|||||||
data["stream_options"] = merged
|
data["stream_options"] = merged
|
||||||
changed = True
|
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:
|
if changed:
|
||||||
return json.dumps(data).encode()
|
return json.dumps(data).encode()
|
||||||
return body
|
return body
|
||||||
@@ -819,6 +867,7 @@ class BaseUpstreamProvider:
|
|||||||
k: v for k, v in obj.items() if k != "choices"
|
k: v for k, v in obj.items() if k != "choices"
|
||||||
}
|
}
|
||||||
usage_chunk_data["choices"] = []
|
usage_chunk_data["choices"] = []
|
||||||
|
self._add_normalized_usage_fields(usage_chunk_data)
|
||||||
# Forward the content now, without usage, so token
|
# Forward the content now, without usage, so token
|
||||||
# usage is reported exactly once (in the trailer).
|
# usage is reported exactly once (in the trailer).
|
||||||
forward = {k: v for k, v in obj.items() if k != "usage"}
|
forward = {k: v for k, v in obj.items() if k != "usage"}
|
||||||
@@ -830,6 +879,7 @@ class BaseUpstreamProvider:
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
usage_chunk_data = obj
|
usage_chunk_data = obj
|
||||||
|
self._add_normalized_usage_fields(usage_chunk_data)
|
||||||
return
|
return
|
||||||
yield prefix + b"data: " + json.dumps(obj).encode() + b"\n\n"
|
yield prefix + b"data: " + json.dumps(obj).encode() + b"\n\n"
|
||||||
else:
|
else:
|
||||||
@@ -867,6 +917,8 @@ class BaseUpstreamProvider:
|
|||||||
if fresh_key:
|
if fresh_key:
|
||||||
cost_data: dict
|
cost_data: dict
|
||||||
try:
|
try:
|
||||||
|
if usage_chunk_data is not None:
|
||||||
|
self._add_normalized_usage_fields(usage_chunk_data)
|
||||||
adjustment_input = (
|
adjustment_input = (
|
||||||
usage_chunk_data
|
usage_chunk_data
|
||||||
if usage_chunk_data is not None
|
if usage_chunk_data is not None
|
||||||
@@ -1018,7 +1070,10 @@ class BaseUpstreamProvider:
|
|||||||
response_json["id"] = f"chatcmpl-{uuid.uuid4()}"
|
response_json["id"] = f"chatcmpl-{uuid.uuid4()}"
|
||||||
|
|
||||||
cost_data = await adjust_payment_for_tokens(
|
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)
|
await session.refresh(key)
|
||||||
@@ -1436,7 +1491,10 @@ class BaseUpstreamProvider:
|
|||||||
response_json["id"] = f"chatcmpl-{uuid.uuid4()}"
|
response_json["id"] = f"chatcmpl-{uuid.uuid4()}"
|
||||||
|
|
||||||
cost_data = await adjust_payment_for_tokens(
|
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)
|
await session.refresh(key)
|
||||||
@@ -1631,7 +1689,10 @@ class BaseUpstreamProvider:
|
|||||||
"usage": None,
|
"usage": None,
|
||||||
}
|
}
|
||||||
cost_data = await adjust_payment_for_tokens(
|
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
|
usage_finalized = True
|
||||||
return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
|
return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
|
||||||
@@ -1850,7 +1911,10 @@ class BaseUpstreamProvider:
|
|||||||
response_json["usage"] = {"input_tokens": input_tokens}
|
response_json["usage"] = {"input_tokens": input_tokens}
|
||||||
|
|
||||||
cost_data = await adjust_payment_for_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)
|
self.inject_cost_metadata(response_json, cost_data, key)
|
||||||
@@ -1952,7 +2016,10 @@ class BaseUpstreamProvider:
|
|||||||
response_json["model"] = requested_model
|
response_json["model"] = requested_model
|
||||||
|
|
||||||
cost_data = await adjust_payment_for_tokens(
|
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)
|
self.inject_cost_metadata(response_json, cost_data, key)
|
||||||
|
|
||||||
@@ -3025,44 +3092,46 @@ class BaseUpstreamProvider:
|
|||||||
extra={"model": model, "has_usage": "usage" in response_data},
|
extra={"model": model, "has_usage": "usage" in response_data},
|
||||||
)
|
)
|
||||||
|
|
||||||
async with create_session() as session:
|
match await calculate_cost(
|
||||||
match await calculate_cost(response_data, max_cost_for_model, session):
|
response_data,
|
||||||
case MaxCostData() as cost:
|
max_cost_for_model,
|
||||||
logger.debug(
|
):
|
||||||
"Using max cost pricing",
|
case MaxCostData() as cost:
|
||||||
extra={"model": model, "max_cost_msats": cost.total_msats},
|
logger.debug(
|
||||||
)
|
"Using max cost pricing",
|
||||||
return cost
|
extra={"model": model, "max_cost_msats": cost.total_msats},
|
||||||
case CostData() as cost:
|
)
|
||||||
logger.debug(
|
return cost
|
||||||
"Using token-based pricing",
|
case CostData() as cost:
|
||||||
extra={
|
logger.debug(
|
||||||
"model": model,
|
"Using token-based pricing",
|
||||||
"total_cost_msats": cost.total_msats,
|
extra={
|
||||||
"input_msats": cost.input_msats,
|
"model": model,
|
||||||
"output_msats": cost.output_msats,
|
"total_cost_msats": cost.total_msats,
|
||||||
},
|
"input_msats": cost.input_msats,
|
||||||
)
|
"output_msats": cost.output_msats,
|
||||||
return cost
|
},
|
||||||
case CostDataError() as error:
|
)
|
||||||
logger.error(
|
return cost
|
||||||
"Cost calculation error",
|
case CostDataError() as error:
|
||||||
extra={
|
logger.error(
|
||||||
"model": model,
|
"Cost calculation error",
|
||||||
"error_message": error.message,
|
extra={
|
||||||
"error_code": error.code,
|
"model": model,
|
||||||
},
|
"error_message": error.message,
|
||||||
)
|
"error_code": error.code,
|
||||||
raise HTTPException(
|
},
|
||||||
status_code=400,
|
)
|
||||||
detail={
|
raise HTTPException(
|
||||||
"error": {
|
status_code=400,
|
||||||
"message": error.message,
|
detail={
|
||||||
"type": "invalid_request_error",
|
"error": {
|
||||||
"code": error.code,
|
"message": error.message,
|
||||||
}
|
"type": "invalid_request_error",
|
||||||
},
|
"code": error.code,
|
||||||
)
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def send_refund(
|
async def send_refund(
|
||||||
@@ -4562,14 +4631,19 @@ class BaseUpstreamProvider:
|
|||||||
def _apply_provider_fee_to_model(self, model: Model) -> Model:
|
def _apply_provider_fee_to_model(self, model: Model) -> Model:
|
||||||
"""Apply provider fee to model's USD pricing and calculate max costs.
|
"""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:
|
Args:
|
||||||
model: Model object to update
|
model: Model object to update
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Model with provider fee applied to pricing and max costs calculated
|
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(
|
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(
|
temp_model = Model(
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -0,0 +1,243 @@
|
|||||||
|
"""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_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.
|
||||||
|
# Client-visible input_msats includes cache-read input cost for display,
|
||||||
|
# while cache_read_msats keeps the detailed breakdown.
|
||||||
|
assert result.input_msats == 1900
|
||||||
|
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
|
||||||
@@ -1,10 +1,11 @@
|
|||||||
"""Tests for cache token handling in cost calculation.
|
"""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
|
import os
|
||||||
from unittest.mock import AsyncMock, patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
@@ -16,12 +17,6 @@ from routstr.core.settings import settings
|
|||||||
from routstr.payment.cost_calculation import CostData, MaxCostData, calculate_cost
|
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)
|
@pytest.fixture(autouse=True)
|
||||||
def mock_fixed_pricing(monkeypatch: pytest.MonkeyPatch) -> None:
|
def mock_fixed_pricing(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
"""Mock settings and price to use fixed pricing."""
|
"""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
|
# Test 1: OpenAI Cache Format
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""OpenAI includes cached_tokens in prompt_tokens, subtract them."""
|
||||||
response = {
|
response = {
|
||||||
"model": "gpt-4",
|
"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 isinstance(result, CostData)
|
||||||
assert result.input_tokens == 1000 # 2000 - 1000
|
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
|
# Test 2: Anthropic Cache Format
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""Anthropic cache tokens are separate (additive) from input_tokens."""
|
||||||
response = {
|
response = {
|
||||||
"model": "claude-3-5-sonnet",
|
"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,
|
"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 isinstance(result, CostData)
|
||||||
assert result.input_tokens == 500
|
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)
|
# Test 3: Invalid Cache (Edge Case)
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""Handle buggy upstream reporting cached > prompt_tokens."""
|
||||||
response = {
|
response = {
|
||||||
"model": "gpt-4",
|
"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
|
# Should not go negative
|
||||||
assert isinstance(result, CostData)
|
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
|
# Test 4: Malformed Token Values
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""Handle non-numeric cache token values."""
|
||||||
response = {
|
response = {
|
||||||
"model": "gpt-4",
|
"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
|
# Both should coerce to 0
|
||||||
assert isinstance(result, CostData)
|
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
|
# Test 5: Anthropic Cache Not Subtracted
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""Anthropic cache fields should NOT be subtracted from input_tokens."""
|
||||||
response = {
|
response = {
|
||||||
"model": "claude-3-5-sonnet",
|
"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
|
"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
|
# Anthropic: input_tokens stays as-is
|
||||||
assert isinstance(result, CostData)
|
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
|
# Test 6: Only Cache Read, No Regular Input
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""Handle response with only cache read tokens."""
|
||||||
response = {
|
response = {
|
||||||
"model": "gpt-4",
|
"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 isinstance(result, CostData)
|
||||||
assert result.input_tokens == 0 # max(0, 0 - 1000)
|
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
|
# Test 7: Only Cache Creation
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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)."""
|
"""Handle response with only cache creation tokens (Anthropic)."""
|
||||||
response = {
|
response = {
|
||||||
"model": "claude-3-5-sonnet",
|
"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,
|
"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 isinstance(result, CostData)
|
||||||
assert result.input_tokens == 500
|
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
|
# Test 8: Both Cache Read and Creation
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""Handle response with both cache read and creation."""
|
||||||
response = {
|
response = {
|
||||||
"model": "claude-3-5-sonnet",
|
"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,
|
"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 isinstance(result, CostData)
|
||||||
assert result.input_tokens == 300
|
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
|
# Test 9: Token Field Fallback
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""Verify fallback order for token extraction."""
|
||||||
# When prompt_tokens is not present, fall back to input_tokens
|
# When prompt_tokens is not present, fall back to input_tokens
|
||||||
response = {
|
response = {
|
||||||
@@ -243,7 +238,7 @@ async def test_token_field_fallback_order(mock_session: AsyncMock, mock_fixed_pr
|
|||||||
"completion_tokens": 50,
|
"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 isinstance(result, CostData)
|
||||||
assert result.input_tokens == 250
|
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
|
# Test 10: Float Token Values
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""Handle float token values by converting to int."""
|
||||||
response = {
|
response = {
|
||||||
"model": "gpt-4",
|
"model": "gpt-4",
|
||||||
"usage": {
|
"usage": {
|
||||||
"prompt_tokens": 100.7, # Float
|
"prompt_tokens": 100.7, # Float
|
||||||
"completion_tokens": 50.3, # 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 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.output_tokens == 50 # Floored
|
||||||
assert result.cache_read_input_tokens == 25 # 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
|
# Test 11: Boolean Cache Tokens
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""Handle boolean cache token values by coercing to zero."""
|
||||||
response = {
|
response = {
|
||||||
"model": "gpt-4",
|
"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
|
"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 isinstance(result, CostData)
|
||||||
assert result.cache_read_input_tokens == 0 # Boolean coerced to 0
|
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
|
# Test 12: Zero Cache Tokens
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""Handle explicit zero cache tokens."""
|
||||||
response = {
|
response = {
|
||||||
"model": "gpt-4",
|
"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 isinstance(result, CostData)
|
||||||
assert result.cache_read_input_tokens == 0
|
assert result.cache_read_input_tokens == 0
|
||||||
assert result.input_tokens == 100
|
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
|
# Test 13: Missing Usage Block
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""When usage is missing, return MaxCostData with zero tokens."""
|
||||||
response = {"model": "gpt-4", "choices": [{"message": {"content": "test"}}]}
|
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 isinstance(result, MaxCostData)
|
||||||
assert result.input_tokens == 0
|
assert result.input_tokens == 0
|
||||||
@@ -335,10 +423,10 @@ async def test_missing_usage_block(mock_session: AsyncMock, mock_fixed_pricing:
|
|||||||
# Test 14: Null Usage Block
|
# Test 14: Null Usage Block
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
@pytest.mark.asyncio
|
@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."""
|
"""When usage is null, return MaxCostData with zero tokens."""
|
||||||
response = {"model": "gpt-4", "usage": None}
|
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 isinstance(result, MaxCostData)
|
||||||
assert result.input_tokens == 0
|
assert result.input_tokens == 0
|
||||||
|
|||||||
@@ -500,7 +500,7 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
|
|||||||
captured_cost_call: dict[str, Any] = {}
|
captured_cost_call: dict[str, Any] = {}
|
||||||
|
|
||||||
async def fake_adjust(
|
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:
|
) -> dict:
|
||||||
captured_cost_call["combined_data"] = combined_data
|
captured_cost_call["combined_data"] = combined_data
|
||||||
captured_cost_call["max_cost"] = max_cost
|
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] = {}
|
captured: dict[str, Any] = {}
|
||||||
|
|
||||||
async def fake_adjust(
|
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:
|
) -> dict:
|
||||||
captured["combined_data"] = combined_data
|
captured["combined_data"] = combined_data
|
||||||
return fake_cost
|
return fake_cost
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ comment ever reaches the client. That invariant is exactly what the buggy
|
|||||||
|
|
||||||
import json
|
import json
|
||||||
from collections.abc import AsyncGenerator
|
from collections.abc import AsyncGenerator
|
||||||
|
from typing import cast
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -107,6 +108,76 @@ def _assert_clean(out: list[bytes]) -> list[dict]:
|
|||||||
return objs
|
return objs
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_deepseek_usage_chunk_is_normalized_before_billing() -> None:
|
||||||
|
"""DeepSeek stream trailers keep raw fields and add canonical billing fields."""
|
||||||
|
usage = {
|
||||||
|
"prompt_tokens": 10000,
|
||||||
|
"completion_tokens": 500,
|
||||||
|
"total_tokens": 10500,
|
||||||
|
"prompt_cache_hit_tokens": 9000,
|
||||||
|
"prompt_cache_miss_tokens": 1000,
|
||||||
|
}
|
||||||
|
chunks = [
|
||||||
|
b'data: {"id":"ds","model":"deepseek-chat","choices":[{"delta":{"content":"ok"}}]}\n\n',
|
||||||
|
b"data: "
|
||||||
|
+ json.dumps(
|
||||||
|
{
|
||||||
|
"id": "ds",
|
||||||
|
"model": "deepseek-chat",
|
||||||
|
"choices": [],
|
||||||
|
"usage": usage,
|
||||||
|
}
|
||||||
|
).encode()
|
||||||
|
+ b"\n\n",
|
||||||
|
b"data: [DONE]\n\n",
|
||||||
|
]
|
||||||
|
|
||||||
|
await _drive(chunks)
|
||||||
|
|
||||||
|
adjust_mock = cast(AsyncMock, base.adjust_payment_for_tokens)
|
||||||
|
adjustment_input = adjust_mock.call_args.args[1]
|
||||||
|
billed_usage = adjustment_input["usage"]
|
||||||
|
assert billed_usage["prompt_tokens"] == 10000
|
||||||
|
assert billed_usage["prompt_cache_hit_tokens"] == 9000
|
||||||
|
assert billed_usage["prompt_cache_miss_tokens"] == 1000
|
||||||
|
assert billed_usage["input_tokens"] == 1000
|
||||||
|
assert billed_usage["output_tokens"] == 500
|
||||||
|
assert billed_usage["cache_read_input_tokens"] == 9000
|
||||||
|
assert billed_usage["cache_creation_input_tokens"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_fold_cache_tokens_does_not_double_count_inclusive_prompt_tokens() -> None:
|
||||||
|
"""Visible usage mutation must not inflate OpenAI-compatible prompt totals."""
|
||||||
|
usage = {
|
||||||
|
"prompt_tokens": 10000,
|
||||||
|
"completion_tokens": 500,
|
||||||
|
"cache_read_input_tokens": 9000,
|
||||||
|
"prompt_cache_hit_tokens": 9000,
|
||||||
|
}
|
||||||
|
|
||||||
|
BaseUpstreamProvider._fold_cache_into_input_tokens(usage)
|
||||||
|
|
||||||
|
assert usage["prompt_tokens"] == 10000
|
||||||
|
assert usage["cache_read_input_tokens"] == 9000
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_fold_cache_tokens_still_rolls_up_anthropic_input_tokens() -> None:
|
||||||
|
"""Anthropic native input_tokens excludes cache and still needs rollup."""
|
||||||
|
usage = {
|
||||||
|
"input_tokens": 1000,
|
||||||
|
"output_tokens": 500,
|
||||||
|
"cache_read_input_tokens": 9000,
|
||||||
|
"cache_creation_input_tokens": 200,
|
||||||
|
}
|
||||||
|
|
||||||
|
BaseUpstreamProvider._fold_cache_into_input_tokens(usage)
|
||||||
|
|
||||||
|
assert usage["input_tokens"] == 10200
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_openai_style_plain_stream() -> None:
|
async def test_openai_style_plain_stream() -> None:
|
||||||
"""OpenAI / Groq / Fireworks / xAI / Perplexity: plain data + [DONE]."""
|
"""OpenAI / Groq / Fireworks / xAI / Perplexity: plain data + [DONE]."""
|
||||||
|
|||||||
@@ -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
|
||||||
Reference in New Issue
Block a user