Files
routstr-core/tests/unit/test_tinfoil_integration.py
T

1412 lines
54 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
from .proxy_test_utils import patch_proxy_session
# ---------------------------------------------------------------------------
# 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
assert result["input_source"] == "missing"
assert result["output_source"] == "missing"
@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
assert result["input_source"] == "reported"
assert result["output_source"] == "reported"
assert result["cache_read_source"] == "missing"
assert result["cache_creation_source"] == "missing"
@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
@pytest.mark.parametrize("input_count,output_count", [(0, 0), (10, 5)])
async def test_unpriceable_usage_does_not_charge_authorization_ceiling(
self,
input_count: int,
output_count: int,
) -> 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(
f"prompt={input_count},completion={output_count}",
model_obj,
50_000,
)
assert result["total_msats"] == 0
assert result["input_tokens"] == input_count
assert result["output_tokens"] == output_count
assert result["input_source"] == "reported"
assert result["output_source"] == "reported"
@pytest.mark.asyncio
async def test_partially_malformed_usage_preserves_independent_presence(
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(
"prompt=7,completion=not-a-number",
model_obj,
50_000,
)
assert result["total_msats"] == 0
assert result["input_tokens"] == 7
assert result["output_tokens"] == 0
assert result["input_source"] == "reported"
assert result["output_source"] == "missing"
@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"
actual_model_obj.canonical_slug = "meta/llama-3.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["actual_model_identifier"] == "meta/llama-3.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_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."""
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
assert result["actual_model_unresolved"] is True
# 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=reservation_snapshot),
),
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
patch(
"routstr.upstream.ehbp.forward_with_trailer",
AsyncMock(return_value=upstream_resp),
),
patch_proxy_session(session),
):
response = await proxy_module.proxy(request, "v1/chat/completions")
# 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"