diff --git a/routstr/core/exceptions.py b/routstr/core/exceptions.py index 360b810d..6e7d90a0 100644 --- a/routstr/core/exceptions.py +++ b/routstr/core/exceptions.py @@ -34,6 +34,22 @@ class UpstreamError(Exception): super().__init__(message) +class EhbpTimeoutError(UpstreamError): + """Raised when an EHBP upstream times out waiting for a response. + + Distinct from a generic :class:`UpstreamError` so callers can map the + failure to a ``504 Gateway Timeout`` with a stable ``UPSTREAM_TIMEOUT`` + code instead of a misleading ``500`` internal server error. + """ + + def __init__(self, message: str, status_code: int = 504): + super().__init__( + message, + status_code=status_code, + code="UPSTREAM_TIMEOUT", + ) + + async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse: """Handle HTTP exceptions and include request ID in response.""" request_id = getattr(request.state, "request_id", "unknown") diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index d6b7e7d2..b135c238 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -32,7 +32,7 @@ from ..core.db import ( from ..core.db import ( store_cashu_transaction_with_retry as store_cashu_transaction, ) -from ..core.exceptions import UpstreamError +from ..core.exceptions import EhbpTimeoutError, UpstreamError from ..core.settings import settings from ..payment.cost_calculation import ( CostData, @@ -1042,6 +1042,46 @@ async def forward_ehbp_x_cashu_request( except Exception: raise + except EhbpTimeoutError as e: + logger.warning( + "EHBP X-Cashu upstream timed out", + extra={ + "error": str(e), + "path": path, + "method": request.method, + "redeemed": redeemed, + }, + ) + + if redeemed and amount > 0: + try: + refund_token = await send_cashu_refund(amount, unit, mint, request_id) + error_response = create_error_response( + "upstream_timeout", + str(e), + 504, + request=request, + code="UPSTREAM_TIMEOUT", + ) + error_response.headers["X-Cashu"] = refund_token + return error_response + except Exception as refund_error: + logger.error( + "Failed to refund EHBP X-Cashu token after timeout", + extra={ + "error": str(refund_error), + "original_error": str(e), + }, + ) + + return create_error_response( + "upstream_timeout", + str(e), + 504, + request=request, + code="UPSTREAM_TIMEOUT", + ) + except Exception as e: error_message = str(e) logger.error( diff --git a/routstr/upstream/tinfoil_trailer.py b/routstr/upstream/tinfoil_trailer.py index 0357864f..ec3065d6 100644 --- a/routstr/upstream/tinfoil_trailer.py +++ b/routstr/upstream/tinfoil_trailer.py @@ -20,11 +20,12 @@ from urllib.parse import urlsplit import h11 from ..core import get_logger +from ..core.exceptions import EhbpTimeoutError logger = get_logger(__name__) _READ_BUFSIZE = 65536 -_DEFAULT_TIMEOUT_SECONDS = 30.0 +_DEFAULT_TIMEOUT_SECONDS = 60.0 _DEFAULT_CLOSE_TIMEOUT_SECONDS = 1.0 _DEFAULT_MAX_RESPONSE_BYTES = 25 * 1024 * 1024 _HOP_BY_HOP_HEADERS = { @@ -100,10 +101,15 @@ async def forward_with_trailer( headers = _strip_hop_by_hop_headers(headers) ssl_ctx = ssl.create_default_context() - reader, writer = await asyncio.wait_for( - asyncio.open_connection(host, port, ssl=ssl_ctx), - timeout=timeout_seconds, - ) + try: + reader, writer = await asyncio.wait_for( + asyncio.open_connection(host, port, ssl=ssl_ctx), + timeout=timeout_seconds, + ) + except asyncio.TimeoutError as exc: + raise EhbpTimeoutError( + f"EHBP upstream {host} timed out after {timeout_seconds:g}s connecting" + ) from exc try: # Build HTTP/1.1 request @@ -126,7 +132,13 @@ async def forward_with_trailer( request_data += body writer.write(request_data) - await asyncio.wait_for(writer.drain(), timeout=timeout_seconds) + try: + await asyncio.wait_for(writer.drain(), timeout=timeout_seconds) + except asyncio.TimeoutError as exc: + raise EhbpTimeoutError( + f"EHBP upstream {host} timed out after " + f"{timeout_seconds:g}s sending request" + ) from exc # Parse response with h11 conn = h11.Connection(h11.CLIENT) @@ -140,10 +152,16 @@ async def forward_with_trailer( event = conn.next_event() if event is h11.NEED_DATA: - data = await asyncio.wait_for( - reader.read(_READ_BUFSIZE), - timeout=timeout_seconds, - ) + try: + data = await asyncio.wait_for( + reader.read(_READ_BUFSIZE), + timeout=timeout_seconds, + ) + except asyncio.TimeoutError as exc: + raise EhbpTimeoutError( + f"EHBP upstream {host} timed out after " + f"{timeout_seconds:g}s waiting for response data" + ) from exc conn.receive_data(data if data else b"") continue diff --git a/tests/unit/test_ehbp_timeout.py b/tests/unit/test_ehbp_timeout.py new file mode 100644 index 00000000..bd19871d --- /dev/null +++ b/tests/unit/test_ehbp_timeout.py @@ -0,0 +1,80 @@ +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from routstr.core.exceptions import EhbpTimeoutError +from routstr.upstream import ehbp as ehbp_module + +# --------------------------------------------------------------------------- +# forward_ehbp_x_cashu_request — timeout fails closed with a refund + 504 +# --------------------------------------------------------------------------- + + +async def _request() -> MagicMock: + request = MagicMock() + request.state.request_id = "req-123" + request.method = "POST" + request.query_params = {} + request.headers = {} + request.body = AsyncMock(return_value=b"opaque") + return request + + +@pytest.mark.asyncio +async def test_x_cashu_timeout_refunds_and_returns_504( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + ehbp_module, + "recieve_token", + AsyncMock(return_value=(1000, "msat", None)), + ) + monkeypatch.setattr( + ehbp_module, "store_cashu_transaction", AsyncMock(return_value=None) + ) + monkeypatch.setattr( + ehbp_module, "send_cashu_refund", AsyncMock(return_value="refund-token") + ) + monkeypatch.setattr( + ehbp_module, + "forward_with_trailer", + AsyncMock(side_effect=EhbpTimeoutError("EHBP upstream timed out")), + ) + + profile = MagicMock() + profile.client_target_url_header = None + profile.allow_client_target_override = False + profile.proxy_only_headers = frozenset() + profile.usage_response_header = None + + target = MagicMock() + target.url = "https://inference.tinfoil.sh/v1/chat/completions" + target.headers = {} + target.profile = None + + upstream = MagicMock() + upstream.prepare_headers.return_value = {} + upstream.get_ehbp_forwarding_target.return_value = target + upstream.get_confidential_inference_profile.return_value = profile + upstream.prepare_params.return_value = {} + + model_obj = MagicMock() + model_obj.id = "tinfoil-kimi-k2-6" + model_obj.forwarded_model_id = "kimi-k2-6" + + response = await ehbp_module.forward_ehbp_x_cashu_request( + request=await _request(), + x_cashu_token="cashu-token", + path="v1/chat/completions", + max_cost_for_model=5000, + model_obj=model_obj, + upstream=upstream, + ) + + assert response.status_code == 504 + assert response.headers["X-Cashu"] == "refund-token" + ehbp_module.send_cashu_refund.assert_awaited_once_with( + 1000, "msat", None, "req-123" + ) diff --git a/tests/unit/test_tinfoil_trailer.py b/tests/unit/test_tinfoil_trailer.py index ef1c96f1..5a6d68a7 100644 --- a/tests/unit/test_tinfoil_trailer.py +++ b/tests/unit/test_tinfoil_trailer.py @@ -1,9 +1,11 @@ from __future__ import annotations +import asyncio from unittest.mock import AsyncMock, MagicMock import pytest +from routstr.core.exceptions import EhbpTimeoutError, UpstreamError from routstr.upstream.tinfoil_trailer import forward_with_trailer @@ -28,6 +30,14 @@ class FakeWriter: self.written += data +class HangingReader: + """A reader that never returns data, used to trigger a read timeout.""" + + async def read(self, _size: int) -> bytes: + await asyncio.sleep(3600) + return b"" + + @pytest.mark.asyncio async def test_forward_with_trailer_captures_usage_trailer( monkeypatch: pytest.MonkeyPatch, @@ -136,3 +146,53 @@ async def test_forward_with_trailer_enforces_response_size_limit( ) writer.close.assert_called_once() + + +@pytest.mark.asyncio +async def test_forward_with_trailer_connect_timeout_raises_ehbp_timeout( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def _hang_connect(*_args: object, **_kwargs: object) -> object: + raise asyncio.TimeoutError + + monkeypatch.setattr( + "routstr.upstream.tinfoil_trailer.asyncio.open_connection", _hang_connect + ) + + with pytest.raises(EhbpTimeoutError, match="connecting"): + await forward_with_trailer( + method="POST", + url="https://enclave.tinfoil.sh/v1/chat/completions", + headers={}, + body=b"opaque", + ) + + +@pytest.mark.asyncio +async def test_forward_with_trailer_read_timeout_raises_ehbp_timeout( + monkeypatch: pytest.MonkeyPatch, +) -> None: + reader = HangingReader() + writer = FakeWriter() + monkeypatch.setattr( + "routstr.upstream.tinfoil_trailer.asyncio.open_connection", + AsyncMock(return_value=(reader, writer)), + ) + + with pytest.raises(EhbpTimeoutError, match="waiting for response data"): + await forward_with_trailer( + method="POST", + url="https://enclave.tinfoil.sh/v1/chat/completions", + headers={}, + body=b"opaque", + timeout_seconds=0.01, + ) + + writer.close.assert_called_once() + + +def test_ehbp_timeout_error_metadata() -> None: + exc = EhbpTimeoutError("boom") + assert exc.status_code == 504 + assert exc.code == "UPSTREAM_TIMEOUT" + assert isinstance(exc, UpstreamError)