mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
Merge pull request #736 from Routstr/tinfoil-cache-pricing
feat(ehbp): bill Tinfoil cached prompt tokens at the cached rate
This commit is contained in:
@@ -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=<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 +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=<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.
|
||||
|
||||
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`).
|
||||
|
||||
+123
-32
@@ -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=<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 +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"]
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user