Compare commits

...

4 Commits

Author SHA1 Message Date
9qeklajc
27f1cc3c42 enforce load mint keyset 2026-05-10 08:32:19 +02:00
9qeklajc
752d4f3803 Merge pull request #496 from Routstr/fix-provider-forwarding-upstream
fix provder forwarded path & wrong html response
2026-05-09 14:47:35 +02:00
9qeklajc
9d905758ec added page not found redirection 2026-05-09 14:07:45 +02:00
9qeklajc
6e9932e0ac fix provder forwarded path & wrong html response 2026-05-09 12:05:51 +02:00
7 changed files with 351 additions and 9 deletions

View File

@@ -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")

View File

@@ -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(

View File

@@ -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"

View File

@@ -48,6 +48,8 @@ async def recieve_token(
if token_obj.mint not in settings.cashu_mints:
return await swap_to_primary_mint(token_obj, wallet)
await wallet.load_mint(keyset_id=token_obj.keysets[0])
wallet.verify_proofs_dleq(token_obj.proofs)
await wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True)

View File

@@ -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 "<!DOCTYPE html>" 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")

View File

@@ -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"<!DOCTYPE html><html><head><title>Error</title></head>"
b"<body><pre>Cannot POST /messages</pre></body></html>"
)
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"

View File

@@ -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/<endpoint>``
instead of ``/<endpoint>`` (which the upstream Routstr 404s with HTML)."""
provider = RoutstrUpstreamProvider(
base_url="https://privateprovider.xyz", api_key="key"
)
assert provider.normalize_request_path("v1/messages") == "v1/messages"
assert provider.normalize_request_path("/v1/messages") == "v1/messages"
assert (
provider.normalize_request_path("v1/chat/completions")
== "v1/chat/completions"
)
def test_build_request_url_for_v1_messages() -> None:
"""Forwarding ``/v1/messages`` must hit the upstream's ``/v1/messages``."""
provider = RoutstrUpstreamProvider(
base_url="https://privateprovider.xyz", api_key="key"
)
normalized = provider.normalize_request_path("v1/messages")
assert (
provider.build_request_url(normalized)
== "https://privateprovider.xyz/v1/messages"
)
def test_supports_anthropic_messages_natively() -> None:
"""Routstr nodes serve ``/v1/messages`` directly, so the proxy must
forward as-is instead of round-tripping through litellm."""
assert RoutstrUpstreamProvider.supports_anthropic_messages is True