From 623c55227c9343670902e6b1a52cfec558b8da9a Mon Sep 17 00:00:00 2001 From: redshift <213178690+1ftredsh@users.noreply.github.com> Date: Tue, 8 Sep 2026 13:19:32 +0200 Subject: [PATCH] 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. --- routstr/upstream/ehbp.py | 20 +- tests/unit/test_tinfoil_integration.py | 295 ++++++++++++++++++++----- 2 files changed, 250 insertions(+), 65 deletions(-) diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index b2822366..b15d5bad 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -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", ) diff --git a/tests/unit/test_tinfoil_integration.py b/tests/unit/test_tinfoil_integration.py index b7d52dfe..1058a19a 100644 --- a/tests/unit/test_tinfoil_integration.py +++ b/tests/unit/test_tinfoil_integration.py @@ -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"