diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index bf5c2aa2..647fe6b5 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -3,11 +3,27 @@ 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`` +* 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`` + ``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 @@ -52,37 +68,60 @@ def _first_token_count(usage_data: dict, *fields: str) -> int: 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 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. + 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 - 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)) + cache_read, cache_write = _extract_cache_tokens(usage_data) - # 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) + # ``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, diff --git a/tests/unit/test_cost_calculation_caching.py b/tests/unit/test_cost_calculation_caching.py index 2eb7a7c0..41d85385 100644 --- a/tests/unit/test_cost_calculation_caching.py +++ b/tests/unit/test_cost_calculation_caching.py @@ -256,13 +256,14 @@ async def test_float_token_values_coerced_to_int(mock_fixed_pricing: None) -> No "usage": { "prompt_tokens": 100.7, # Float "completion_tokens": 50.3, # Float - "cache_read_input_tokens": 25.9, # Float + "prompt_tokens_details": {"cached_tokens": 25.9}, # Float } } result = await calculate_cost(response, max_cost=100000) assert isinstance(result, CostData) - assert result.input_tokens == 100 # Floored + # cached_tokens are part of prompt_tokens (OpenAI dialect) → subtracted: 100 - 25 + assert result.input_tokens == 75 # Floored assert result.output_tokens == 50 # Floored assert result.cache_read_input_tokens == 25 # Floored diff --git a/tests/unit/test_usage_normalization.py b/tests/unit/test_usage_normalization.py index 0412d55e..f819c2cb 100644 --- a/tests/unit/test_usage_normalization.py +++ b/tests/unit/test_usage_normalization.py @@ -94,6 +94,45 @@ def patch_sats_usd_price() -> None: # type: ignore[misc] {"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: