Add Tinfoil cached prompt pricing support

Parse cachedInputTokenPricePer1M and cached_prompt_tokens so EHBP
relays bill cache reads at the discounted cached rate instead of the
full input rate. Models without a cached rate fall back to full input
pricing, and cache writes stay at full input since Tinfoil has no
separate cache-write price.
This commit is contained in:
redshift
2026-09-17 20:10:44 +02:00
parent 303ff32ecc
commit b5c2b1b10b
4 changed files with 226 additions and 37 deletions
+18 -5
View File
@@ -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=<name>]` 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=<name>][,cost_usd=<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=<prompt_tokens>,completion=<completion_tokens>,total=<total_tokens>,model=<served_model>
prompt=<prompt_tokens>,completion=<completion_tokens>,total=<total_tokens>[,cached_prompt_tokens=<n>,uncached_prompt_tokens=<n>][,model=<served_model>][,cost_usd=<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:
+80 -31
View File
@@ -133,22 +133,30 @@ def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None:
The header format is::
prompt=<n>,completion=<n>,total=<n>[,model=<name>]
prompt=<n>,completion=<n>,total=<n>[,cached_prompt_tokens=<n>,
uncached_prompt_tokens=<n>][,model=<name>][,cost_usd=<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": "<name>"}`` 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"]
+12 -1
View File
@@ -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,
),
)
)
+116
View File
@@ -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")