Files
routstr-core/tests/unit/test_tinfoil_integration.py
T
redshift 3c813daedc fix(ehbp): resolve the served model within the tinfoil namespace before pricing
The cache discount never applied in production even though everything
downstream of it worked: the enclave reported cache splits
(cached_prompt_tokens=160512 of 161652) and the parse put them into
cache_read_input_tokens, but the applied cache-read rate was the full input
rate.

Root cause: the SDK strips the routstr "tinfoil-" namespace prefix for the
encrypted body (getTinfoilUpstreamModelId), so the enclave always reports the
bare upstream id ("deepseek-v4-1-flash") in X-Tinfoil-Usage-Metrics, while
the catalog registers the model as "tinfoil-deepseek-v4-1-flash" and
forwarded_model_id carries the prefix too. _normalize_upstream_model_id only
lowercases, so every request took the "served model differs" path and
re-derived pricing via get_model_instance on the bare id — which resolves
globally to a cheaper cross-provider model whose pricing has no cache rate
(input_cache_read=0), falling back to the full input price.

Billing was therefore doubly wrong: no cache discount, and E2EE requests
undercharged at the cross-provider rate instead of the Tinfoil rate.

Fix: when the requested model is namespaced "tinfoil-" and the served id is
not, resolve the served id within the same namespace first (with a bare-id
fallback). A same-model report then maps back onto the requested Tinfoil
model (keeping its pricing), and a genuine failover lands on the
actually-served Tinfoil model — while the bare-id lookup that previously
hijacked pricing is only used as a last resort.
2026-09-17 22:20:25 +02:00

1274 lines
48 KiB
Python

"""Unit tests for Tinfoil direct integration.
Covers the EHBP usage-metrics header parser, the proxy header stripping, the
enclave URL override, and the TinfoilUpstreamProvider model fetching/forwarding
target logic.
"""
from __future__ import annotations
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from routstr import proxy as proxy_module
from routstr.core.db import ApiKey
from routstr.upstream.ehbp import (
_PROXY_ONLY_HEADERS,
EHBPForwardingTarget,
_compute_ehbp_actual_cost,
_is_ehbp_key_config_response,
_passthrough_key_config_response,
_prepare_ehbp_upstream_headers,
_resolve_ehbp_target_url,
_strip_proxy_headers,
forward_ehbp_x_cashu_request,
parse_tinfoil_usage_metrics,
)
from routstr.upstream.tinfoil import (
TinfoilModel,
TinfoilUpstreamProvider,
)
from routstr.upstream.tinfoil_trailer import TrailerResponse
# ---------------------------------------------------------------------------
# parse_tinfoil_usage_metrics
# ---------------------------------------------------------------------------
class TestParseTinfoilUsageMetrics:
def test_full_header(self) -> None:
result = parse_tinfoil_usage_metrics("prompt=67,completion=42,total=109")
assert result == {
"prompt_tokens": 67,
"completion_tokens": 42,
"total_tokens": 109,
}
def test_without_total(self) -> None:
result = parse_tinfoil_usage_metrics("prompt=10,completion=5")
assert result == {"prompt_tokens": 10, "completion_tokens": 5}
def test_none(self) -> None:
assert parse_tinfoil_usage_metrics(None) is None
def test_empty(self) -> None:
assert parse_tinfoil_usage_metrics("") is None
def test_malformed(self) -> None:
assert parse_tinfoil_usage_metrics("garbage") is None
def test_missing_completion(self) -> None:
assert parse_tinfoil_usage_metrics("prompt=10") is None
def test_extra_whitespace(self) -> None:
result = parse_tinfoil_usage_metrics(
"prompt = 100 , completion = 200 , total = 300"
)
assert result == {
"prompt_tokens": 100,
"completion_tokens": 200,
"total_tokens": 300,
}
def test_with_model_field(self) -> None:
result = parse_tinfoil_usage_metrics(
"prompt=42,completion=10,total=52,model=llama3-3-70b"
)
assert result == {
"prompt_tokens": 42,
"completion_tokens": 10,
"total_tokens": 52,
"model": "llama3-3-70b",
}
def test_with_model_no_total(self) -> None:
result = parse_tinfoil_usage_metrics(
"prompt=67,completion=42,model=gpt-oss-120b"
)
assert result == {
"prompt_tokens": 67,
"completion_tokens": 42,
"model": "gpt-oss-120b",
}
def test_model_with_dashes_and_numbers(self) -> None:
result = parse_tinfoil_usage_metrics(
"prompt=1,completion=1,total=2,model=kimi-k2-6"
)
assert result is not None
assert result["model"] == "kimi-k2-6"
def test_model_with_extra_fields(self) -> None:
result = parse_tinfoil_usage_metrics(
"prompt=69,completion=20,total=89,"
"cached_prompt_tokens=64,uncached_prompt_tokens=5,"
"model=kimi-k2-6"
)
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["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."""
result = parse_tinfoil_usage_metrics("prompt=67,completion=42,total=109")
assert result == {
"prompt_tokens": 67,
"completion_tokens": 42,
"total_tokens": 109,
}
assert "model" not in result
# ---------------------------------------------------------------------------
# _strip_proxy_headers
# ---------------------------------------------------------------------------
class TestStripProxyHeaders:
def test_strips_all_proxy_only(self) -> None:
headers = {
"x-routstr-model": "tinfoil-llama3-3-70b",
"X-Tinfoil-Enclave-Url": "https://inference.tinfoil.sh",
"X-Tinfoil-Request-Usage-Metrics": "true",
"Authorization": "Bearer secret",
"Ehbp-Encapsulated-Key": "abc123",
}
clean = _strip_proxy_headers(headers)
assert "x-routstr-model" not in clean
assert "X-Tinfoil-Enclave-Url" not in clean
assert "X-Tinfoil-Request-Usage-Metrics" not in clean
assert clean["Authorization"] == "Bearer secret"
assert clean["Ehbp-Encapsulated-Key"] == "abc123"
def test_all_proxy_only_headers_covered(self) -> None:
assert _PROXY_ONLY_HEADERS == {
"x-routstr-model",
"x-tinfoil-enclave-url",
"x-tinfoil-request-usage-metrics",
}
class TestPrepareEHBPUpstreamHeaders:
def test_strips_client_proxy_headers_before_merging_target_headers(self) -> None:
headers = {
"x-routstr-model": "tinfoil-llama3-3-70b",
"X-Tinfoil-Enclave-Url": "https://enclave.tinfoil.sh",
"X-Tinfoil-Request-Usage-Metrics": "false",
"Authorization": "Bearer upstream-key",
"Ehbp-Encapsulated-Key": "abc123",
}
target_headers = {"X-Tinfoil-Request-Usage-Metrics": "true"}
clean = _prepare_ehbp_upstream_headers(headers, target_headers)
assert "x-routstr-model" not in clean
assert "X-Tinfoil-Enclave-Url" not in clean
assert clean["Authorization"] == "Bearer upstream-key"
assert clean["Ehbp-Encapsulated-Key"] == "abc123"
assert clean["X-Tinfoil-Request-Usage-Metrics"] == "true"
# ---------------------------------------------------------------------------
# _resolve_ehbp_target_url
# ---------------------------------------------------------------------------
class TestResolveEhbpTargetUrl:
def test_override_with_enclave_url_for_tinfoil(self) -> None:
result = _resolve_ehbp_target_url(
"https://default.example.com/v1/chat/completions",
"v1/chat/completions",
{"X-Tinfoil-Enclave-Url": "https://enclave.tinfoil.sh"},
"tinfoil",
)
assert result == "https://enclave.tinfoil.sh/v1/chat/completions"
def test_override_lowercase_header_for_tinfoil(self) -> None:
result = _resolve_ehbp_target_url(
"https://default.example.com/v1/chat/completions",
"v1/chat/completions",
{"x-tinfoil-enclave-url": "https://enclave.tinfoil.sh"},
"tinfoil",
)
assert result == "https://enclave.tinfoil.sh/v1/chat/completions"
def test_no_override(self) -> None:
default = "https://inference.tinfoil.sh/v1/chat/completions"
result = _resolve_ehbp_target_url(
default,
"v1/chat/completions",
{},
"tinfoil",
)
assert result == default
def test_non_tinfoil_provider_ignores_enclave_url(self) -> None:
default = "https://api.ppq.ai/private/v1/chat/completions"
result = _resolve_ehbp_target_url(
default,
"v1/chat/completions",
{"X-Tinfoil-Enclave-Url": "https://enclave.tinfoil.sh"},
"ppqai",
)
assert result == default
@pytest.mark.parametrize(
"bad_url",
[
"http://enclave.tinfoil.sh",
"https://attacker.example",
"https://tinfoil.sh.attacker.example",
"https://127.0.0.1",
"https://enclave.tinfoil.sh:8443",
"https://user:pass@enclave.tinfoil.sh",
],
)
def test_tinfoil_rejects_unsafe_enclave_url(self, bad_url: str) -> None:
from routstr.core.exceptions import UpstreamError
with pytest.raises(UpstreamError):
_resolve_ehbp_target_url(
"https://default.example.com/v1/chat/completions",
"v1/chat/completions",
{"X-Tinfoil-Enclave-Url": bad_url},
"tinfoil",
)
# ---------------------------------------------------------------------------
# _compute_ehbp_actual_cost
# ---------------------------------------------------------------------------
class TestComputeEhbpActualCost:
@pytest.mark.asyncio
async def test_no_usage_does_not_charge_authorization_ceiling(self) -> None:
model_obj = MagicMock()
model_obj.id = "llama3-3-70b"
model_obj.forwarded_model_id = "llama3-3-70b"
result = await _compute_ehbp_actual_cost(None, model_obj, 100_000)
assert result["total_msats"] == 0
assert result["input_tokens"] == 0
assert result["output_tokens"] == 0
@pytest.mark.asyncio
async def test_usage_parsed_and_clamped(self) -> None:
model_obj = MagicMock()
model_obj.id = "llama3-3-70b"
model_obj.forwarded_model_id = "llama3-3-70b"
# The actual cost from calculate_cost will be small; we just verify
# it's clamped to min_request_msat at minimum.
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=10,
output_msats=20,
total_msats=30,
total_usd=0.0001,
input_tokens=67,
output_tokens=42,
)
result = await _compute_ehbp_actual_cost(
"prompt=67,completion=42,total=109",
model_obj,
100_000,
)
assert result["total_msats"] == 30
assert result["total_msats"] <= 100_000
assert result["input_tokens"] == 67
assert result["output_tokens"] == 42
assert result["total_tokens"] == 109
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,
) -> None:
model_obj = MagicMock()
model_obj.id = "llama3-3-70b"
model_obj.forwarded_model_id = "llama3-3-70b"
with patch(
"routstr.upstream.ehbp.calculate_cost",
new_callable=AsyncMock,
) as mock_calc:
from routstr.payment.cost_calculation import MaxCostData
mock_calc.return_value = MaxCostData(
base_msats=0,
input_msats=0,
output_msats=0,
total_msats=0,
total_usd=0.0,
input_tokens=0,
output_tokens=0,
)
result = await _compute_ehbp_actual_cost(
"prompt=0,completion=0",
model_obj,
50_000,
)
assert result["total_msats"] == 0
assert result["input_tokens"] == 0
assert result["output_tokens"] == 0
@pytest.mark.asyncio
async def test_model_match_no_actual_model_key(self) -> None:
"""When the served model matches the requested one, no actual_model key."""
model_obj = MagicMock()
model_obj.id = "llama3-3-70b"
model_obj.forwarded_model_id = "llama3-3-70b"
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=10,
total_msats=15,
total_usd=0.0,
input_tokens=42,
output_tokens=10,
)
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=llama3-3-70b",
model_obj,
100_000,
)
assert "actual_model" not in result
# calculate_cost called with requested model
call_args = mock_calc.call_args
assert call_args[0][0]["model"] == "llama3-3-70b"
@pytest.mark.asyncio
async def test_alias_match_no_actual_model_key(self) -> None:
"""When the served upstream model matches forwarded_model_id through
a client-facing alias, no actual_model key is set."""
model_obj = MagicMock()
model_obj.id = "tinfoil-glm-5-2" # client-facing alias
model_obj.forwarded_model_id = "glm-5-2" # actual upstream ID
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=10,
total_msats=15,
total_usd=0.0,
input_tokens=42,
output_tokens=10,
)
# Tinfoil header returns the actual upstream model ID
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=glm-5-2",
model_obj,
100_000,
)
assert "actual_model" not in result
# calculate_cost called with the client-facing model ID (whose
# pricing includes the correct upstream rates)
call_args = mock_calc.call_args
assert call_args[0][0]["model"] == "tinfoil-glm-5-2"
@pytest.mark.asyncio
async def test_real_mismatch_uses_actual_model_for_pricing(self) -> None:
"""When the served model differs from the expected upstream model,
the actual model's pricing is used."""
model_obj = MagicMock()
model_obj.id = "tinfoil-gpt-oss-120b" # client-facing alias
model_obj.forwarded_model_id = "gpt-oss-120b" # expected upstream
actual_model_obj = MagicMock()
actual_model_obj.id = "tinfoil-llama3-3-70b" # client-facing of actual
actual_model_obj.forwarded_model_id = "llama3-3-70b"
with (
patch(
"routstr.proxy.get_model_instance",
return_value=actual_model_obj,
),
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,
)
# Tinfoil served llama3-3-70b instead of gpt-oss-120b
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=llama3-3-70b",
model_obj,
100_000,
)
assert result["actual_model"] == "llama3-3-70b"
assert result["total_msats"] == 60
# calculate_cost called with the actual model's client-facing ID
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_model_mismatch_unknown_model_falls_back(self) -> None:
"""When the served model is not in the registry, use requested model."""
model_obj = MagicMock()
model_obj.id = "gpt-oss-120b"
model_obj.forwarded_model_id = "gpt-oss-120b"
with (
patch(
"routstr.proxy.get_model_instance",
return_value=None,
),
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=42,
output_tokens=10,
)
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=nonexistent",
model_obj,
100_000,
)
assert "actual_model" not in result
# calculate_cost called with the requested model (fallback)
call_args = mock_calc.call_args
assert call_args[0][0]["model"] == "gpt-oss-120b"
@pytest.mark.asyncio
async def test_old_format_no_model_uses_requested(self) -> None:
"""Old format without model field uses requested model for pricing."""
model_obj = MagicMock()
model_obj.id = "llama3-3-70b"
model_obj.forwarded_model_id = "llama3-3-70b"
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=10,
total_msats=15,
total_usd=0.0,
input_tokens=67,
output_tokens=42,
)
result = await _compute_ehbp_actual_cost(
"prompt=67,completion=42,total=109",
model_obj,
100_000,
)
assert "actual_model" not in result
call_args = mock_calc.call_args
assert call_args[0][0]["model"] == "llama3-3-70b"
@pytest.mark.asyncio
async def test_case_insensitive_model_match(self) -> None:
"""Casing differences between the header and forwarded_model_id
should not trigger a spurious mismatch."""
model_obj = MagicMock()
model_obj.id = "tinfoil-glm-5-2"
model_obj.forwarded_model_id = "glm-5-2" # lowercase
with (
patch("routstr.proxy.get_model_instance") as mock_get_model,
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=42,
output_tokens=10,
)
# Header returns uppercase — same model, different casing
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=GLM-5-2",
model_obj,
100_000,
)
assert "actual_model" not in result
# No mismatch: requested model pricing used
call_args = mock_calc.call_args
assert call_args[0][0]["model"] == "tinfoil-glm-5-2"
mock_get_model.assert_not_called()
@pytest.mark.asyncio
async def test_date_versioned_alias_resolves_to_requested(self) -> None:
"""When the served model is a date-versioned alias that resolves back
to the requested model, no mismatch is propagated."""
model_obj = MagicMock()
model_obj.id = "tinfoil-glm-5-2"
model_obj.forwarded_model_id = "glm-5-2"
resolved_model_obj = MagicMock()
resolved_model_obj.id = "other-provider-glm-5-2"
resolved_model_obj.forwarded_model_id = "glm-5-2"
with (
patch(
"routstr.proxy.get_model_instance",
return_value=resolved_model_obj,
) as mock_get_model,
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=42,
output_tokens=10,
)
# Tinfoil returns a date-versioned ID with different casing.
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=GLM-5-2-20260415",
model_obj,
100_000,
)
# Registry resolution, rather than unconditional suffix removal,
# establishes that this alias represents the expected model.
assert "actual_model" not in result
assert mock_calc.call_args[0][0]["model"] == "tinfoil-glm-5-2"
mock_get_model.assert_called_once_with("GLM-5-2-20260415")
@pytest.mark.asyncio
async def test_configured_date_version_is_preserved_as_identity(self) -> None:
"""A date suffix in forwarded_model_id is meaningful and preserved."""
model_obj = MagicMock()
model_obj.id = "tinfoil-glm-5-2-20260415"
model_obj.forwarded_model_id = "glm-5-2-20260415"
with (
patch("routstr.proxy.get_model_instance") as mock_get_model,
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=42,
output_tokens=10,
)
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=GLM-5-2-20260415",
model_obj,
100_000,
)
assert "actual_model" not in result
assert mock_calc.call_args[0][0]["model"] == "tinfoil-glm-5-2-20260415"
mock_get_model.assert_not_called()
@pytest.mark.asyncio
async def test_different_client_alias_same_upstream_identity(self) -> None:
"""A global alias winner from another provider is not a failover when
its forwarded model ID matches the requested upstream identity."""
model_obj = MagicMock()
model_obj.id = "tinfoil-glm-5-2"
model_obj.forwarded_model_id = "glm-5-2"
resolved_model_obj = MagicMock()
resolved_model_obj.id = "other-provider-glm-5-2"
resolved_model_obj.forwarded_model_id = "GLM-5-2"
with (
patch(
"routstr.proxy.get_model_instance",
return_value=resolved_model_obj,
) as mock_get_model,
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=42,
output_tokens=10,
)
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=provider-alias",
model_obj,
100_000,
)
mock_get_model.assert_called_once_with("provider-alias")
assert "actual_model" not in result
assert mock_calc.call_args[0][0]["model"] == "tinfoil-glm-5-2"
# ---------------------------------------------------------------------------
# TinfoilUpstreamProvider
# ---------------------------------------------------------------------------
class TestTinfoilUpstreamProvider:
def test_provider_type_and_defaults(self) -> None:
assert TinfoilUpstreamProvider.provider_type == "tinfoil"
assert (
TinfoilUpstreamProvider.default_base_url == "https://inference.tinfoil.sh"
)
assert TinfoilUpstreamProvider.supports_ehbp is True
def test_transform_model_name(self) -> None:
provider = TinfoilUpstreamProvider(api_key="test")
assert provider.transform_model_name("tinfoil/llama3-3-70b") == "llama3-3-70b"
assert provider.transform_model_name("llama3-3-70b") == "llama3-3-70b"
def test_get_ehbp_forwarding_target_includes_usage_header(self) -> None:
provider = TinfoilUpstreamProvider(api_key="test")
model_obj = MagicMock()
model_obj.id = "llama3-3-70b"
model_obj.forwarded_model_id = "llama3-3-70b"
target = provider.get_ehbp_forwarding_target("v1/chat/completions", model_obj)
assert target.headers["X-Tinfoil-Request-Usage-Metrics"] == "true"
assert "v1/chat/completions" in target.url
def test_get_provider_metadata(self) -> None:
meta = TinfoilUpstreamProvider.get_provider_metadata()
assert meta["id"] == "tinfoil"
assert meta["name"] == "Tinfoil"
assert meta["fixed_base_url"] is True
def test_tinfoil_model_pricing_parses(self) -> None:
data = {
"id": "llama3-3-70b",
"context_window": 128000,
"created": 1721764788,
"pricing": {
"inputTokenPricePer1M": 1.75,
"outputTokenPricePer1M": 2.75,
"requestPrice": 0,
},
"endpoints": ["/v1/chat/completions"],
"type": "chat",
}
tf = TinfoilModel.parse_obj(data)
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:
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": "llama3-3-70b",
"context_window": 128000,
"created": 1721764788,
"multimodal": False,
"pricing": {
"inputTokenPricePer1M": 1.75,
"outputTokenPricePer1M": 2.75,
"requestPrice": 0,
},
"endpoints": ["/v1/chat/completions"],
"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].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")
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(side_effect=Exception("network error"))
mock_client_cls.return_value = mock_client
models = await provider.fetch_models()
assert models == []
# ---------------------------------------------------------------------------
# EHBP key-config mismatch passthrough
# ---------------------------------------------------------------------------
def _key_config_trailer_response(
status_code: int = 422,
content_type: str = "application/problem+json",
body: bytes = b'{"type":"urn:ietf:params:ehbp:error:key-config","title":"failed to read decrypted request body"}',
) -> TrailerResponse:
return TrailerResponse(
status_code=status_code,
headers=[
("content-type", content_type),
("content-length", str(len(body))),
],
body=body,
trailers=[],
)
class TestIsEhbpKeyConfigResponse:
def test_genuine_key_config_422(self) -> None:
resp = _key_config_trailer_response()
assert _is_ehbp_key_config_response(resp) is True
def test_200_is_not_key_config(self) -> None:
resp = _key_config_trailer_response(status_code=200)
assert _is_ehbp_key_config_response(resp) is False
def test_400_is_not_key_config(self) -> None:
resp = _key_config_trailer_response(status_code=400)
assert _is_ehbp_key_config_response(resp) is False
def test_json_content_type_is_not_key_config(self) -> None:
"""A 422 with application/json (e.g. proxy-wrapped error) must NOT be
treated as key-config — only the original problem+json counts."""
resp = _key_config_trailer_response(content_type="application/json")
assert _is_ehbp_key_config_response(resp) is False
def test_problem_json_with_different_type_is_not_key_config(self) -> None:
"""A 422 problem+json with a different error type is not key-config."""
body = b'{"type":"urn:ietf:params:ehbp:error:other","title":"other"}'
resp = _key_config_trailer_response(body=body)
assert _is_ehbp_key_config_response(resp) is False
def test_empty_body_is_not_key_config(self) -> None:
resp = _key_config_trailer_response(body=b"")
assert _is_ehbp_key_config_response(resp) is False
def test_invalid_json_body_is_not_key_config(self) -> None:
resp = _key_config_trailer_response(body=b"not json")
assert _is_ehbp_key_config_response(resp) is False
def test_problem_json_with_charset(self) -> None:
resp = _key_config_trailer_response(
content_type="application/problem+json; charset=utf-8"
)
assert _is_ehbp_key_config_response(resp) is True
def test_content_type_parameter_disguising_other_media_type(self) -> None:
"""A substring check would accept this; the media type must match
exactly, mirroring the ehbp client's isProblemJSONContentType."""
resp = _key_config_trailer_response(
content_type="text/html; x=application/problem+json"
)
assert _is_ehbp_key_config_response(resp) is False
def test_uppercase_media_type_with_params_matches(self) -> None:
resp = _key_config_trailer_response(
content_type="Application/Problem+JSON; charset=UTF-8"
)
assert _is_ehbp_key_config_response(resp) is True
def test_missing_content_type_is_not_key_config(self) -> None:
resp = TrailerResponse(
status_code=422,
headers=[],
body=b'{"type":"urn:ietf:params:ehbp:error:key-config"}',
)
assert _is_ehbp_key_config_response(resp) is False
class TestPassthroughKeyConfigResponse:
def test_status_and_content_type(self) -> None:
resp = _key_config_trailer_response()
result = _passthrough_key_config_response(resp)
assert result.status_code == 422
assert result.media_type == "application/problem+json"
def test_body_passed_through(self) -> None:
original_body = b'{"type":"urn:ietf:params:ehbp:error:key-config","title":"failed to read decrypted request body"}'
resp = _key_config_trailer_response(body=original_body)
result = _passthrough_key_config_response(resp)
assert result.body == original_body
def test_ehbp_nonce_header_dropped(self) -> None:
"""The nonce must not survive the passthrough: the stock ehbp client
checks for the nonce before the key-config mismatch, so a forwarded
nonce would send it down the decrypt path on this plaintext error
body and the re-attestation loop would never fire."""
resp = TrailerResponse(
status_code=422,
headers=[
("content-type", "application/problem+json"),
("ehbp-response-nonce", "abc123"),
("content-length", "999"),
("server", "nginx"),
("x-request-id", "some-id"),
],
body=b'{"type":"urn:ietf:params:ehbp:error:key-config","title":"test"}',
)
result = _passthrough_key_config_response(resp)
assert "ehbp-response-nonce" not in result.headers
assert "server" not in result.headers
assert "x-request-id" not in result.headers
# Content-length is recomputed from the actual body, not forwarded.
assert result.headers["content-length"] == str(len(resp.body))
# ---------------------------------------------------------------------------
# Key-config passthrough at the forwarding call sites
# ---------------------------------------------------------------------------
def _ehbp_tinfoil_upstream() -> MagicMock:
"""A minimal EHBP-capable upstream stub shaped like the Tinfoil provider."""
upstream = MagicMock()
upstream.provider_type = "tinfoil"
upstream.supports_ehbp = True
upstream.prepare_headers = MagicMock(side_effect=lambda h: h)
upstream.get_confidential_inference_profile = MagicMock(return_value=None)
upstream.get_ehbp_forwarding_target = MagicMock(
return_value=EHBPForwardingTarget(
url="https://inference.tinfoil.sh/private/v1/chat/completions"
)
)
upstream.prepare_params = MagicMock(return_value={})
return upstream
@pytest.mark.asyncio
async def test_bearer_key_config_422_releases_reservation_and_passes_through() -> None:
"""The bearer path returns the enclave's problem+json verbatim AND the
reservation is released.
The early return inside ``forward_ehbp_request`` skips the UpstreamError
handler, so the release depends on the proxy's non-200 branch treating 422
as non-retryable. Nothing else pins that; this does.
"""
key = ApiKey(hashed_key="keyconfig", balance=10_000)
session = MagicMock()
reservation_snapshot = MagicMock()
revert_mock = AsyncMock(return_value=True)
request = MagicMock()
request.method = "POST"
request.headers = {
"authorization": "Bearer sk-keyconfig",
"ehbp-encapsulated-key": "abc123",
"x-routstr-model": "tinfoil/llama3-3-70b",
}
request.body = AsyncMock(return_value=b"sealed-body")
request.query_params = {}
model_obj = MagicMock()
model_obj.id = "tinfoil/llama3-3-70b"
upstream = _ehbp_tinfoil_upstream()
# The enclave may include a nonce even on the 422 — the passthrough must
# drop it, or stock ehbp clients (nonce checked before key-config) would
# try to decrypt this plaintext body instead of re-attesting.
upstream_resp = _key_config_trailer_response()
upstream_resp.headers.append(("ehbp-response-nonce", "nonce-value"))
with (
patch.object(
proxy_module, "get_candidates", return_value=[(model_obj, upstream)]
),
patch.object(
proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000)
),
patch.object(
proxy_module,
"calculate_discounted_max_cost",
AsyncMock(return_value=1_000),
),
patch.object(proxy_module, "check_token_balance", MagicMock()),
patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)),
patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)),
patch.object(
proxy_module,
"get_reservation_snapshot",
AsyncMock(return_value=reservation_snapshot),
),
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
patch(
"routstr.upstream.ehbp.forward_with_trailer",
AsyncMock(return_value=upstream_resp),
),
):
response = await proxy_module.proxy(
request, "v1/chat/completions", session=session
)
# The reservation was released despite the early passthrough return.
revert_mock.assert_awaited_once_with(key, session, 1_000, reservation_snapshot)
# The client receives the enclave's problem+json verbatim...
assert response.status_code == 422
assert response.headers["content-type"] == "application/problem+json"
assert response.body == upstream_resp.body
# ...without the nonce.
assert "ehbp-response-nonce" not in response.headers
@pytest.mark.asyncio
async def test_x_cashu_key_config_422_refunds_and_sets_x_cashu_header() -> None:
"""The x-cashu path refunds the full redeemed amount and attaches the
refund token to the passthrough response."""
request = MagicMock()
request.method = "POST"
request.headers = {
"ehbp-encapsulated-key": "abc123",
"x-routstr-model": "tinfoil/llama3-3-70b",
}
request.query_params = {}
request.body = AsyncMock(return_value=b"sealed-body")
request.state.request_id = "req-1"
model_obj = MagicMock()
model_obj.id = "tinfoil/llama3-3-70b"
upstream = _ehbp_tinfoil_upstream()
upstream_resp = _key_config_trailer_response()
refund_mock = AsyncMock(return_value="cashuArefund")
store_mock = AsyncMock()
with (
patch(
"routstr.upstream.ehbp.recieve_token",
AsyncMock(return_value=(50_000, "msat", "https://mint.example")),
),
patch("routstr.upstream.ehbp.store_cashu_transaction", store_mock),
patch("routstr.upstream.ehbp.send_cashu_refund", refund_mock),
patch(
"routstr.upstream.ehbp.forward_with_trailer",
AsyncMock(return_value=upstream_resp),
),
):
response = await forward_ehbp_x_cashu_request(
request=request,
x_cashu_token="cashuAtoken",
path="v1/chat/completions",
max_cost_for_model=1_000,
model_obj=model_obj,
upstream=upstream,
)
# Full refund of the redeemed amount (the enclave never processed it).
refund_mock.assert_awaited_once_with(
50_000, "msat", "https://mint.example", "req-1"
)
# The redemption itself was recorded.
store_mock.assert_awaited_once()
assert store_mock.await_args is not None
assert store_mock.await_args.kwargs.get("typ") == "in"
# Passthrough shape with the refund attached.
assert response.status_code == 422
assert response.headers["content-type"] == "application/problem+json"
assert response.body == upstream_resp.body
assert response.headers["x-cashu"] == "cashuArefund"