mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fix(ehbp): drop nonce on key-config passthrough, exact media type, call-site tests
Address review feedback on the key-config 422 passthrough: - Drop Ehbp-Response-Nonce (and content-length) from the passthrough. A nonce only carries meaning on an encrypted response body, and the stock ehbp client (shouldDecryptResponse) checks the nonce BEFORE the key-config mismatch -- a forwarded nonce would push that client into the decrypt path on the plaintext error body, so re-attestation would never fire. routstr-sdk checks key-config first, but the passthrough must stay correct for any EHBP client. Content-length is recomputed from the body instead of forwarded. - Match the media type exactly (split on ';') like the ehbp client's isProblemJSONContentType, instead of a substring check that also accepted e.g. 'text/html; x=application/problem+json'. - Add call-site tests: the bearer path (driven through the proxy handler) pins that a key-config 422 passes through as problem+json AND releases the reservation (the early return skips the UpstreamError handler, so the release depends on 422 remaining non-retryable); the x-cashu path pins the full refund and the X-Cashu header on the passthrough response. - Apply ruff format to the touched test file.
This commit is contained in:
+10
-10
@@ -81,7 +81,8 @@ def _is_ehbp_key_config_response(resp: TrailerResponse) -> bool:
|
||||
if k.lower() == "content-type":
|
||||
ct = v.lower()
|
||||
break
|
||||
if "application/problem+json" not in ct:
|
||||
media_type = ct.split(";", 1)[0].strip()
|
||||
if media_type != "application/problem+json":
|
||||
return False
|
||||
try:
|
||||
body = json.loads(resp.body)
|
||||
@@ -94,19 +95,18 @@ def _passthrough_key_config_response(resp: TrailerResponse) -> Response:
|
||||
"""Return the enclave's key-config 422 with its original body and content
|
||||
type so the EHBP client's ``KeyConfigMismatchError`` detection fires.
|
||||
|
||||
Only EHBP protocol headers are forwarded; everything else (hop-by-hop,
|
||||
upstream-internal) is filtered out.
|
||||
Only the content type is forwarded. ``Ehbp-Response-Nonce`` must be
|
||||
dropped: a nonce only carries meaning for an *encrypted* response body,
|
||||
and the stock ``ehbp`` client (``shouldDecryptResponse``) checks for the
|
||||
nonce *before* checking for a key-config mismatch — forwarding it would
|
||||
send that client down the decrypt path on this plaintext error body, so
|
||||
the re-attestation loop would never fire. Content-length is recomputed
|
||||
from the body, and upstream-internal headers are filtered out.
|
||||
"""
|
||||
passthrough_headers: dict[str, str] = {
|
||||
"content-type": "application/problem+json",
|
||||
}
|
||||
for k, v in resp.headers:
|
||||
if k.lower() in ("ehbp-response-nonce", "content-length"):
|
||||
passthrough_headers[k] = v
|
||||
return Response(
|
||||
content=resp.body,
|
||||
status_code=422,
|
||||
headers=passthrough_headers,
|
||||
headers={"content-type": "application/problem+json"},
|
||||
media_type="application/problem+json",
|
||||
)
|
||||
|
||||
|
||||
@@ -11,14 +11,18 @@ 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 (
|
||||
@@ -109,9 +113,7 @@ class TestParseTinfoilUsageMetrics:
|
||||
|
||||
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"
|
||||
)
|
||||
result = parse_tinfoil_usage_metrics("prompt=67,completion=42,total=109")
|
||||
assert result == {
|
||||
"prompt_tokens": 67,
|
||||
"completion_tokens": 42,
|
||||
@@ -288,7 +290,9 @@ class TestComputeEhbpActualCost:
|
||||
assert result["output_msats"] == 20
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unpriceable_usage_does_not_charge_authorization_ceiling(self) -> None:
|
||||
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"
|
||||
@@ -393,13 +397,16 @@ class TestComputeEhbpActualCost:
|
||||
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:
|
||||
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(
|
||||
@@ -430,13 +437,16 @@ class TestComputeEhbpActualCost:
|
||||
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:
|
||||
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(
|
||||
@@ -495,12 +505,13 @@ class TestComputeEhbpActualCost:
|
||||
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:
|
||||
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(
|
||||
@@ -536,13 +547,16 @@ class TestComputeEhbpActualCost:
|
||||
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:
|
||||
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(
|
||||
@@ -573,12 +587,13 @@ class TestComputeEhbpActualCost:
|
||||
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:
|
||||
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(
|
||||
@@ -597,10 +612,7 @@ class TestComputeEhbpActualCost:
|
||||
)
|
||||
|
||||
assert "actual_model" not in result
|
||||
assert (
|
||||
mock_calc.call_args[0][0]["model"]
|
||||
== "tinfoil-glm-5-2-20260415"
|
||||
)
|
||||
assert mock_calc.call_args[0][0]["model"] == "tinfoil-glm-5-2-20260415"
|
||||
mock_get_model.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -615,13 +627,16 @@ class TestComputeEhbpActualCost:
|
||||
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:
|
||||
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(
|
||||
@@ -653,8 +668,7 @@ class TestTinfoilUpstreamProvider:
|
||||
def test_provider_type_and_defaults(self) -> None:
|
||||
assert TinfoilUpstreamProvider.provider_type == "tinfoil"
|
||||
assert (
|
||||
TinfoilUpstreamProvider.default_base_url
|
||||
== "https://inference.tinfoil.sh"
|
||||
TinfoilUpstreamProvider.default_base_url == "https://inference.tinfoil.sh"
|
||||
)
|
||||
assert TinfoilUpstreamProvider.supports_ehbp is True
|
||||
|
||||
@@ -669,9 +683,7 @@ class TestTinfoilUpstreamProvider:
|
||||
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 target.headers["X-Tinfoil-Request-Usage-Metrics"] == "true"
|
||||
assert "v1/chat/completions" in target.url
|
||||
|
||||
def test_get_provider_metadata(self) -> None:
|
||||
@@ -812,6 +824,20 @@ class TestIsEhbpKeyConfigResponse:
|
||||
)
|
||||
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,
|
||||
@@ -834,18 +860,177 @@ class TestPassthroughKeyConfigResponse:
|
||||
result = _passthrough_key_config_response(resp)
|
||||
assert result.body == original_body
|
||||
|
||||
def test_ehbp_nonce_header_forwarded(self) -> None:
|
||||
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 result.headers["ehbp-response-nonce"] == "abc123"
|
||||
assert "ehbp-response-nonce" not in result.headers
|
||||
assert "server" not in result.headers
|
||||
assert "x-request-id" not in result.headers
|
||||
# Content-length is recomputed from the actual body, not forwarded.
|
||||
assert result.headers["content-length"] == str(len(resp.body))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Key-config passthrough at the forwarding call sites
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _ehbp_tinfoil_upstream() -> MagicMock:
|
||||
"""A minimal EHBP-capable upstream stub shaped like the Tinfoil provider."""
|
||||
upstream = MagicMock()
|
||||
upstream.provider_type = "tinfoil"
|
||||
upstream.supports_ehbp = True
|
||||
upstream.prepare_headers = MagicMock(side_effect=lambda h: h)
|
||||
upstream.get_confidential_inference_profile = MagicMock(return_value=None)
|
||||
upstream.get_ehbp_forwarding_target = MagicMock(
|
||||
return_value=EHBPForwardingTarget(
|
||||
url="https://inference.tinfoil.sh/private/v1/chat/completions"
|
||||
)
|
||||
)
|
||||
upstream.prepare_params = MagicMock(return_value={})
|
||||
return upstream
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bearer_key_config_422_releases_reservation_and_passes_through() -> None:
|
||||
"""The bearer path returns the enclave's problem+json verbatim AND the
|
||||
reservation is released.
|
||||
|
||||
The early return inside ``forward_ehbp_request`` skips the UpstreamError
|
||||
handler, so the release depends on the proxy's non-200 branch treating 422
|
||||
as non-retryable. Nothing else pins that; this does.
|
||||
"""
|
||||
key = ApiKey(hashed_key="keyconfig", balance=10_000)
|
||||
session = MagicMock()
|
||||
reservation_snapshot = MagicMock()
|
||||
revert_mock = AsyncMock(return_value=True)
|
||||
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = {
|
||||
"authorization": "Bearer sk-keyconfig",
|
||||
"ehbp-encapsulated-key": "abc123",
|
||||
"x-routstr-model": "tinfoil/llama3-3-70b",
|
||||
}
|
||||
request.body = AsyncMock(return_value=b"sealed-body")
|
||||
request.query_params = {}
|
||||
|
||||
model_obj = MagicMock()
|
||||
model_obj.id = "tinfoil/llama3-3-70b"
|
||||
upstream = _ehbp_tinfoil_upstream()
|
||||
|
||||
# The enclave may include a nonce even on the 422 — the passthrough must
|
||||
# drop it, or stock ehbp clients (nonce checked before key-config) would
|
||||
# try to decrypt this plaintext body instead of re-attesting.
|
||||
upstream_resp = _key_config_trailer_response()
|
||||
upstream_resp.headers.append(("ehbp-response-nonce", "nonce-value"))
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
proxy_module, "get_candidates", return_value=[(model_obj, upstream)]
|
||||
),
|
||||
patch.object(
|
||||
proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000)
|
||||
),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
"calculate_discounted_max_cost",
|
||||
AsyncMock(return_value=1_000),
|
||||
),
|
||||
patch.object(proxy_module, "check_token_balance", MagicMock()),
|
||||
patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)),
|
||||
patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
"get_reservation_snapshot",
|
||||
AsyncMock(return_value=reservation_snapshot),
|
||||
),
|
||||
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
|
||||
patch(
|
||||
"routstr.upstream.ehbp.forward_with_trailer",
|
||||
AsyncMock(return_value=upstream_resp),
|
||||
),
|
||||
):
|
||||
response = await proxy_module.proxy(
|
||||
request, "v1/chat/completions", session=session
|
||||
)
|
||||
|
||||
# The reservation was released despite the early passthrough return.
|
||||
revert_mock.assert_awaited_once_with(key, session, 1_000, reservation_snapshot)
|
||||
# The client receives the enclave's problem+json verbatim...
|
||||
assert response.status_code == 422
|
||||
assert response.headers["content-type"] == "application/problem+json"
|
||||
assert response.body == upstream_resp.body
|
||||
# ...without the nonce.
|
||||
assert "ehbp-response-nonce" not in response.headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_x_cashu_key_config_422_refunds_and_sets_x_cashu_header() -> None:
|
||||
"""The x-cashu path refunds the full redeemed amount and attaches the
|
||||
refund token to the passthrough response."""
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = {
|
||||
"ehbp-encapsulated-key": "abc123",
|
||||
"x-routstr-model": "tinfoil/llama3-3-70b",
|
||||
}
|
||||
request.query_params = {}
|
||||
request.body = AsyncMock(return_value=b"sealed-body")
|
||||
request.state.request_id = "req-1"
|
||||
|
||||
model_obj = MagicMock()
|
||||
model_obj.id = "tinfoil/llama3-3-70b"
|
||||
upstream = _ehbp_tinfoil_upstream()
|
||||
|
||||
upstream_resp = _key_config_trailer_response()
|
||||
refund_mock = AsyncMock(return_value="cashuArefund")
|
||||
store_mock = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.ehbp.recieve_token",
|
||||
AsyncMock(return_value=(50_000, "msat", "https://mint.example")),
|
||||
),
|
||||
patch("routstr.upstream.ehbp.store_cashu_transaction", store_mock),
|
||||
patch("routstr.upstream.ehbp.send_cashu_refund", refund_mock),
|
||||
patch(
|
||||
"routstr.upstream.ehbp.forward_with_trailer",
|
||||
AsyncMock(return_value=upstream_resp),
|
||||
),
|
||||
):
|
||||
response = await forward_ehbp_x_cashu_request(
|
||||
request=request,
|
||||
x_cashu_token="cashuAtoken",
|
||||
path="v1/chat/completions",
|
||||
max_cost_for_model=1_000,
|
||||
model_obj=model_obj,
|
||||
upstream=upstream,
|
||||
)
|
||||
|
||||
# Full refund of the redeemed amount (the enclave never processed it).
|
||||
refund_mock.assert_awaited_once_with(
|
||||
50_000, "msat", "https://mint.example", "req-1"
|
||||
)
|
||||
# The redemption itself was recorded.
|
||||
store_mock.assert_awaited_once()
|
||||
assert store_mock.await_args.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"
|
||||
|
||||
Reference in New Issue
Block a user