mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fix: bill cached input tokens at their real rates across vendor dialects
Cached prompt tokens were billed at the full input rate whenever a vendor's usage dialect or cache pricing was unknown, overcharging DeepSeek topups ~5-10x on agentic workloads (hits are 10x cheaper upstream) and silently mispricing OpenAI cached reads and Anthropic cache writes the same way. Two root causes, two fixes: - Usage dialects: DeepSeek reports prompt_cache_hit_tokens / prompt_cache_miss_tokens, which billing never parsed. Usage normalization now lives in payment/usage.py as a union parser over the known, non-colliding dialects (OpenAI prompt_tokens_details, Anthropic additive cache fields, DeepSeek hit/miss), producing one canonical NormalizedUsage. Providers expose it as an overridable BaseUpstreamProvider.normalize_usage hook — the escape hatch for future vendors whose fields genuinely conflict — and every settlement call site passes the provider's result through, so calculate_cost holds no vendor knowledge of its own. - Cache rates: the OpenRouter model feed omits input_cache_read/-write for most DeepSeek models (and e.g. openai/gpt-4o), so billing fell back to the full input rate. Missing rates are now backfilled from litellm's bundled cost map before the provider fee is applied; the input-rate fallback remains only as the documented last resort. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
a16bc1220c
commit
eaf74edbba
+11
-2
@@ -20,6 +20,7 @@ from .payment.cost_calculation import (
|
|||||||
MaxCostData,
|
MaxCostData,
|
||||||
calculate_cost,
|
calculate_cost,
|
||||||
)
|
)
|
||||||
|
from .payment.usage import NormalizedUsage
|
||||||
from .wallet import credit_balance, deserialize_token_from_string
|
from .wallet import credit_balance, deserialize_token_from_string
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
@@ -675,12 +676,20 @@ 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,
|
||||||
|
usage: NormalizedUsage | None = None,
|
||||||
) -> 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.
|
||||||
|
|
||||||
|
``usage`` carries the upstream provider's normalized token usage (its
|
||||||
|
``normalize_usage`` hook); when omitted, the response's usage object is
|
||||||
|
normalized with the default union parser.
|
||||||
"""
|
"""
|
||||||
billing_key = await get_billing_key(key, session)
|
billing_key = await get_billing_key(key, session)
|
||||||
model = response_data.get("model", "unknown")
|
model = response_data.get("model", "unknown")
|
||||||
@@ -763,7 +772,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, session, usage=usage):
|
||||||
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)",
|
||||||
|
|||||||
@@ -6,6 +6,15 @@ from ..core import get_logger
|
|||||||
from ..core.db import AsyncSession
|
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 NormalizedUsage, normalize_usage, parse_token_count
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"CostData",
|
||||||
|
"CostDataError",
|
||||||
|
"MaxCostData",
|
||||||
|
"calculate_cost",
|
||||||
|
"parse_token_count",
|
||||||
|
]
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
@@ -34,13 +43,20 @@ 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,
|
||||||
|
session: AsyncSession,
|
||||||
|
usage: NormalizedUsage | None = None,
|
||||||
) -> 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.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
response_data: Response data containing usage information
|
response_data: Response data containing usage information
|
||||||
max_cost: Maximum cost in millisats
|
max_cost: Maximum cost in millisats
|
||||||
|
usage: Pre-normalized usage from the upstream provider's
|
||||||
|
``normalize_usage`` hook. When omitted, the response's usage
|
||||||
|
object is normalized with the default union parser. This function
|
||||||
|
holds no vendor-dialect knowledge of its own.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Cost data or error information
|
Cost data or error information
|
||||||
@@ -54,8 +70,10 @@ async def calculate_cost(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
# Check for usage data
|
if usage is None:
|
||||||
if "usage" not in response_data or response_data["usage"] is None:
|
usage = normalize_usage(response_data.get("usage"))
|
||||||
|
|
||||||
|
if usage is None:
|
||||||
logger.warning(
|
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 +102,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 +218,6 @@ async def calculate_cost(
|
|||||||
# ============================================================================
|
# ============================================================================
|
||||||
|
|
||||||
|
|
||||||
def parse_token_count(value: object) -> int:
|
|
||||||
"""Parse a token count from various formats (int, float, str, bool)."""
|
|
||||||
if isinstance(value, bool):
|
|
||||||
return 0
|
|
||||||
if isinstance(value, int):
|
|
||||||
return max(0, value)
|
|
||||||
if isinstance(value, float):
|
|
||||||
return max(0, int(value))
|
|
||||||
if isinstance(value, str):
|
|
||||||
try:
|
|
||||||
return max(0, int(float(value)))
|
|
||||||
except ValueError:
|
|
||||||
return 0
|
|
||||||
return 0
|
|
||||||
|
|
||||||
|
|
||||||
def _coerce_usd(value: object) -> float:
|
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 +230,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.
|
||||||
|
|
||||||
|
|||||||
@@ -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,92 @@
|
|||||||
|
"""Vendor-agnostic normalization of upstream usage objects.
|
||||||
|
|
||||||
|
Upstream providers report token usage in vendor dialects that differ in field
|
||||||
|
names and in whether cached tokens are included in the input count:
|
||||||
|
|
||||||
|
* OpenAI: ``prompt_tokens_details.cached_tokens``, included in ``prompt_tokens``
|
||||||
|
* Anthropic: ``cache_read_input_tokens`` / ``cache_creation_input_tokens``,
|
||||||
|
additive to (not included in) ``input_tokens``
|
||||||
|
* DeepSeek: ``prompt_cache_hit_tokens`` / ``prompt_cache_miss_tokens``, with
|
||||||
|
``prompt_tokens = hit + miss``
|
||||||
|
|
||||||
|
``normalize_usage`` maps all of them onto one canonical ``NormalizedUsage``
|
||||||
|
shape so billing code needs no vendor knowledge. The known dialects' field
|
||||||
|
names do not collide, so a single union parser is safe; vendors whose fields
|
||||||
|
would genuinely conflict override ``BaseUpstreamProvider.normalize_usage``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from pydantic.v1 import BaseModel
|
||||||
|
|
||||||
|
|
||||||
|
class NormalizedUsage(BaseModel):
|
||||||
|
"""Canonical token usage: input_tokens never includes cached tokens."""
|
||||||
|
|
||||||
|
input_tokens: int = 0
|
||||||
|
output_tokens: int = 0
|
||||||
|
cache_read_tokens: int = 0
|
||||||
|
cache_write_tokens: int = 0
|
||||||
|
|
||||||
|
|
||||||
|
def parse_token_count(value: object) -> int:
|
||||||
|
"""Parse a token count from various formats (int, float, str, bool)."""
|
||||||
|
if isinstance(value, bool):
|
||||||
|
return 0
|
||||||
|
if isinstance(value, int):
|
||||||
|
return max(0, value)
|
||||||
|
if isinstance(value, float):
|
||||||
|
return max(0, int(value))
|
||||||
|
if isinstance(value, str):
|
||||||
|
try:
|
||||||
|
return max(0, int(float(value)))
|
||||||
|
except ValueError:
|
||||||
|
return 0
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def _first_token_count(usage_data: dict, *fields: str) -> int:
|
||||||
|
"""Return the first positive token count among the given fields."""
|
||||||
|
for field in fields:
|
||||||
|
value = parse_token_count(usage_data.get(field, 0))
|
||||||
|
if value > 0:
|
||||||
|
return value
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_usage(usage_data: object) -> NormalizedUsage | None:
|
||||||
|
"""Map a vendor usage dict onto the canonical shape, or None if absent.
|
||||||
|
|
||||||
|
Cached tokens are subtracted from the input count exactly once, only for
|
||||||
|
dialects that include them in it (OpenAI, DeepSeek). Precedence between
|
||||||
|
cache fields: Anthropic explicit > OpenAI details > DeepSeek hit/miss.
|
||||||
|
"""
|
||||||
|
if not isinstance(usage_data, dict):
|
||||||
|
return None
|
||||||
|
|
||||||
|
input_tokens = _first_token_count(usage_data, "prompt_tokens", "input_tokens")
|
||||||
|
output_tokens = _first_token_count(
|
||||||
|
usage_data, "completion_tokens", "output_tokens"
|
||||||
|
)
|
||||||
|
cache_write = parse_token_count(usage_data.get("cache_creation_input_tokens", 0))
|
||||||
|
|
||||||
|
# Anthropic: cache reads are additive, input_tokens stays untouched
|
||||||
|
cache_read = parse_token_count(usage_data.get("cache_read_input_tokens", 0))
|
||||||
|
|
||||||
|
if not cache_read:
|
||||||
|
# OpenAI: cached tokens are included in prompt_tokens
|
||||||
|
prompt_details = usage_data.get("prompt_tokens_details")
|
||||||
|
if isinstance(prompt_details, dict):
|
||||||
|
cache_read = parse_token_count(prompt_details.get("cached_tokens", 0))
|
||||||
|
# DeepSeek: prompt_tokens = prompt_cache_hit_tokens + prompt_cache_miss_tokens
|
||||||
|
if not cache_read:
|
||||||
|
cache_read = parse_token_count(
|
||||||
|
usage_data.get("prompt_cache_hit_tokens", 0)
|
||||||
|
)
|
||||||
|
if cache_read:
|
||||||
|
input_tokens = max(0, input_tokens - cache_read)
|
||||||
|
|
||||||
|
return NormalizedUsage(
|
||||||
|
input_tokens=input_tokens,
|
||||||
|
output_tokens=output_tokens,
|
||||||
|
cache_read_tokens=cache_read,
|
||||||
|
cache_write_tokens=cache_write,
|
||||||
|
)
|
||||||
@@ -35,9 +35,11 @@ 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 NormalizedUsage, 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 .count_tokens import count_tokens_locally
|
from .count_tokens import count_tokens_locally
|
||||||
@@ -101,6 +103,17 @@ class BaseUpstreamProvider:
|
|||||||
self._models_cache = []
|
self._models_cache = []
|
||||||
self._models_by_id = {}
|
self._models_by_id = {}
|
||||||
|
|
||||||
|
def normalize_usage(self, usage_data: object) -> NormalizedUsage | None:
|
||||||
|
"""Map this provider's usage dialect onto the canonical shape.
|
||||||
|
|
||||||
|
The default union parser covers the known dialects (OpenAI, Anthropic,
|
||||||
|
DeepSeek), whose field names do not collide. Override in a subclass
|
||||||
|
when a vendor's usage fields would conflict with another dialect or
|
||||||
|
need bespoke interpretation; billing code consumes only the canonical
|
||||||
|
result and holds no vendor knowledge.
|
||||||
|
"""
|
||||||
|
return normalize_usage(usage_data)
|
||||||
|
|
||||||
def get_litellm_provider_prefix(self) -> str:
|
def get_litellm_provider_prefix(self) -> str:
|
||||||
"""Resolve the litellm provider prefix for this provider instance.
|
"""Resolve the litellm provider prefix for this provider instance.
|
||||||
|
|
||||||
@@ -880,6 +893,9 @@ class BaseUpstreamProvider:
|
|||||||
adjustment_input,
|
adjustment_input,
|
||||||
session,
|
session,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
|
usage=self.normalize_usage(
|
||||||
|
adjustment_input.get("usage")
|
||||||
|
),
|
||||||
)
|
)
|
||||||
usage_finalized = True
|
usage_finalized = True
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -1018,7 +1034,11 @@ 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,
|
||||||
|
usage=self.normalize_usage(response_json.get("usage")),
|
||||||
)
|
)
|
||||||
|
|
||||||
await session.refresh(key)
|
await session.refresh(key)
|
||||||
@@ -1273,6 +1293,9 @@ class BaseUpstreamProvider:
|
|||||||
adjustment_input,
|
adjustment_input,
|
||||||
session,
|
session,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
|
usage=self.normalize_usage(
|
||||||
|
adjustment_input.get("usage")
|
||||||
|
),
|
||||||
)
|
)
|
||||||
usage_finalized = True
|
usage_finalized = True
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -1436,7 +1459,11 @@ 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,
|
||||||
|
usage=self.normalize_usage(response_json.get("usage")),
|
||||||
)
|
)
|
||||||
|
|
||||||
await session.refresh(key)
|
await session.refresh(key)
|
||||||
@@ -1631,7 +1658,11 @@ 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=self.normalize_usage(fallback.get("usage")),
|
||||||
)
|
)
|
||||||
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()
|
||||||
@@ -1783,6 +1814,9 @@ class BaseUpstreamProvider:
|
|||||||
combined_data,
|
combined_data,
|
||||||
new_session,
|
new_session,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
|
usage=self.normalize_usage(
|
||||||
|
combined_data.get("usage")
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.inject_cost_metadata(
|
self.inject_cost_metadata(
|
||||||
@@ -1850,7 +1884,11 @@ 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,
|
||||||
|
usage=self.normalize_usage(response_json.get("usage")),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.inject_cost_metadata(response_json, cost_data, key)
|
self.inject_cost_metadata(response_json, cost_data, key)
|
||||||
@@ -1952,7 +1990,11 @@ 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,
|
||||||
|
usage=self.normalize_usage(response_json.get("usage")),
|
||||||
)
|
)
|
||||||
self.inject_cost_metadata(response_json, cost_data, key)
|
self.inject_cost_metadata(response_json, cost_data, key)
|
||||||
|
|
||||||
@@ -2094,6 +2136,7 @@ class BaseUpstreamProvider:
|
|||||||
fallback,
|
fallback,
|
||||||
new_session,
|
new_session,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
|
usage=self.normalize_usage(fallback.get("usage")),
|
||||||
)
|
)
|
||||||
usage_finalized = True
|
usage_finalized = True
|
||||||
return (
|
return (
|
||||||
@@ -2167,6 +2210,9 @@ class BaseUpstreamProvider:
|
|||||||
combined_data,
|
combined_data,
|
||||||
new_session,
|
new_session,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
|
usage=self.normalize_usage(
|
||||||
|
combined_data.get("usage")
|
||||||
|
),
|
||||||
)
|
)
|
||||||
self.inject_cost_metadata(
|
self.inject_cost_metadata(
|
||||||
combined_data, cost_data, fresh_key
|
combined_data, cost_data, fresh_key
|
||||||
@@ -3026,7 +3072,12 @@ class BaseUpstreamProvider:
|
|||||||
)
|
)
|
||||||
|
|
||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
match await calculate_cost(response_data, max_cost_for_model, session):
|
match await calculate_cost(
|
||||||
|
response_data,
|
||||||
|
max_cost_for_model,
|
||||||
|
session,
|
||||||
|
usage=self.normalize_usage(response_data.get("usage")),
|
||||||
|
):
|
||||||
case MaxCostData() as cost:
|
case MaxCostData() as cost:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Using max cost pricing",
|
"Using max cost pricing",
|
||||||
@@ -4562,14 +4613,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,241 @@
|
|||||||
|
"""Tests for cache-aware pricing of cached input tokens.
|
||||||
|
|
||||||
|
Specifies two things:
|
||||||
|
|
||||||
|
1. ``backfill_cache_pricing`` — when the OpenRouter model feed omits cache
|
||||||
|
rates (it does for most DeepSeek models and e.g. openai/gpt-4o), they are
|
||||||
|
filled from litellm's bundled cost map instead of silently billing cache
|
||||||
|
reads at the full input rate. Existing OpenRouter values are never
|
||||||
|
overwritten, and provider fees apply to backfilled rates like any other.
|
||||||
|
2. ``calculate_cost`` — cached tokens are billed at the cache rates from the
|
||||||
|
model's sats_pricing; the full input rate remains only as the documented
|
||||||
|
last resort when no cache rate could be resolved anywhere.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
from unittest.mock import AsyncMock, Mock, patch
|
||||||
|
|
||||||
|
import litellm
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
|
||||||
|
os.environ.setdefault("UPSTREAM_API_KEY", "test")
|
||||||
|
os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to")
|
||||||
|
|
||||||
|
from routstr.core.settings import settings
|
||||||
|
from routstr.payment.cost_calculation import CostData, calculate_cost
|
||||||
|
from routstr.payment.models import (
|
||||||
|
Architecture,
|
||||||
|
Model,
|
||||||
|
Pricing,
|
||||||
|
backfill_cache_pricing,
|
||||||
|
)
|
||||||
|
from routstr.upstream import GenericUpstreamProvider
|
||||||
|
|
||||||
|
|
||||||
|
def _make_model(model_id: str, pricing: Pricing) -> Model:
|
||||||
|
return Model(
|
||||||
|
id=model_id,
|
||||||
|
name=model_id,
|
||||||
|
created=0,
|
||||||
|
description="",
|
||||||
|
context_length=64000,
|
||||||
|
architecture=Architecture(
|
||||||
|
modality="text->text",
|
||||||
|
input_modalities=["text"],
|
||||||
|
output_modalities=["text"],
|
||||||
|
tokenizer="Other",
|
||||||
|
instruct_type=None,
|
||||||
|
),
|
||||||
|
pricing=pricing,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================================
|
||||||
|
# backfill_cache_pricing — litellm as fallback source for missing cache rates
|
||||||
|
# ============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
def test_backfill_deepseek_cache_read_from_litellm() -> None:
|
||||||
|
"""deepseek/deepseek-chat has no input_cache_read on OpenRouter; litellm
|
||||||
|
knows the real rate (10x cheaper than input)."""
|
||||||
|
pricing = Pricing(prompt=2.8e-07, completion=4.2e-07)
|
||||||
|
|
||||||
|
result = backfill_cache_pricing("deepseek/deepseek-chat", pricing)
|
||||||
|
|
||||||
|
expected = litellm.model_cost["deepseek/deepseek-chat"][
|
||||||
|
"cache_read_input_token_cost"
|
||||||
|
]
|
||||||
|
assert result.input_cache_read == expected
|
||||||
|
assert result.input_cache_read < pricing.prompt # sanity: it's a discount
|
||||||
|
|
||||||
|
|
||||||
|
def test_backfill_strips_vendor_prefix_for_litellm_lookup() -> None:
|
||||||
|
"""OpenRouter ids are vendor-prefixed (openai/gpt-4o); litellm keys most
|
||||||
|
non-DeepSeek models without the prefix (gpt-4o)."""
|
||||||
|
pricing = Pricing(prompt=2.5e-06, completion=1e-05)
|
||||||
|
|
||||||
|
result = backfill_cache_pricing("openai/gpt-4o", pricing)
|
||||||
|
|
||||||
|
expected = litellm.model_cost["gpt-4o"]["cache_read_input_token_cost"]
|
||||||
|
assert result.input_cache_read == expected
|
||||||
|
|
||||||
|
|
||||||
|
def test_backfill_fills_cache_write_rate() -> None:
|
||||||
|
"""Anthropic cache writes cost more than input (1.25x); billing them at
|
||||||
|
the input rate undercharges. litellm carries the write rate."""
|
||||||
|
pricing = Pricing(prompt=3e-06, completion=1.5e-05)
|
||||||
|
|
||||||
|
result = backfill_cache_pricing("anthropic/claude-sonnet-4-5", pricing)
|
||||||
|
|
||||||
|
expected = litellm.model_cost["claude-sonnet-4-5"][
|
||||||
|
"cache_creation_input_token_cost"
|
||||||
|
]
|
||||||
|
assert result.input_cache_write == expected
|
||||||
|
assert result.input_cache_write > pricing.prompt # sanity: write premium
|
||||||
|
|
||||||
|
|
||||||
|
def test_backfill_never_overwrites_openrouter_rates() -> None:
|
||||||
|
"""When OpenRouter provides a cache rate, it is authoritative."""
|
||||||
|
pricing = Pricing(
|
||||||
|
prompt=2.1e-07, completion=7.9e-07, input_cache_read=1.3e-07
|
||||||
|
)
|
||||||
|
|
||||||
|
result = backfill_cache_pricing("deepseek/deepseek-chat", pricing)
|
||||||
|
|
||||||
|
assert result.input_cache_read == 1.3e-07
|
||||||
|
|
||||||
|
|
||||||
|
def test_backfill_unknown_model_unchanged() -> None:
|
||||||
|
"""Models litellm doesn't know stay untouched (last-resort fallback to
|
||||||
|
the input rate happens later, at billing time)."""
|
||||||
|
pricing = Pricing(prompt=1e-06, completion=2e-06)
|
||||||
|
|
||||||
|
result = backfill_cache_pricing("artificial-dumbness/dumb-1", pricing)
|
||||||
|
|
||||||
|
assert result.input_cache_read == 0.0
|
||||||
|
assert result.input_cache_write == 0.0
|
||||||
|
|
||||||
|
|
||||||
|
def test_provider_fee_applies_to_backfilled_cache_rates() -> None:
|
||||||
|
"""Backfill happens before the provider fee, so cache rates carry the
|
||||||
|
same markup as every other price component."""
|
||||||
|
provider = GenericUpstreamProvider(
|
||||||
|
base_url="http://upstream.example", provider_fee=2.0
|
||||||
|
)
|
||||||
|
model = _make_model(
|
||||||
|
"deepseek/deepseek-chat", Pricing(prompt=2.8e-07, completion=4.2e-07)
|
||||||
|
)
|
||||||
|
|
||||||
|
adjusted = provider._apply_provider_fee_to_model(model)
|
||||||
|
|
||||||
|
litellm_read = litellm.model_cost["deepseek/deepseek-chat"][
|
||||||
|
"cache_read_input_token_cost"
|
||||||
|
]
|
||||||
|
assert adjusted.pricing.input_cache_read == pytest.approx(litellm_read * 2.0)
|
||||||
|
assert adjusted.pricing.prompt == pytest.approx(2.8e-07 * 2.0)
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================================
|
||||||
|
# calculate_cost — cached tokens billed at cache rates
|
||||||
|
# ============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def patch_sats_usd_price() -> None: # type: ignore[misc]
|
||||||
|
with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-5):
|
||||||
|
yield
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def model_pricing(monkeypatch: pytest.MonkeyPatch) -> Mock:
|
||||||
|
"""Model-based pricing: 1 msat per input token, 2 per output token,
|
||||||
|
0.1 per cache-read token, 1.25 per cache-write token."""
|
||||||
|
monkeypatch.setattr(settings, "fixed_pricing", False)
|
||||||
|
model = Mock()
|
||||||
|
model.sats_pricing = Pricing(
|
||||||
|
prompt=0.001,
|
||||||
|
completion=0.002,
|
||||||
|
input_cache_read=0.0001,
|
||||||
|
input_cache_write=0.00125,
|
||||||
|
)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_deepseek_cache_hits_billed_at_cache_rate(model_pricing: Mock) -> None:
|
||||||
|
"""The reported overcharge scenario: a 10k-token prompt with 90% cache
|
||||||
|
hits costs 2900 msats at honest rates, not the 11000 msats that billing
|
||||||
|
every prompt token at the full input rate would charge."""
|
||||||
|
response = {
|
||||||
|
"model": "deepseek-chat",
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": 10000,
|
||||||
|
"completion_tokens": 500,
|
||||||
|
"prompt_cache_hit_tokens": 9000,
|
||||||
|
"prompt_cache_miss_tokens": 1000,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
with patch("routstr.proxy.get_model_instance", return_value=model_pricing):
|
||||||
|
result = await calculate_cost(response, max_cost=100000, session=AsyncMock())
|
||||||
|
|
||||||
|
assert isinstance(result, CostData)
|
||||||
|
# 1000 input @ 1 msat + 9000 cache reads @ 0.1 msat + 500 output @ 2 msat
|
||||||
|
assert result.input_msats == 1000
|
||||||
|
assert result.cache_read_msats == 900
|
||||||
|
assert result.output_msats == 1000
|
||||||
|
assert result.total_msats == 2900
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_anthropic_cache_write_billed_at_write_rate(model_pricing: Mock) -> None:
|
||||||
|
"""Cache writes carry their premium rate (1.25x input here), instead of
|
||||||
|
being silently billed at the plain input rate."""
|
||||||
|
response = {
|
||||||
|
"model": "claude-sonnet-4-5",
|
||||||
|
"usage": {
|
||||||
|
"input_tokens": 300,
|
||||||
|
"output_tokens": 100,
|
||||||
|
"cache_read_input_tokens": 500,
|
||||||
|
"cache_creation_input_tokens": 2000,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
with patch("routstr.proxy.get_model_instance", return_value=model_pricing):
|
||||||
|
result = await calculate_cost(response, max_cost=100000, session=AsyncMock())
|
||||||
|
|
||||||
|
assert isinstance(result, CostData)
|
||||||
|
# 300 @ 1 + 500 @ 0.1 + 2000 @ 1.25 + 100 @ 2
|
||||||
|
assert result.cache_read_msats == 50
|
||||||
|
assert result.cache_creation_msats == 2500
|
||||||
|
assert result.total_msats == 3050
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_missing_cache_rate_falls_back_to_input_rate(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""Documented last resort: when no cache rate could be resolved anywhere
|
||||||
|
(OpenRouter and litellm both silent), cache reads bill at the input rate —
|
||||||
|
never cheaper, never free."""
|
||||||
|
monkeypatch.setattr(settings, "fixed_pricing", False)
|
||||||
|
model = Mock()
|
||||||
|
model.sats_pricing = Pricing(prompt=0.001, completion=0.002)
|
||||||
|
|
||||||
|
response = {
|
||||||
|
"model": "dumb-1",
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": 10000,
|
||||||
|
"completion_tokens": 500,
|
||||||
|
"prompt_cache_hit_tokens": 9000,
|
||||||
|
"prompt_cache_miss_tokens": 1000,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
with patch("routstr.proxy.get_model_instance", return_value=model):
|
||||||
|
result = await calculate_cost(response, max_cost=100000, session=AsyncMock())
|
||||||
|
|
||||||
|
assert isinstance(result, CostData)
|
||||||
|
# 1000 @ 1 + 9000 @ 1 (fallback) + 500 @ 2
|
||||||
|
assert result.total_msats == 11000
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
"""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
|
||||||
@@ -316,6 +317,98 @@ async def test_zero_cache_tokens(mock_session: AsyncMock, mock_fixed_pricing: No
|
|||||||
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(mock_session: AsyncMock) -> None:
|
||||||
|
"""DeepSeek cache hits are extracted and removed from regular input.
|
||||||
|
|
||||||
|
Payload shape verbatim from the DeepSeek API reference (usage object).
|
||||||
|
"""
|
||||||
|
response = {
|
||||||
|
"model": "deepseek-chat",
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": 10000, # = hit + miss
|
||||||
|
"completion_tokens": 500,
|
||||||
|
"total_tokens": 10500,
|
||||||
|
"prompt_cache_hit_tokens": 9000,
|
||||||
|
"prompt_cache_miss_tokens": 1000,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||||
|
|
||||||
|
assert isinstance(result, CostData)
|
||||||
|
assert result.input_tokens == 1000 # only the cache misses
|
||||||
|
assert result.cache_read_input_tokens == 9000
|
||||||
|
assert result.output_tokens == 500
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_deepseek_all_tokens_cached(mock_session: AsyncMock) -> None:
|
||||||
|
"""A fully cached DeepSeek prompt bills zero regular input tokens."""
|
||||||
|
response = {
|
||||||
|
"model": "deepseek-chat",
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": 5000,
|
||||||
|
"completion_tokens": 100,
|
||||||
|
"prompt_cache_hit_tokens": 5000,
|
||||||
|
"prompt_cache_miss_tokens": 0,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||||
|
|
||||||
|
assert isinstance(result, CostData)
|
||||||
|
assert result.input_tokens == 0
|
||||||
|
assert result.cache_read_input_tokens == 5000
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_dialect_precedence_never_double_subtracts(mock_session: AsyncMock) -> None:
|
||||||
|
"""If a vendor emits both OpenAI-style and DeepSeek-style cache fields for
|
||||||
|
the same cached tokens, they are counted once, not subtracted twice."""
|
||||||
|
response = {
|
||||||
|
"model": "deepseek-chat",
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": 10000,
|
||||||
|
"completion_tokens": 500,
|
||||||
|
"prompt_tokens_details": {"cached_tokens": 9000},
|
||||||
|
"prompt_cache_hit_tokens": 9000,
|
||||||
|
"prompt_cache_miss_tokens": 1000,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||||
|
|
||||||
|
assert isinstance(result, CostData)
|
||||||
|
assert result.input_tokens == 1000 # 10000 - 9000, applied exactly once
|
||||||
|
assert result.cache_read_input_tokens == 9000
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_deepseek_malformed_hit_tokens_coerce_to_zero(mock_session: AsyncMock) -> None:
|
||||||
|
"""Malformed DeepSeek cache fields degrade to billing all input at full
|
||||||
|
rate instead of crashing or going negative."""
|
||||||
|
response = {
|
||||||
|
"model": "deepseek-chat",
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": 1000,
|
||||||
|
"completion_tokens": 50,
|
||||||
|
"prompt_cache_hit_tokens": "garbage",
|
||||||
|
"prompt_cache_miss_tokens": -5,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||||
|
|
||||||
|
assert isinstance(result, CostData)
|
||||||
|
assert result.input_tokens == 1000
|
||||||
|
assert result.cache_read_input_tokens == 0
|
||||||
|
|
||||||
|
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
# Test 13: Missing Usage Block
|
# Test 13: Missing Usage Block
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -0,0 +1,206 @@
|
|||||||
|
"""Tests for vendor-agnostic usage normalization.
|
||||||
|
|
||||||
|
Specifies the seam that keeps vendor usage dialects out of generic billing
|
||||||
|
code: a canonical ``NormalizedUsage`` shape produced by
|
||||||
|
``routstr.payment.usage.normalize_usage`` (union parser for the known,
|
||||||
|
non-colliding dialects) and exposed as an overridable
|
||||||
|
``BaseUpstreamProvider.normalize_usage`` hook for vendors whose fields would
|
||||||
|
genuinely conflict. ``calculate_cost`` accepts a pre-normalized usage and then
|
||||||
|
needs no vendor knowledge of its own.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
|
||||||
|
os.environ.setdefault("UPSTREAM_API_KEY", "test")
|
||||||
|
os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to")
|
||||||
|
|
||||||
|
from routstr.core.settings import settings
|
||||||
|
from routstr.payment.cost_calculation import CostData, calculate_cost
|
||||||
|
from routstr.payment.usage import NormalizedUsage, normalize_usage
|
||||||
|
from routstr.upstream import BaseUpstreamProvider, GenericUpstreamProvider
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def mock_fixed_pricing(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
monkeypatch.setattr(settings, "fixed_pricing", True)
|
||||||
|
monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 0.001)
|
||||||
|
monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 0.001)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def patch_sats_usd_price() -> None: # type: ignore[misc]
|
||||||
|
with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-5):
|
||||||
|
yield
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================================
|
||||||
|
# The union parser: one canonical shape for all known dialects
|
||||||
|
# ============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"usage,expected",
|
||||||
|
[
|
||||||
|
# OpenAI: cached_tokens included in prompt_tokens → subtracted
|
||||||
|
(
|
||||||
|
{
|
||||||
|
"prompt_tokens": 2000,
|
||||||
|
"completion_tokens": 100,
|
||||||
|
"prompt_tokens_details": {"cached_tokens": 800},
|
||||||
|
},
|
||||||
|
NormalizedUsage(
|
||||||
|
input_tokens=1200,
|
||||||
|
output_tokens=100,
|
||||||
|
cache_read_tokens=800,
|
||||||
|
cache_write_tokens=0,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
# DeepSeek: hit/miss fields, prompt_tokens = hit + miss → hit subtracted
|
||||||
|
(
|
||||||
|
{
|
||||||
|
"prompt_tokens": 10000,
|
||||||
|
"completion_tokens": 500,
|
||||||
|
"prompt_cache_hit_tokens": 9000,
|
||||||
|
"prompt_cache_miss_tokens": 1000,
|
||||||
|
},
|
||||||
|
NormalizedUsage(
|
||||||
|
input_tokens=1000,
|
||||||
|
output_tokens=500,
|
||||||
|
cache_read_tokens=9000,
|
||||||
|
cache_write_tokens=0,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
# Anthropic: cache fields additive, input_tokens NOT reduced
|
||||||
|
(
|
||||||
|
{
|
||||||
|
"input_tokens": 300,
|
||||||
|
"output_tokens": 100,
|
||||||
|
"cache_read_input_tokens": 500,
|
||||||
|
"cache_creation_input_tokens": 2000,
|
||||||
|
},
|
||||||
|
NormalizedUsage(
|
||||||
|
input_tokens=300,
|
||||||
|
output_tokens=100,
|
||||||
|
cache_read_tokens=500,
|
||||||
|
cache_write_tokens=2000,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
# Plain OpenAI without caching
|
||||||
|
(
|
||||||
|
{"prompt_tokens": 100, "completion_tokens": 50},
|
||||||
|
NormalizedUsage(input_tokens=100, output_tokens=50),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_normalize_usage_dialects(usage: dict, expected: NormalizedUsage) -> None:
|
||||||
|
"""Each known vendor dialect maps onto the same canonical shape."""
|
||||||
|
assert normalize_usage(usage) == expected
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalize_usage_absent_usage() -> None:
|
||||||
|
"""Missing/invalid usage yields None so callers can bill at max cost."""
|
||||||
|
assert normalize_usage(None) is None
|
||||||
|
assert normalize_usage("not a dict") is None # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalize_usage_never_negative() -> None:
|
||||||
|
"""Buggy upstreams reporting more cached than prompt tokens clamp to 0."""
|
||||||
|
result = normalize_usage(
|
||||||
|
{
|
||||||
|
"prompt_tokens": 100,
|
||||||
|
"completion_tokens": 50,
|
||||||
|
"prompt_cache_hit_tokens": 150,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
assert result is not None
|
||||||
|
assert result.input_tokens == 0
|
||||||
|
assert result.cache_read_tokens == 150
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================================
|
||||||
|
# The provider hook: default delegates to the union parser, subclasses
|
||||||
|
# override for vendors whose fields genuinely conflict
|
||||||
|
# ============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
class ArtificialDumbnessProvider(BaseUpstreamProvider):
|
||||||
|
"""Fictional vendor whose usage dialect collides with nothing we know."""
|
||||||
|
|
||||||
|
provider_type = "artificial-dumbness"
|
||||||
|
|
||||||
|
def normalize_usage(self, usage_data: object) -> NormalizedUsage | None:
|
||||||
|
if not isinstance(usage_data, dict):
|
||||||
|
return None
|
||||||
|
return NormalizedUsage(
|
||||||
|
input_tokens=usage_data.get("dumb_tokens_in", 0),
|
||||||
|
output_tokens=usage_data.get("dumb_tokens_out", 0),
|
||||||
|
cache_read_tokens=usage_data.get("dumb_tokens_reused", 0),
|
||||||
|
cache_write_tokens=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_base_provider_hook_delegates_to_union_parser() -> None:
|
||||||
|
"""Without an override, the provider hook equals the module parser, so
|
||||||
|
DeepSeek-dialect upstreams work through GenericUpstreamProvider with no
|
||||||
|
subclass."""
|
||||||
|
provider = GenericUpstreamProvider(base_url="http://upstream.example")
|
||||||
|
usage = {
|
||||||
|
"prompt_tokens": 10000,
|
||||||
|
"completion_tokens": 500,
|
||||||
|
"prompt_cache_hit_tokens": 9000,
|
||||||
|
"prompt_cache_miss_tokens": 1000,
|
||||||
|
}
|
||||||
|
assert provider.normalize_usage(usage) == normalize_usage(usage)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_provider_override_is_honored_by_calculate_cost() -> None:
|
||||||
|
"""The escape hatch: a vendor-specific override feeds calculate_cost
|
||||||
|
through the `usage` parameter, and generic billing needs no knowledge of
|
||||||
|
the vendor's field names."""
|
||||||
|
provider = ArtificialDumbnessProvider(base_url="http://ad.example", api_key="k")
|
||||||
|
response = {
|
||||||
|
"model": "dumb-1",
|
||||||
|
"usage": {
|
||||||
|
"dumb_tokens_in": 700,
|
||||||
|
"dumb_tokens_out": 60,
|
||||||
|
"dumb_tokens_reused": 4300,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result = await calculate_cost(
|
||||||
|
response,
|
||||||
|
max_cost=100000,
|
||||||
|
session=AsyncMock(),
|
||||||
|
usage=provider.normalize_usage(response["usage"]),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert isinstance(result, CostData)
|
||||||
|
assert result.input_tokens == 700
|
||||||
|
assert result.output_tokens == 60
|
||||||
|
assert result.cache_read_input_tokens == 4300
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_explicit_usage_param_wins_over_response_extraction() -> None:
|
||||||
|
"""When a normalized usage is passed, calculate_cost must not re-derive
|
||||||
|
token counts from the raw response."""
|
||||||
|
response = {
|
||||||
|
"model": "dumb-1",
|
||||||
|
"usage": {"prompt_tokens": 999999, "completion_tokens": 999999},
|
||||||
|
}
|
||||||
|
|
||||||
|
result = await calculate_cost(
|
||||||
|
response,
|
||||||
|
max_cost=100000,
|
||||||
|
session=AsyncMock(),
|
||||||
|
usage=NormalizedUsage(input_tokens=10, output_tokens=5),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert isinstance(result, CostData)
|
||||||
|
assert result.input_tokens == 10
|
||||||
|
assert result.output_tokens == 5
|
||||||
Reference in New Issue
Block a user