diff --git a/docs/tinfoil-direct-integration.md b/docs/tinfoil-direct-integration.md index d41228ec..685df512 100644 --- a/docs/tinfoil-direct-integration.md +++ b/docs/tinfoil-direct-integration.md @@ -385,7 +385,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 +398,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 +449,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 +474,15 @@ 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. + The `model` field carries the actual model name served by the enclave. Routstr uses this to: diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index 61cd219c..74051b00 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -133,22 +133,30 @@ def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None: 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 +166,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 +352,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 +365,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 +395,12 @@ 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( @@ -499,6 +537,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 @@ -855,7 +898,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 +937,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_tinfoil_integration.py b/tests/unit/test_tinfoil_integration.py index 1319e5a0..bf6fa0c5 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, @@ -709,6 +767,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 +819,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")