diff --git a/docs/tinfoil-direct-integration.md b/docs/tinfoil-direct-integration.md index d41228ec..171adf43 100644 --- a/docs/tinfoil-direct-integration.md +++ b/docs/tinfoil-direct-integration.md @@ -260,7 +260,13 @@ During local PPQ testing, PPQ responses included this CORS exposure header: Access-Control-Expose-Headers: Ehbp-Response-Nonce, X-Private-Usage-Metrics, X-Encrypted-Usage-Metrics, X-Tinfoil-Usage-Metrics ``` -However, the actual tested non-streaming response did not include any of these usage headers, even when `X-Tinfoil-Request-Usage-Metrics: true` was sent. +However, when this was tested against PPQ's `/private/` endpoint +(`private/gpt-oss-120b`) on 2026-06-21, the non-streaming response did not +include any of these usage headers, even when +`X-Tinfoil-Request-Usage-Metrics: true` was sent. That observation does *not* +hold for the direct Tinfoil enclave upstream that Routstr ships: see +[Usage metrics header format](#usage-metrics-header-format) below, where the +response header and the streaming trailer are both verified present. The decrypted body did include normal OpenAI usage, but only the decrypting Tinfoil client can see that body. @@ -385,7 +391,9 @@ and `routstr/upstream/ehbp.py`. - Base URL: `https://inference.tinfoil.sh` - Fetches models from the public `GET /v1/models` endpoint (no auth needed). - Parses Tinfoil's pricing (`inputTokenPricePer1M`, `outputTokenPricePer1M`, - `requestPrice`) into the standard `Model`/`Pricing` schema. + `cachedInputTokenPricePer1M`, `requestPrice`) into the standard + `Model`/`Pricing` schema. Cached reads use the cached rate when present, + otherwise the full input rate; cache writes always use the full input rate. - `supports_ehbp = True` — acts as a blind EHBP relay. - `get_ehbp_forwarding_target()` returns a target that includes `X-Tinfoil-Request-Usage-Metrics: true`. @@ -396,9 +404,12 @@ and `routstr/upstream/ehbp.py`. - `routstr/upstream/ehbp.py`: - `parse_tinfoil_usage_metrics()` parses - `prompt=N,completion=N[,total=N][,model=]` into an OpenAI-style - usage dict. The `model` field (added in tinfoilsh/confidential-model-router - PR #385) is extracted as a string. + `prompt=N,completion=N[,total=N][,cached_prompt_tokens=N, + uncached_prompt_tokens=N][,model=][,cost_usd=]` into an + OpenAI-style usage dict. Cache reads map to ``cache_read_input_tokens`` + so ``calculate_cost`` can bill them at the cached rate. The ``model`` + field (added in tinfoilsh/confidential-model-router PR #385) is extracted + as a string; ``cost_usd`` is parsed for logging only. - `_resolve_ehbp_target_url()` overrides the forwarding URL with `X-Tinfoil-Enclave-Url` when the SDK sends it. - `_strip_proxy_headers()` removes `X-Routstr-Model`, @@ -444,6 +455,8 @@ Routstr returns cost info as response headers: | `X-Routstr-Cost-Usd` | Bearer | USD equivalent of the computed usage | | `X-Routstr-Input-Cost-Msats` | Bearer, X-Cashu | msats attributed to input tokens | | `X-Routstr-Output-Cost-Msats` | Bearer, X-Cashu | msats attributed to output tokens | +| `X-Routstr-Cache-Read-Msats` | Bearer, X-Cashu | msats attributed to cached input (cache reads) | +| `X-Routstr-Cache-Creation-Msats` | Bearer, X-Cashu | msats attributed to cache creation (0 for Tinfoil today) | The client/Tinfoil SDK can read these headers from the HTTP response without needing to decrypt the body. A duplicate or rejected finalization can therefore @@ -467,9 +480,23 @@ Tinfoil returns usage metrics in the `X-Tinfoil-Usage-Metrics` response header true` is sent. As of tinfoilsh/confidential-model-router PR #385, the format is: ``` -prompt=,completion=,total=,model= +prompt=,completion=,total=[,cached_prompt_tokens=,uncached_prompt_tokens=][,model=][,cost_usd=] ``` +`prompt` is the inclusive prompt total; `cached_prompt_tokens` is the portion +already in Tinfoil's prefix cache and is billed at the model's +`cachedInputTokenPricePer1M` rate (or the full input rate when the model has +no cached rate). `cost_usd` is Tinfoil's own computed request cost and is +currently parsed for observability only — Routstr bills from token counts. + +Note that the header/trailer value is not always a single occurrence: for +streaming responses the trailer is emitted twice, so a client that reads the +trailer directly may see the same `prompt=...,completion=...,...` string twice +in one field, comma-joined. Parsers must be tolerant of the duplicate rather +than assuming exactly one occurrence. `parse_tinfoil_usage_metrics()` is +unaffected: it assigns each `key=value` part as it walks the comma-separated +value, and both occurrences carry identical numbers. + The `model` field carries the actual model name served by the enclave. Routstr uses this to: @@ -496,4 +523,6 @@ back to the requested model's pricing. finalizers for bearer and X-Cashu requests. This provides actual-cost billing today, at the cost of full time-to-last-byte latency for streaming responses. - Whether Tinfoil's `/v1/responses` endpoint also returns usage metrics - headers or trailers. + headers or trailers. Verified: yes — `/v1/responses` returns + `X-Tinfoil-Usage-Metrics` as a plaintext response header, with the same field + set as `/v1/chat/completions` (including `cost_usd`). diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index 61cd219c..90b92719 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -127,28 +127,43 @@ _PROXY_ONLY_HEADERS = frozenset( } ) +# Namespace prefix the routstr catalog applies to Tinfoil models +# (e.g. ``tinfoil-deepseek-v4-1-flash``). The SDK strips this prefix for the +# encrypted body (``getTinfoilUpstreamModelId`` in client/TinfoilSecure.ts), so +# the enclave always reports the *bare* upstream model id in the usage-metrics +# header even though the routstr model id and ``forwarded_model_id`` carry it. +TINFOIL_MODEL_PREFIX = "tinfoil-" + def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None: """Parse ``X-Tinfoil-Usage-Metrics`` into an OpenAI-style usage dict. The header format is:: - prompt=,completion=,total=[,model=] + prompt=,completion=,total=[,cached_prompt_tokens=, + uncached_prompt_tokens=][,model=][,cost_usd=] + + ``prompt`` is the inclusive prompt total and ``cached_prompt_tokens`` is + the cache-read portion included within it. Routstr maps these to + ``prompt_tokens`` and ``cache_read_input_tokens`` so ``normalize_usage`` + can subtract the cached read from the prompt total (OpenAI-family + semantics). ``cost_usd`` is parsed as a float and kept for logging/ + cross-checking only — billing uses the token path. The ``model`` field (added in tinfoilsh/confidential-model-router PR #385) - is extracted as a string and included in the returned dict under the - ``"model"`` key so callers can compare the served model against the - requested one and adjust pricing. + is extracted as a string so callers can compare the served model against + the requested one and adjust pricing. - Returns a dict like ``{"prompt_tokens": n, "completion_tokens": n, - "model": ""}`` suitable for :func:`calculate_cost` (which ignores - the extra ``model`` key in the usage sub-dict), or ``None`` when the + Returns a dict suitable for :func:`calculate_cost`, or ``None`` when the header is absent or malformed. """ if not header_value: return None - parts: dict[str, int] = {} + + int_parts: dict[str, int] = {} model: str | None = None + cost_usd: float | None = None + for item in header_value.split(","): key, sep, value = item.partition("=") if not sep: @@ -158,30 +173,44 @@ def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None: if key == "model": model = value continue + if key == "cost_usd": + try: + cost_usd = float(value) + except (ValueError, TypeError): + cost_usd = None + continue try: - parts[key] = int(value) + int_parts[key] = int(value) except (ValueError, TypeError): continue - prompt = parts.get("prompt") - completion = parts.get("completion") - if prompt is not None and completion is not None: - result: dict[str, int | str] = { - "prompt_tokens": prompt, - "completion_tokens": completion, - } - if "total" in parts: - result["total_tokens"] = parts["total"] - if model: - result["model"] = model - return result - logger.warning( - "Failed to parse X-Tinfoil-Usage-Metrics header", - extra={ - "header_value": header_value, - "parsed_parts": parts, - }, - ) - return None + + prompt = int_parts.get("prompt") + completion = int_parts.get("completion") + if prompt is None or completion is None: + logger.warning( + "Failed to parse X-Tinfoil-Usage-Metrics header", + extra={ + "header_value": header_value, + "parsed_parts": int_parts, + }, + ) + return None + + result: dict[str, int | float | str] = { + "prompt_tokens": prompt, + "completion_tokens": completion, + } + if "total" in int_parts: + result["total_tokens"] = int_parts["total"] + if "cached_prompt_tokens" in int_parts: + result["cache_read_input_tokens"] = int_parts["cached_prompt_tokens"] + if "uncached_prompt_tokens" in int_parts: + result["uncached_prompt_tokens"] = int_parts["uncached_prompt_tokens"] + if cost_usd is not None: + result["cost_usd"] = cost_usd + if model: + result["model"] = model + return result def _get_header_case_insensitive( @@ -330,6 +359,11 @@ def _build_cost_info( output_tokens: int = 0, input_msats: int = 0, output_msats: int = 0, + cache_read_input_tokens: int = 0, + cache_creation_input_tokens: int = 0, + cache_read_msats: int = 0, + cache_creation_msats: int = 0, + total_usd: float = 0.0, actual_model: str | None = None, ) -> dict: """Build a cost-info dict with token counts and per-token-type costs. @@ -338,13 +372,18 @@ def _build_cost_info( one), it is included in the returned dict so callers can use it for billing finalization and logging. """ - result: dict[str, int | str | None] = { + result: dict[str, int | float | str | None] = { "total_msats": total_msats, "input_tokens": input_tokens, "output_tokens": output_tokens, "total_tokens": input_tokens + output_tokens, "input_msats": input_msats, "output_msats": output_msats, + "cache_read_input_tokens": cache_read_input_tokens, + "cache_creation_input_tokens": cache_creation_input_tokens, + "cache_read_msats": cache_read_msats, + "cache_creation_msats": cache_creation_msats, + "total_usd": total_usd, } if actual_model: result["actual_model"] = actual_model @@ -363,6 +402,10 @@ def _inject_cost_response_headers(headers: dict[str, str], cost_info: dict) -> N headers["X-Routstr-Computed-Cost-Msats"] = str(cost_info["computed_msats"]) headers["X-Routstr-Input-Cost-Msats"] = str(cost_info["input_msats"]) headers["X-Routstr-Output-Cost-Msats"] = str(cost_info["output_msats"]) + headers["X-Routstr-Cache-Read-Msats"] = str(cost_info.get("cache_read_msats", 0)) + headers["X-Routstr-Cache-Creation-Msats"] = str( + cost_info.get("cache_creation_msats", 0) + ) async def _compute_ehbp_actual_cost( @@ -400,6 +443,14 @@ async def _compute_ehbp_actual_cost( # look up the actual model's pricing. actual_model: str | None = usage_dict.pop("model", None) # type: ignore[arg-type] pricing_model_id = model_obj.id + # Bill the model we actually routed to. Passing only the model *string* + # to calculate_cost makes it re-derive pricing from the global alias map, + # which resolves the id to the best-ranked candidate — not the serving + # one. Tinfoil's catalog id (e.g. ``deepseek-v4-1-flash``) is also a + # cross-provider alias, and that cheaper candidate has no cache rate, so + # the cache discount silently disappeared (and the request was + # undercharged). Hand calculate_cost the identity it cannot reconstruct. + pricing_model_obj: Model = model_obj expected_upstream_model = model_obj.forwarded_model_id or model_obj.id expected_identity = _normalize_upstream_model_id(expected_upstream_model) served_identity = _normalize_upstream_model_id(actual_model) @@ -415,7 +466,24 @@ async def _compute_ehbp_actual_cost( # the global model map. The resolved object can belong to a different # provider and therefore have a different client-facing ``id`` while # still representing the same upstream model. - actual_model_obj = get_model_instance(actual_model) + # + # The enclave reports the *bare* upstream id, but the routstr model is + # namespaced ``tinfoil-`` (and the SDK strips that prefix for the + # encrypted body). Resolve the served id within the same namespace + # first: a same-model report then maps back onto the requested Tinfoil + # model, and a genuine failover lands on the actually-served Tinfoil + # model — instead of the cheaper cross-provider model the bare id + # would resolve to in the global map. + namespaced_served = actual_model + if ( + expected_upstream_model.startswith(TINFOIL_MODEL_PREFIX) + and not actual_model.startswith(TINFOIL_MODEL_PREFIX) + ): + namespaced_served = TINFOIL_MODEL_PREFIX + actual_model + + actual_model_obj = get_model_instance(namespaced_served) + if actual_model_obj is None and namespaced_served != actual_model: + actual_model_obj = get_model_instance(actual_model) if actual_model_obj is None: logger.warning( "EHBP served model not found in registry, falling back " @@ -444,6 +512,7 @@ async def _compute_ehbp_actual_cost( }, ) pricing_model_id = actual_model_obj.id + pricing_model_obj = actual_model_obj else: # A different registry/client alias resolved to the same # upstream model; retain the requested model's pricing. @@ -456,6 +525,7 @@ async def _compute_ehbp_actual_cost( cost = await calculate_cost( {"model": pricing_model_id, "usage": usage_dict}, max_cost_for_model, + pricing_model_obj, ) except Exception as e: logger.warning( @@ -499,6 +569,11 @@ async def _compute_ehbp_actual_cost( output_tokens=cost.output_tokens, input_msats=cost.input_msats, output_msats=cost.output_msats, + cache_read_input_tokens=cost.cache_read_input_tokens, + cache_creation_input_tokens=cost.cache_creation_input_tokens, + cache_read_msats=cost.cache_read_msats, + cache_creation_msats=cost.cache_creation_msats, + total_usd=cost.total_usd, actual_model=actual_model, ) # CostDataError @@ -632,6 +707,16 @@ async def finalize_ehbp_actual_cost_payment( "cost_charged": total_cost_msats, "input_tokens": cost_info.get("input_tokens", 0), "output_tokens": cost_info.get("output_tokens", 0), + # Cache splits are only knowable when the enclave reports + # ``cached_prompt_tokens``; absent that they are a measured zero on + # the token counts the provider did report (not an unknown), so the + # event key set stays stable for usage-analytics consumers. + "cache_read_input_tokens": cost_info.get("cache_read_input_tokens", 0), + "cache_creation_input_tokens": cost_info.get( + "cache_creation_input_tokens", 0 + ), + "cache_read_msats": cost_info.get("cache_read_msats", 0), + "cache_creation_msats": cost_info.get("cache_creation_msats", 0), "balance": key.balance, "reserved_balance": key.reserved_balance, "total_spent": key.total_spent, @@ -855,7 +940,7 @@ async def forward_ehbp_request( **cost_info, "total_msats": charged_msats, "charged_msats": charged_msats, - "total_usd": 0.0, + "total_usd": cost_info.get("total_usd", 0.0), } if computed_msats != charged_msats: cost_data["computed_msats"] = computed_msats @@ -894,6 +979,12 @@ async def forward_ehbp_request( + cost_data.get("output_tokens", 0), "input_msats": cost_data.get("input_msats", 0), "output_msats": cost_data.get("output_msats", 0), + "cache_read_input_tokens": cost_data.get("cache_read_input_tokens", 0), + "cache_creation_input_tokens": cost_data.get( + "cache_creation_input_tokens", 0 + ), + "cache_read_msats": cost_data.get("cache_read_msats", 0), + "cache_creation_msats": cost_data.get("cache_creation_msats", 0), } if "computed_msats" in cost_data: cost_info["computed_msats"] = cost_data["computed_msats"] diff --git a/routstr/upstream/tinfoil.py b/routstr/upstream/tinfoil.py index 6eeb1147..0928cbc9 100644 --- a/routstr/upstream/tinfoil.py +++ b/routstr/upstream/tinfoil.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Optional import httpx from fastapi import Request @@ -28,6 +28,7 @@ logger = get_logger(__name__) class TinfoilModelPricing(BaseModel): inputTokenPricePer1M: float = 0.0 outputTokenPricePer1M: float = 0.0 + cachedInputTokenPricePer1M: Optional[float] = None requestPrice: float = 0.0 @@ -186,6 +187,14 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider): output_price = tf.pricing.outputTokenPricePer1M request_price = tf.pricing.requestPrice + # Tinfoil bills cache reads at the cached rate when the + # model exposes one, otherwise at the full input rate. + # Cache writes are never priced separately — a miss is + # just regular input prefill. + cached_price = tf.pricing.cachedInputTokenPricePer1M + if cached_price is None or cached_price <= 0.0: + cached_price = input_price + modality = "text->text" input_modalities = ["text"] output_modalities = ["text"] @@ -214,6 +223,8 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider): image=0.0, web_search=0.0, internal_reasoning=0.0, + input_cache_read=cached_price / 1_000_000, + input_cache_write=input_price / 1_000_000, ), ) ) diff --git a/tests/unit/test_ehbp_finalize_payment.py b/tests/unit/test_ehbp_finalize_payment.py index 2f905e63..caf0a5aa 100644 --- a/tests/unit/test_ehbp_finalize_payment.py +++ b/tests/unit/test_ehbp_finalize_payment.py @@ -1,6 +1,8 @@ from __future__ import annotations -from typing import Any, AsyncGenerator +import logging +from contextlib import contextmanager +from typing import Any, AsyncGenerator, Iterator from unittest.mock import AsyncMock, MagicMock import pytest @@ -106,6 +108,111 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve assert updated.total_spent == 1_200 +@contextmanager +def _capture_payments_logs() -> Iterator[list[logging.LogRecord]]: + """Collect ``routstr.payments`` records for the duration of the block. + + ``setup_logging()`` sets ``propagate=False`` on the ``routstr`` logger, so + pytest's ``caplog`` (attached at the root) never sees these records; a + handler on the payments logger itself does. + """ + records: list[logging.LogRecord] = [] + + class _RecordingHandler(logging.Handler): + def emit(self, record: logging.LogRecord) -> None: + records.append(record) + + payments_logger = logging.getLogger("routstr.payments") + handler = _RecordingHandler(level=logging.INFO) + previous_level = payments_logger.level + payments_logger.addHandler(handler) + payments_logger.setLevel(logging.INFO) + try: + yield records + finally: + payments_logger.removeHandler(handler) + payments_logger.setLevel(previous_level) + + +@pytest.mark.asyncio +async def test_finalize_actual_cost_payment_logs_cache_tokens( + session: AsyncSession, +) -> None: + """The FINALIZE event carries the cache splits, not just input/output.""" + key = ApiKey(hashed_key="ehbp-cache-logging", balance=10_000) + session.add(key) + await session.commit() + await pay_for_request(key, 3_000, session) + reservation = await get_reservation_snapshot(key, session) + + with _capture_payments_logs() as records: + charged = await finalize_ehbp_actual_cost_payment( + key, + session, + reserved_cost_for_model=3_000, + model_id="tinfoil/glm-5-2", + cost_info={ + "total_msats": 1_200, + "input_tokens": 5, + "output_tokens": 20, + "input_msats": 500, + "output_msats": 700, + "cache_read_input_tokens": 64, + "cache_creation_input_tokens": 0, + "cache_read_msats": 12, + "cache_creation_msats": 0, + }, + reservation_snapshot=reservation, + ) + + assert charged == 1_200 + finalize_records = [ + record for record in records if record.getMessage() == "FINALIZE" + ] + assert len(finalize_records) == 1 + record = finalize_records[0] + # finalize_type/input_tokens/... are attached via logging's extra= payload. + assert record.finalize_type == "ehbp_usage" # type: ignore[attr-defined] + assert record.input_tokens == 5 # type: ignore[attr-defined] + assert record.output_tokens == 20 # type: ignore[attr-defined] + assert record.cache_read_input_tokens == 64 # type: ignore[attr-defined] + assert record.cache_creation_input_tokens == 0 # type: ignore[attr-defined] + assert record.cache_read_msats == 12 # type: ignore[attr-defined] + assert record.cache_creation_msats == 0 # type: ignore[attr-defined] + + +@pytest.mark.asyncio +async def test_finalize_actual_cost_payment_logs_zero_cache_when_absent( + session: AsyncSession, +) -> None: + """Providers that report no cache split still emit a stable key set.""" + key = ApiKey(hashed_key="ehbp-no-cache-logging", balance=10_000) + session.add(key) + await session.commit() + await pay_for_request(key, 3_000, session) + reservation = await get_reservation_snapshot(key, session) + + with _capture_payments_logs() as records: + await finalize_ehbp_actual_cost_payment( + key, + session, + reserved_cost_for_model=3_000, + model_id="tinfoil/glm-5-2", + cost_info={ + "total_msats": 1_200, + "input_tokens": 10, + "output_tokens": 20, + }, + reservation_snapshot=reservation, + ) + + record = next(record for record in records if record.getMessage() == "FINALIZE") + assert record.cache_read_input_tokens == 0 # type: ignore[attr-defined] + assert record.cache_creation_input_tokens == 0 # type: ignore[attr-defined] + assert record.cache_read_msats == 0 # type: ignore[attr-defined] + assert record.cache_creation_msats == 0 # type: ignore[attr-defined] + + @pytest.mark.asyncio async def test_unmeasured_ehbp_releases_reservation( session: AsyncSession, diff --git a/tests/unit/test_tinfoil_integration.py b/tests/unit/test_tinfoil_integration.py index 1319e5a0..36e76765 100644 --- a/tests/unit/test_tinfoil_integration.py +++ b/tests/unit/test_tinfoil_integration.py @@ -109,7 +109,25 @@ class TestParseTinfoilUsageMetrics: assert result["prompt_tokens"] == 69 assert result["completion_tokens"] == 20 assert result["total_tokens"] == 89 + assert result["cache_read_input_tokens"] == 64 + assert result["uncached_prompt_tokens"] == 5 assert result["model"] == "kimi-k2-6" + assert "cost_usd" not in result + + def test_with_cached_and_cost_usd(self) -> None: + result = parse_tinfoil_usage_metrics( + "prompt=69,completion=20,total=89," + "cached_prompt_tokens=64,uncached_prompt_tokens=5," + "model=glm-5-2,cost_usd=0.000123456" + ) + assert result is not None + assert result["prompt_tokens"] == 69 + assert result["completion_tokens"] == 20 + assert result["total_tokens"] == 89 + assert result["cache_read_input_tokens"] == 64 + assert result["uncached_prompt_tokens"] == 5 + assert result["cost_usd"] == 0.000123456 + assert result["model"] == "glm-5-2" def test_old_format_still_works(self) -> None: """Headers without the model field (pre-PR #385) still parse.""" @@ -289,6 +307,46 @@ class TestComputeEhbpActualCost: assert result["input_msats"] == 10 assert result["output_msats"] == 20 + @pytest.mark.asyncio + async def test_cache_fields_propagated(self) -> None: + model_obj = MagicMock() + model_obj.id = "tinfoil-glm-5-2" + model_obj.forwarded_model_id = "glm-5-2" + with patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc: + from routstr.payment.cost_calculation import CostData + + mock_calc.return_value = CostData( + base_msats=0, + input_msats=5, + output_msats=20, + total_msats=25, + total_usd=0.0003, + input_tokens=5, + output_tokens=20, + cache_read_input_tokens=64, + cache_creation_input_tokens=0, + cache_read_msats=12, + cache_creation_msats=0, + ) + result = await _compute_ehbp_actual_cost( + "prompt=69,completion=20,total=89," + "cached_prompt_tokens=64,uncached_prompt_tokens=5," + "model=glm-5-2", + model_obj, + 100_000, + ) + assert result["total_msats"] == 25 + assert result["input_tokens"] == 5 + assert result["output_tokens"] == 20 + assert result["cache_read_input_tokens"] == 64 + assert result["cache_creation_input_tokens"] == 0 + assert result["cache_read_msats"] == 12 + assert result["cache_creation_msats"] == 0 + assert result["total_usd"] == 0.0003 + @pytest.mark.asyncio async def test_unpriceable_usage_does_not_charge_authorization_ceiling( self, @@ -430,6 +488,230 @@ class TestComputeEhbpActualCost: call_args = mock_calc.call_args assert call_args[0][0]["model"] == "tinfoil-llama3-3-70b" + @pytest.mark.asyncio + async def test_namespaced_prefix_served_bare_keeps_requested_pricing( + self, + ) -> None: + """Production shape: the catalog model is ``tinfoil-X`` and its + ``forwarded_model_id`` carries the prefix, the SDK strips the prefix + for the encrypted body, and the enclave reports bare ``X``. Pricing + must stay on the requested Tinfoil model (correct rate + cache + discount), not the cheaper cross-provider model the bare id resolves + to in the global map.""" + model_obj = MagicMock() + model_obj.id = "tinfoil-deepseek-v4-1-flash" + model_obj.forwarded_model_id = "tinfoil-deepseek-v4-1-flash" + + tinfoil_model = MagicMock() + tinfoil_model.id = "tinfoil-deepseek-v4-1-flash" + tinfoil_model.forwarded_model_id = "tinfoil-deepseek-v4-1-flash" + + # The cheaper cross-provider model the bare id resolves to globally. + cross_provider_model = MagicMock() + cross_provider_model.id = "deepseek-v4-1-flash" + cross_provider_model.forwarded_model_id = "deepseek-v4-1-flash" + + registry = { + "tinfoil-deepseek-v4-1-flash": tinfoil_model, + "deepseek-v4-1-flash": cross_provider_model, + } + + with ( + patch( + "routstr.proxy.get_model_instance", + side_effect=lambda name: registry.get(name), + ), + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): + from routstr.payment.cost_calculation import CostData + + mock_calc.return_value = CostData( + base_msats=0, + input_msats=5, + output_msats=10, + total_msats=15, + total_usd=0.0, + input_tokens=5, + output_tokens=10, + cache_read_input_tokens=64, + cache_creation_input_tokens=0, + cache_read_msats=1, + cache_creation_msats=0, + ) + result = await _compute_ehbp_actual_cost( + "prompt=69,completion=10,total=79," + "cached_prompt_tokens=64,uncached_prompt_tokens=5," + "model=deepseek-v4-1-flash", + model_obj, + 100_000, + ) + # No mismatch: pricing stays on the requested Tinfoil model. + assert "actual_model" not in result + call_args = mock_calc.call_args + assert call_args[0][0]["model"] == "tinfoil-deepseek-v4-1-flash" + + @pytest.mark.asyncio + async def test_namespaced_prefix_failover_uses_served_tinfoil_model( + self, + ) -> None: + """A genuine failover (asked ``tinfoil-glm-5-3``, enclave served + ``glm-5-3-flash``) must bill the served *Tinfoil* model, not the + cheaper cross-provider alias the bare id resolves to.""" + model_obj = MagicMock() + model_obj.id = "tinfoil-glm-5-3" + model_obj.forwarded_model_id = "tinfoil-glm-5-3" + + served_tinfoil = MagicMock() + served_tinfoil.id = "tinfoil-glm-5-3-flash" + served_tinfoil.forwarded_model_id = "tinfoil-glm-5-3-flash" + + cross_provider = MagicMock() + cross_provider.id = "glm-5-3-flash" + cross_provider.forwarded_model_id = "glm-5-3-flash" + + registry = { + "tinfoil-glm-5-3-flash": served_tinfoil, + "glm-5-3-flash": cross_provider, + } + + with ( + patch( + "routstr.proxy.get_model_instance", + side_effect=lambda name: registry.get(name), + ), + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): + from routstr.payment.cost_calculation import CostData + + mock_calc.return_value = CostData( + base_msats=0, + input_msats=20, + output_msats=40, + total_msats=60, + total_usd=0.0, + input_tokens=42, + output_tokens=10, + ) + result = await _compute_ehbp_actual_cost( + "prompt=42,completion=10,total=52,model=glm-5-3-flash", + model_obj, + 100_000, + ) + assert result["actual_model"] == "glm-5-3-flash" + # Billed on the served *Tinfoil* model, not the bare-id alias. + call_args = mock_calc.call_args + assert call_args[0][0]["model"] == "tinfoil-glm-5-3-flash" + + @pytest.mark.asyncio + async def test_calculate_cost_receives_routed_model_obj(self) -> None: + """The routed ``Model`` is handed to ``calculate_cost`` so pricing is + billed directly. Without it, ``calculate_cost`` re-derives pricing from + the response's model *string* through the global alias map, which + resolves a bare id to the best-ranked (cheaper) cross-provider + candidate rather than the serving one.""" + model_obj = MagicMock() + model_obj.id = "deepseek-v4-1-flash" + model_obj.forwarded_model_id = "tinfoil-deepseek-v4-1-flash" + + resolved = MagicMock() + resolved.id = "tinfoil-deepseek-v4-1-flash" + resolved.forwarded_model_id = "tinfoil-deepseek-v4-1-flash" + + with ( + patch( + "routstr.proxy.get_model_instance", + return_value=resolved, + ), + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): + from routstr.payment.cost_calculation import CostData + + mock_calc.return_value = CostData( + base_msats=0, + input_msats=5, + output_msats=10, + total_msats=15, + total_usd=0.0, + input_tokens=5, + output_tokens=10, + ) + await _compute_ehbp_actual_cost( + "prompt=12952,completion=1,total=12953," + "cached_prompt_tokens=12800,uncached_prompt_tokens=152," + "model=deepseek-v4-1-flash,cost_usd=0.00176715", + model_obj, + 100_000, + ) + # The routed model object itself must be passed through. + assert mock_calc.call_args[0][2] is model_obj + + @pytest.mark.asyncio + async def test_routed_model_cache_rate_beats_bare_id_alias(self) -> None: + """Production regression: the routed Tinfoil model's id *is* a bare + cross-provider alias, so re-deriving pricing from the echoed model + string silently swapped in the cheaper candidate's full input rate and + the cache discount vanished. Billing must use the routed model's own + discounted cache rate.""" + from routstr.payment.models import Pricing + + model_obj = MagicMock() + model_obj.id = "deepseek-v4-1-flash" + model_obj.forwarded_model_id = "tinfoil-deepseek-v4-1-flash" + # ~688 msat/1k input, ~112 msat/1k cached read (the good Tinfoil rate). + model_obj.sats_pricing = Pricing( + prompt=6.88e-4, + completion=2.0e-3, + input_cache_read=1.12e-4, + ) + + # The cross-provider candidate the bare id resolves to globally: no + # cache rate at all, so a re-derivation charges the full input rate. + cross_provider_model = MagicMock() + cross_provider_model.id = "deepseek-v4-1-flash" + cross_provider_model.forwarded_model_id = "deepseek-v4-1-flash" + cross_provider_model.sats_pricing = Pricing( + prompt=4.9455e-4, + completion=2.0e-3, + input_cache_read=0.0, + ) + + registry = { + "tinfoil-deepseek-v4-1-flash": model_obj, + "deepseek-v4-1-flash": cross_provider_model, + } + + with ( + patch( + "routstr.proxy.get_model_instance", + side_effect=lambda name: registry.get(name), + ), + patch( + "routstr.payment.cost_calculation.sats_usd_price", + return_value=5.0e-5, + ), + ): + result = await _compute_ehbp_actual_cost( + "prompt=12952,completion=1,total=12953," + "cached_prompt_tokens=12800,uncached_prompt_tokens=152," + "model=deepseek-v4-1-flash,cost_usd=0.00176715", + model_obj, + 100_000, + ) + + assert result["cache_read_input_tokens"] == 12800 + # 12800 cached tokens at the discounted (~112 msat/1k) rate, not the + # full input rate (which would be ~8800 msat here). + assert result["cache_read_msats"] == pytest.approx(1434, abs=10) + @pytest.mark.asyncio async def test_model_mismatch_unknown_model_falls_back(self) -> None: """When the served model is not in the registry, use requested model.""" @@ -709,6 +991,20 @@ class TestTinfoilUpstreamProvider: assert tf.id == "llama3-3-70b" assert tf.pricing.inputTokenPricePer1M == 1.75 assert tf.pricing.outputTokenPricePer1M == 2.75 + assert tf.pricing.cachedInputTokenPricePer1M is None + + def test_tinfoil_model_pricing_parses_cached_rate(self) -> None: + data = { + "id": "glm-5-2", + "pricing": { + "inputTokenPricePer1M": 1.5, + "outputTokenPricePer1M": 5.25, + "cachedInputTokenPricePer1M": 0.375, + "requestPrice": 0, + }, + } + tf = TinfoilModel.parse_obj(data) + assert tf.pricing.cachedInputTokenPricePer1M == 0.375 @pytest.mark.asyncio async def test_fetch_models_parses_response(self) -> None: @@ -747,8 +1043,52 @@ class TestTinfoilUpstreamProvider: assert models[0].id == "llama3-3-70b" assert models[0].pricing.prompt == 1.75 / 1_000_000 assert models[0].pricing.completion == 2.75 / 1_000_000 + # No cachedInputTokenPricePer1M means cache reads are billed at the + # full input rate (and cache writes too — no separate write price). + assert models[0].pricing.input_cache_read == 1.75 / 1_000_000 + assert models[0].pricing.input_cache_write == 1.75 / 1_000_000 assert models[0].context_length == 128000 + @pytest.mark.asyncio + async def test_fetch_models_maps_cached_pricing(self) -> None: + provider = TinfoilUpstreamProvider(api_key="test") + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.raise_for_status = MagicMock() + mock_response.json.return_value = { + "data": [ + { + "id": "glm-5-2", + "context_window": 393216, + "created": 1775088000, + "multimodal": False, + "pricing": { + "inputTokenPricePer1M": 1.5, + "outputTokenPricePer1M": 5.25, + "cachedInputTokenPricePer1M": 0.375, + "requestPrice": 0, + }, + "endpoints": ["/v1/chat/completions", "/v1/responses"], + "type": "chat", + } + ] + } + + with patch("routstr.upstream.tinfoil.httpx.AsyncClient") as mock_client_cls: + mock_client = MagicMock() + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client.get = AsyncMock(return_value=mock_response) + mock_client_cls.return_value = mock_client + + models = await provider.fetch_models() + + assert len(models) == 1 + assert models[0].pricing.prompt == 1.5 / 1_000_000 + assert models[0].pricing.completion == 5.25 / 1_000_000 + assert models[0].pricing.input_cache_read == 0.375 / 1_000_000 + assert models[0].pricing.input_cache_write == 1.5 / 1_000_000 + @pytest.mark.asyncio async def test_fetch_models_handles_error(self) -> None: provider = TinfoilUpstreamProvider(api_key="test")