diff --git a/routstr/proxy.py b/routstr/proxy.py index 9d2dbf09..ad301a45 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -1,8 +1,9 @@ import json +from pathlib import Path from typing import Any from fastapi import APIRouter, Depends, HTTPException, Request -from fastapi.responses import Response, StreamingResponse +from fastapi.responses import HTMLResponse, JSONResponse, Response, StreamingResponse from sqlmodel import select from .algorithm import create_model_mappings @@ -150,10 +151,51 @@ async def refresh_model_maps_periodically() -> None: ) +_API_PATH_PREFIXES = ("v1/", "responses") + +_NOT_FOUND_HTML_FILE = Path(__file__).parent.parent / "ui_out" / "404.html" + + +def _read_not_found_html() -> str | None: + try: + return _NOT_FOUND_HTML_FILE.read_text(encoding="utf-8") + except OSError: + return None + + +_NOT_FOUND_HTML: str | None = _read_not_found_html() + + +def _build_not_found_response(request: Request, path: str) -> Response: + """Return a 404 for unknown paths. + """ + accept = request.headers.get("accept", "").lower() + prefers_json = "application/json" in accept and "text/html" not in accept + request_id = getattr(request.state, "request_id", "unknown") + + if not prefers_json and _NOT_FOUND_HTML is not None: + return HTMLResponse(content=_NOT_FOUND_HTML, status_code=404) + + return JSONResponse( + status_code=404, + content={ + "error": { + "message": f"Path '/{path}' not found", + "type": "not_found", + "code": 404, + }, + "request_id": request_id, + }, + ) + + @proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) async def proxy( request: Request, path: str, session: AsyncSession = Depends(get_session) ) -> Response | StreamingResponse: + if not path.startswith(_API_PATH_PREFIXES): + return _build_not_found_response(request, path) + headers = dict(request.headers) is_responses_api = path.startswith("v1/responses") or path.startswith("responses") diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 644bb25b..3b44abb2 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -47,6 +47,17 @@ from .litellm_routing import detect_litellm_prefix logger = get_logger(__name__) +def _is_json_content_type(content_type: str | None) -> bool: + """Return True when the upstream response should be parsed as JSON. + """ + if not content_type: + return False + main = content_type.split(";", 1)[0].strip().lower() + if main in ("application/json", "text/json"): + return True + return main.startswith("application/") and main.endswith("+json") + + class TopupData(BaseModel): """Universal top-up data schema for Lightning Network invoices.""" @@ -523,7 +534,8 @@ class BaseUpstreamProvider: async def forward_upstream_error_response( self, request: Request, path: str, upstream_response: httpx.Response ) -> Response: - """Log upstream errors and forward the upstream response unchanged.""" + """Log upstream errors and forward the response in a JSON envelope. + """ status_code = upstream_response.status_code headers = dict(upstream_response.headers) content_type = headers.get("content-type") or headers.get("Content-Type", "") @@ -545,9 +557,10 @@ class BaseUpstreamProvider: message, upstream_code = self._extract_upstream_error_message(body_bytes) body_preview = body_bytes.decode("utf-8", errors="ignore").strip()[:500] + is_json_body = _is_json_content_type(content_type) logger.warning( - "Forwarding upstream error response as-is", + "Forwarding upstream error response", extra={ "path": path, "provider": self.provider_type, @@ -559,6 +572,7 @@ class BaseUpstreamProvider: "body_preview": body_preview, "body_read_error": body_read_error, "method": request.method, + "json_normalized": not is_json_body, }, ) @@ -586,17 +600,40 @@ class BaseUpstreamProvider: ): headers.pop(header_name, None) - if not content_type: - headers.pop("content-type", None) - headers.pop("Content-Type", None) + if is_json_body: + if not content_type: + headers.pop("content-type", None) + headers.pop("Content-Type", None) + media_type = content_type or None + return Response( + content=body_bytes, + status_code=status_code, + headers=headers, + media_type=media_type, + ) - media_type = content_type or None + # Non-JSON upstream error (HTML, plain text, empty, ...). Wrap it in + # the standard JSON envelope so callers don't need a second parser. + for header_name in ("content-type", "Content-Type"): + headers.pop(header_name, None) + + envelope = { + "error": { + "message": message or "Upstream returned a non-JSON error response", + "type": "upstream_error", + "code": upstream_code or status_code, + "upstream_status": status_code, + "upstream_content_type": content_type or None, + "upstream_body_preview": body_preview or None, + }, + "request_id": getattr(request.state, "request_id", None), + } return Response( - content=body_bytes, + content=json.dumps(envelope).encode(), status_code=status_code, headers=headers, - media_type=media_type, + media_type="application/json", ) async def handle_streaming_chat_completion( diff --git a/routstr/upstream/routstr.py b/routstr/upstream/routstr.py index ab6dd0bd..c9962301 100644 --- a/routstr/upstream/routstr.py +++ b/routstr/upstream/routstr.py @@ -18,6 +18,10 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider): provider_type = "routstr" default_base_url = None platform_url = None + # Upstream Routstr nodes serve `/v1/messages` natively, so forward the + # request as-is instead of round-tripping through litellm's + # Anthropic→OpenAI translator. + supports_anthropic_messages = True def __init__( self, @@ -43,6 +47,13 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider): ) self.settings = provider_settings or {} + def normalize_request_path( + self, path: str, model_obj: "Model | None" = None + ) -> str: + """Preserve the ``v1/`` prefix when forwarding to an upstream Routstr. + """ + return path.lstrip("/") + @classmethod def from_db_row( cls, provider_row: "UpstreamProviderRow" diff --git a/tests/unit/test_proxy_not_found.py b/tests/unit/test_proxy_not_found.py new file mode 100644 index 00000000..54edd29c --- /dev/null +++ b/tests/unit/test_proxy_not_found.py @@ -0,0 +1,64 @@ +"""Tests for the built-in 404 handler in routstr.proxy.""" + +from __future__ import annotations + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from routstr import proxy +from routstr.proxy import proxy_router + + +def _make_app() -> FastAPI: + app = FastAPI() + app.include_router(proxy_router) + return app + + +@pytest.mark.skipif( + proxy._NOT_FOUND_HTML is None, + reason="UI bundle (ui_out/404.html) not present in this environment", +) +def test_unknown_path_returns_html_404_for_browser() -> None: + client = TestClient(_make_app()) + response = client.get("/some/random/page", headers={"accept": "text/html"}) + assert response.status_code == 404 + assert response.headers["content-type"].startswith("text/html") + assert "404" in response.text + + +def test_unknown_path_returns_json_404_for_api_client() -> None: + client = TestClient(_make_app()) + response = client.get( + "/some/random/page", headers={"accept": "application/json"} + ) + assert response.status_code == 404 + assert response.headers["content-type"].startswith("application/json") + payload = response.json() + assert payload["error"]["type"] == "not_found" + assert payload["error"]["code"] == 404 + assert "/some/random/page" in payload["error"]["message"] + + +def test_root_path_returns_404_for_proxy_router() -> None: + client = TestClient(_make_app()) + response = client.get("/", headers={"accept": "application/json"}) + assert response.status_code == 404 + + +def test_v1_path_is_not_intercepted_by_404_handler() -> None: + """Paths starting with v1/ must reach the proxy logic, not the 404 handler.""" + client = TestClient(_make_app(), raise_server_exceptions=False) + response = client.get("/v1/anything") + if response.status_code == 404: + # Any 404 here must come from inner proxy logic, not our HTML page. + assert "" not in response.text + + +def test_json_returned_when_ui_html_missing(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(proxy, "_NOT_FOUND_HTML", None) + client = TestClient(_make_app()) + response = client.get("/some/random/page", headers={"accept": "text/html"}) + assert response.status_code == 404 + assert response.headers["content-type"].startswith("application/json") diff --git a/tests/unit/test_upstream_error_response.py b/tests/unit/test_upstream_error_response.py new file mode 100644 index 00000000..62b9a32c --- /dev/null +++ b/tests/unit/test_upstream_error_response.py @@ -0,0 +1,150 @@ +"""Tests for ``BaseUpstreamProvider.forward_upstream_error_response``. + +Upstream services (e.g. an Express server that doesn't expose ``/messages``) +sometimes return a non-JSON error body. The proxy must surface those errors +in a consistent JSON envelope so clients don't have to parse HTML. +""" + +from __future__ import annotations + +import json +from typing import Any +from unittest.mock import Mock + +import httpx +import pytest + +from routstr.upstream.base import BaseUpstreamProvider, _is_json_content_type + + +def _make_request(request_id: str = "req-123") -> Mock: + request = Mock(spec=["method", "state"]) + request.method = "POST" + request.state = Mock() + request.state.request_id = request_id + return request + + +def _make_upstream_response( + *, + body: bytes, + status_code: int = 404, + content_type: str | None = "text/html", + extra_headers: dict[str, str] | None = None, +) -> httpx.Response: + headers: dict[str, str] = {} + if content_type is not None: + headers["content-type"] = content_type + if extra_headers: + headers.update(extra_headers) + return httpx.Response(status_code=status_code, headers=headers, content=body) + + +@pytest.fixture +def provider() -> BaseUpstreamProvider: + return BaseUpstreamProvider( + base_url="https://privateprovider.xyz", api_key="k", provider_fee=1.0 + ) + + +@pytest.mark.parametrize( + "content_type,expected", + [ + ("application/json", True), + ("application/json; charset=utf-8", True), + ("text/json", True), + ("application/problem+json", True), + ("application/vnd.api+json", True), + ("text/html", False), + ("text/html; charset=utf-8", False), + ("text/plain", False), + ("", False), + (None, False), + ], +) +def test_is_json_content_type(content_type: str | None, expected: bool) -> None: + assert _is_json_content_type(content_type) is expected + + +@pytest.mark.asyncio +async def test_html_error_is_normalized_to_json_envelope( + provider: BaseUpstreamProvider, +) -> None: + html_body = ( + b"
Cannot POST /messages" + ) + upstream = _make_upstream_response(body=html_body, status_code=404) + + response = await provider.forward_upstream_error_response( + _make_request(), "v1/messages", upstream + ) + + assert response.status_code == 404 + assert response.media_type == "application/json" + payload: dict[str, Any] = json.loads(bytes(response.body)) + assert payload["error"]["type"] == "upstream_error" + assert payload["error"]["upstream_status"] == 404 + assert payload["error"]["upstream_content_type"] == "text/html" + assert "Cannot POST /messages" in payload["error"]["upstream_body_preview"] + assert payload["request_id"] == "req-123" + # The upstream's text/html content-type must not survive — Response() + # sets the JSON content-type for us via media_type. + assert response.headers["content-type"].startswith("application/json") + + +@pytest.mark.asyncio +async def test_plain_text_error_is_normalized( + provider: BaseUpstreamProvider, +) -> None: + upstream = _make_upstream_response( + body=b"Service Unavailable", status_code=503, content_type="text/plain" + ) + + response = await provider.forward_upstream_error_response( + _make_request(), "v1/messages", upstream + ) + + assert response.status_code == 503 + assert response.media_type == "application/json" + payload = json.loads(bytes(response.body)) + assert payload["error"]["message"] == "Service Unavailable" + + +@pytest.mark.asyncio +async def test_empty_body_with_non_json_content_type_normalizes( + provider: BaseUpstreamProvider, +) -> None: + upstream = _make_upstream_response( + body=b"", status_code=502, content_type="text/html" + ) + + response = await provider.forward_upstream_error_response( + _make_request(), "v1/messages", upstream + ) + + assert response.status_code == 502 + assert response.media_type == "application/json" + payload = json.loads(bytes(response.body)) + assert payload["error"]["type"] == "upstream_error" + assert payload["error"]["upstream_body_preview"] is None + + +@pytest.mark.asyncio +async def test_json_error_body_is_passed_through_unchanged( + provider: BaseUpstreamProvider, +) -> None: + json_body = json.dumps( + {"error": {"message": "Invalid model", "type": "invalid_request_error"}} + ).encode() + upstream = _make_upstream_response( + body=json_body, status_code=400, content_type="application/json" + ) + + response = await provider.forward_upstream_error_response( + _make_request(), "v1/messages", upstream + ) + + assert response.status_code == 400 + assert bytes(response.body) == json_body + assert response.media_type == "application/json" diff --git a/tests/unit/test_upstream_routstr.py b/tests/unit/test_upstream_routstr.py index 8cb2427d..dc0570ff 100644 --- a/tests/unit/test_upstream_routstr.py +++ b/tests/unit/test_upstream_routstr.py @@ -74,3 +74,39 @@ async def test_get_balance_returns_none_on_connect_timeout( balance = await provider.get_balance() assert balance is None + + +def test_normalize_request_path_keeps_v1_prefix() -> None: + """Routstr upstream stores ``base_url`` without ``/v1``; the prefix + must stay on the path so ``build_request_url`` produces ``/v1/