mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
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:
@@ -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
@@ -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"]
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user