mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
The proxy accepts an API endpoint with or without a leading `v1/`
(`_canonical_api_path`) and forwards the caller's path verbatim, so a
client that spells the endpoint `chat/completions` reaches the upstream
as `chat/completions`. That is fine for providers whose `default_base_url`
carries the version prefix (openai, groq, fireworks, ...) — the client's
`v1/` is stripped and the base URL re-adds its own — but Tinfoil's base URL
is unversioned and its router serves only `/v1/...`. Every Tinfoil request
from such a client therefore got
404 {"error":{"message":"Not found.","type":"invalid_request_error"}}
from `https://inference.tinfoil.sh/chat/completions`, while the same
request with the prefix succeeded. Tinfoil's own error text blamed the
model id, which sent the search in the wrong direction.
Give the EHBP path builders the same hooks the non-EHBP forwarding path
uses: `normalize_request_path` strips the client's optional `v1/`, and a
new `ehbp_path_prefix` re-adds the prefix the provider's enclave actually
serves (`v1` for Tinfoil, `private/v1` for PPQ.AI, whose target had the
same latent bug). Both spellings now reach the same upstream URL.
`_resolve_ehbp_target_url` re-appended the caller's raw path to the
client-supplied enclave URL, re-introducing the spelling the provider had
just normalized away; it now takes the path from the target URL the
provider built, so the override swaps the host only.
1502 lines
58 KiB
Python
1502 lines
58 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",
|
|
{"X-Tinfoil-Enclave-Url": "https://enclave.tinfoil.sh"},
|
|
"tinfoil",
|
|
)
|
|
assert result == "https://enclave.tinfoil.sh/v1/chat/completions"
|
|
|
|
def test_override_keeps_the_providers_version_prefix(self) -> None:
|
|
"""The override swaps the host only, never the path.
|
|
|
|
The provider has already re-added the version prefix the enclave
|
|
requires, so a client that spelled the endpoint ``chat/completions``
|
|
must still land on ``/v1/chat/completions``: this was the path that
|
|
turned every SDK-style request into a paid upstream 404 against
|
|
``router-0.tinfoil.sh``.
|
|
"""
|
|
provider = TinfoilUpstreamProvider(api_key="test")
|
|
model_obj = MagicMock()
|
|
model_obj.id = "tinfoil-deepseek-v4-1-flash"
|
|
model_obj.forwarded_model_id = "deepseek-v4-1-flash"
|
|
target = provider.get_ehbp_forwarding_target("chat/completions", model_obj)
|
|
result = _resolve_ehbp_target_url(
|
|
target.url,
|
|
{"X-Tinfoil-Enclave-Url": "https://router-0.tinfoil.sh"},
|
|
"tinfoil",
|
|
)
|
|
assert result == "https://router-0.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",
|
|
{"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,
|
|
{},
|
|
"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,
|
|
{"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",
|
|
{"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_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
|
|
# 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
|
|
|
|
@pytest.mark.parametrize("client_path", ["v1/chat/completions", "chat/completions"])
|
|
def test_get_ehbp_forwarding_target_is_versioned_for_both_spellings(
|
|
self, client_path: str
|
|
) -> None:
|
|
"""The node accepts ``chat/completions`` and ``v1/chat/completions`` as
|
|
the same endpoint, so both must reach Tinfoil's versioned route.
|
|
|
|
A bare client path used to produce
|
|
``https://inference.tinfoil.sh/chat/completions``, which the router
|
|
answers with 404 ``{"error":{"message":"Not found."}}``.
|
|
"""
|
|
provider = TinfoilUpstreamProvider(api_key="test")
|
|
model_obj = MagicMock()
|
|
model_obj.id = "tinfoil-deepseek-v4-1-flash"
|
|
model_obj.forwarded_model_id = "deepseek-v4-1-flash"
|
|
target = provider.get_ehbp_forwarding_target(client_path, model_obj)
|
|
assert target.url == "https://inference.tinfoil.sh/v1/chat/completions"
|
|
|
|
def test_get_ehbp_forwarding_target_does_not_double_the_prefix(self) -> None:
|
|
"""A path that already carries ``v1/`` is normalized, not prefixed again."""
|
|
provider = TinfoilUpstreamProvider(api_key="test")
|
|
model_obj = MagicMock()
|
|
assert (
|
|
provider.build_ehbp_request_path("v1/chat/completions", model_obj)
|
|
== "v1/chat/completions"
|
|
)
|
|
assert (
|
|
provider.build_ehbp_request_path("chat/completions", model_obj)
|
|
== "v1/chat/completions"
|
|
)
|
|
|
|
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"
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# The forwarded upstream path (regression: bare /chat/completions -> 404)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("client_path", ["v1/chat/completions", "chat/completions"])
|
|
async def test_proxy_reaches_the_enclaves_versioned_route_for_both_spellings(
|
|
client_path: str,
|
|
) -> None:
|
|
"""The caller's spelling must not decide which upstream path is used.
|
|
|
|
The proxy treats ``chat/completions`` and ``v1/chat/completions`` as the
|
|
same endpoint (``_canonical_api_path``) and forwards the caller's path
|
|
verbatim, so the version prefix the enclave requires has to be re-added
|
|
downstream by the provider. Before that, a client that spelled the endpoint
|
|
without ``v1/`` reached ``https://inference.tinfoil.sh/chat/completions``
|
|
and got a 404 ``{"error":{"message":"Not found."}}`` from the router for
|
|
every Tinfoil model, while the same request with the prefix succeeded.
|
|
"""
|
|
key = ApiKey(hashed_key="keyconfig", balance=10_000)
|
|
session = MagicMock()
|
|
reservation_snapshot = MagicMock()
|
|
|
|
request = MagicMock()
|
|
request.method = "POST"
|
|
request.headers = {
|
|
"authorization": "Bearer sk-keyconfig",
|
|
"ehbp-encapsulated-key": "abc123",
|
|
"x-routstr-model": "tinfoil-deepseek-v4-1-flash",
|
|
}
|
|
request.body = AsyncMock(return_value=b"sealed-body")
|
|
request.query_params = {}
|
|
|
|
model_obj = MagicMock()
|
|
model_obj.id = "tinfoil-deepseek-v4-1-flash"
|
|
model_obj.forwarded_model_id = "deepseek-v4-1-flash"
|
|
upstream = TinfoilUpstreamProvider(api_key="upstream-key")
|
|
|
|
# A 422 key-config response short-circuits the billing path while still
|
|
# exercising the URL the request was actually sent to.
|
|
forward_mock = AsyncMock(return_value=_key_config_trailer_response())
|
|
|
|
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", AsyncMock(return_value=True)
|
|
),
|
|
patch("routstr.upstream.ehbp.forward_with_trailer", forward_mock),
|
|
patch_proxy_session(session),
|
|
):
|
|
await proxy_module.proxy(request, client_path)
|
|
|
|
assert forward_mock.await_args is not None
|
|
assert (
|
|
forward_mock.await_args.kwargs["url"]
|
|
== "https://inference.tinfoil.sh/v1/chat/completions"
|
|
)
|