diff --git a/routstr/core/exceptions.py b/routstr/core/exceptions.py index d82bb190..ce6ac46e 100644 --- a/routstr/core/exceptions.py +++ b/routstr/core/exceptions.py @@ -5,7 +5,11 @@ from fastapi.encoders import jsonable_encoder from fastapi.exceptions import RequestValidationError from fastapi.responses import JSONResponse -from .error_scope import ERROR_SCOPE_UPSTREAM, UPSTREAM_ERROR_STATUS +from .error_scope import ( + ERROR_SCOPE_UPSTREAM, + UPSTREAM_ERROR_STATUS, + UPSTREAM_UNAVAILABLE, +) from .logging import get_logger logger = get_logger(__name__) @@ -73,6 +77,28 @@ class EhbpTimeoutError(UpstreamError): ) +class EhbpConnectionError(UpstreamError): + """Raised when an EHBP upstream cannot be reached. + + Covers transport failures while establishing the provider connection: DNS + resolution, TCP refused/reset, or a TLS error that is not a handshake + timeout. Distinct from a generic :class:`UpstreamError` so the failure is + attributed to the provider hop (``UPSTREAM_UNAVAILABLE``, reported as + ``424``) instead of being flattened into a misleading node-scoped ``500``. + + ``details`` carries optional structured, redaction-safe context and is + forwarded to the client by ``create_upstream_error_response``. + """ + + def __init__(self, message: str, details: dict[str, object] | None = None): + super().__init__( + message, + status_code=UPSTREAM_ERROR_STATUS, + code=UPSTREAM_UNAVAILABLE, + details=details, + ) + + def _error_message_from_detail(detail: object) -> str | None: """Extract a message from an HTTPException ``detail``, capped at 200 chars.""" if isinstance(detail, dict): diff --git a/routstr/upstream/tinfoil_trailer.py b/routstr/upstream/tinfoil_trailer.py index 1b4bb60d..2246a250 100644 --- a/routstr/upstream/tinfoil_trailer.py +++ b/routstr/upstream/tinfoil_trailer.py @@ -20,7 +20,7 @@ from urllib.parse import urlsplit import h11 from ..core import get_logger -from ..core.exceptions import EhbpTimeoutError +from ..core.exceptions import EhbpConnectionError, EhbpTimeoutError logger = get_logger(__name__) @@ -110,6 +110,21 @@ async def forward_with_trailer( raise EhbpTimeoutError( f"EHBP upstream {host} timed out after {timeout_seconds:g}s connecting" ) from exc + except ConnectionAbortedError as exc: + # CPython's ssl module aborts a TLS handshake that outlives its + # internal timer ("SSL handshake is taking longer than N seconds") + # with ConnectionAbortedError. That is a connect timeout on the + # provider hop, not a local node fault, so surface it as a timeout. + raise EhbpTimeoutError( + f"EHBP upstream {host} TLS handshake timed out while connecting" + ) from exc + except (ssl.SSLError, ConnectionError, OSError) as exc: + # DNS failure, connection refused/reset, or a non-timeout TLS error: + # the provider could not be reached. Attribute it to the upstream hop + # rather than letting it become a node-scoped 500. + raise EhbpConnectionError( + f"Unable to connect to EHBP upstream {host}: {type(exc).__name__}" + ) from exc try: # Build HTTP/1.1 request diff --git a/tests/unit/test_ehbp_timeout.py b/tests/unit/test_ehbp_timeout.py index 97a7b7ee..8ec300c2 100644 --- a/tests/unit/test_ehbp_timeout.py +++ b/tests/unit/test_ehbp_timeout.py @@ -10,7 +10,11 @@ from routstr.core.error_scope import ( ERROR_SCOPE_UPSTREAM, UPSTREAM_ERROR_STATUS, ) -from routstr.core.exceptions import EhbpTimeoutError, UpstreamError +from routstr.core.exceptions import ( + EhbpConnectionError, + EhbpTimeoutError, + UpstreamError, +) from routstr.upstream import ehbp as ehbp_module # --------------------------------------------------------------------------- @@ -141,3 +145,47 @@ async def test_bearer_timeout_propagates_424( assert exc_info.value.code == "UPSTREAM_TIMEOUT" assert exc_info.value.scope == ERROR_SCOPE_UPSTREAM assert isinstance(exc_info.value, UpstreamError) + + +@pytest.mark.asyncio +async def test_bearer_connection_error_propagates_upstream_scope( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A connect failure must stay upstream-scoped instead of becoming a 500. + + ``forward_with_trailer`` classifies TLS/connection failures as + :class:`EhbpConnectionError`; ``forward_ehbp_request``'s ``except + UpstreamError: raise`` must let it through so ``proxy.py`` answers 424 with + the upstream scope header rather than a node-scoped 500. + """ + monkeypatch.setattr( + ehbp_module, + "forward_with_trailer", + AsyncMock( + side_effect=EhbpConnectionError( + "Unable to connect to EHBP upstream inference.tinfoil.sh: " + "ConnectionAbortedError" + ) + ), + ) + upstream, model_obj = _ehbp_upstream_mocks() + key = MagicMock() + key.hashed_key = "abcdef1234567890" + + with pytest.raises(EhbpConnectionError) as exc_info: + await ehbp_module.forward_ehbp_request( + request=await _request(), + path="v1/chat/completions", + headers={}, + request_body=b"opaque", + upstream=upstream, + key=key, + max_cost_for_model=5000, + session=MagicMock(), + model_obj=model_obj, + ) + + assert exc_info.value.status_code == UPSTREAM_ERROR_STATUS + assert exc_info.value.code == "UPSTREAM_UNAVAILABLE" + assert exc_info.value.scope == ERROR_SCOPE_UPSTREAM + assert isinstance(exc_info.value, UpstreamError) diff --git a/tests/unit/test_tinfoil_trailer.py b/tests/unit/test_tinfoil_trailer.py index 2b3d22f4..12107aca 100644 --- a/tests/unit/test_tinfoil_trailer.py +++ b/tests/unit/test_tinfoil_trailer.py @@ -1,12 +1,13 @@ from __future__ import annotations import asyncio +import ssl from unittest.mock import AsyncMock, MagicMock import pytest from routstr.core.error_scope import ERROR_SCOPE_UPSTREAM, UPSTREAM_ERROR_STATUS -from routstr.core.exceptions import EhbpTimeoutError, UpstreamError +from routstr.core.exceptions import EhbpConnectionError, EhbpTimeoutError, UpstreamError from routstr.upstream.tinfoil_trailer import forward_with_trailer @@ -169,6 +170,68 @@ async def test_forward_with_trailer_connect_timeout_raises_ehbp_timeout( ) +@pytest.mark.asyncio +async def test_forward_with_trailer_tls_handshake_timeout_raises_ehbp_timeout( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The stdlib TLS handshake timer surfaces as ConnectionAbortedError. + + CPython aborts a slow handshake with ``ConnectionAbortedError`` rather + than ``asyncio.TimeoutError``, so the connect handler must classify it as + an upstream timeout — otherwise it escapes to the node-scoped 500 in + ``forward_ehbp_request``. + """ + + async def _handshake_timeout(*_args: object, **_kwargs: object) -> object: + raise ConnectionAbortedError( + "SSL handshake is taking longer than 60.0 seconds: aborting the connection" + ) + + monkeypatch.setattr( + "routstr.upstream.tinfoil_trailer.asyncio.open_connection", + _handshake_timeout, + ) + + with pytest.raises(EhbpTimeoutError, match="TLS handshake timed out"): + await forward_with_trailer( + method="POST", + url="https://inference.tinfoil.sh/v1/chat/completions", + headers={}, + body=b"opaque", + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "exc", + [ + ConnectionRefusedError("connection refused"), + ConnectionResetError("connection reset"), + ssl.SSLError("certificate verify failed"), + OSError("name resolution failed"), + ], +) +async def test_forward_with_trailer_connection_failure_raises_ehbp_connection( + monkeypatch: pytest.MonkeyPatch, exc: Exception +) -> None: + """Non-timeout connect failures must be upstream-scoped, not node 500s.""" + + async def _fail_connect(*_args: object, **_kwargs: object) -> object: + raise exc + + monkeypatch.setattr( + "routstr.upstream.tinfoil_trailer.asyncio.open_connection", _fail_connect + ) + + with pytest.raises(EhbpConnectionError, match="Unable to connect"): + await forward_with_trailer( + method="POST", + url="https://inference.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, @@ -207,3 +270,18 @@ def test_ehbp_timeout_error_forwards_details() -> None: assert exc.details == {"phase": "connect"} assert exc.status_code == UPSTREAM_ERROR_STATUS assert exc.code == "UPSTREAM_TIMEOUT" + + +def test_ehbp_connection_error_metadata() -> None: + exc = EhbpConnectionError("boom") + assert exc.status_code == UPSTREAM_ERROR_STATUS + assert exc.code == "UPSTREAM_UNAVAILABLE" + assert exc.details is None + assert exc.scope == ERROR_SCOPE_UPSTREAM + assert isinstance(exc, UpstreamError) + + +def test_ehbp_connection_error_forwards_details() -> None: + exc = EhbpConnectionError("boom", details={"provider": "tinfoil"}) + assert exc.details == {"provider": "tinfoil"} + assert exc.code == "UPSTREAM_UNAVAILABLE"