mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-31 15:56:14 +00:00
Compare commits
15
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3c13be20cb | ||
|
|
9e33d3b100 | ||
|
|
c88665f6d1 | ||
|
|
a8157b3e2d | ||
|
|
a07e6723d9 | ||
|
|
5bc7b741bc | ||
|
|
1db37cf084 | ||
|
|
6c5b103149 | ||
|
|
27ec052c17 | ||
|
|
91e5198a94 | ||
|
|
27f1cc3c42 | ||
|
|
a4d048f2b5 | ||
|
|
752d4f3803 | ||
|
|
9d905758ec | ||
|
|
6e9932e0ac |
+1
-1
@@ -4,7 +4,7 @@ FROM node:23-alpine AS ui-builder
|
||||
WORKDIR /app/ui
|
||||
|
||||
# Install pnpm
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@10.15.0 --activate
|
||||
|
||||
# Copy UI source
|
||||
COPY ui/package.json ui/pnpm-lock.yaml* ./
|
||||
|
||||
+22
-5
@@ -211,8 +211,20 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None:
|
||||
_refund_cache[key] = (expiry, value)
|
||||
|
||||
|
||||
async def _lookup_key_no_create(
|
||||
bearer_value: str, session: AsyncSession
|
||||
) -> ApiKey | None:
|
||||
"""Look up an existing API key without creating one Used by the refund endpoint"""
|
||||
if bearer_value.startswith("sk-"):
|
||||
return await session.get(ApiKey, bearer_value[3:])
|
||||
if bearer_value.startswith("cashu"):
|
||||
hashed = hashlib.sha256(bearer_value.encode()).hexdigest()
|
||||
return await session.get(ApiKey, hashed)
|
||||
return None
|
||||
|
||||
|
||||
async def _restore_balance(
|
||||
session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int
|
||||
session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str
|
||||
) -> None:
|
||||
"""Restore balance after a failed refund mint attempt."""
|
||||
restore_stmt = (
|
||||
@@ -227,7 +239,7 @@ async def _restore_balance(
|
||||
await session.commit()
|
||||
logger.info(
|
||||
"refund_wallet_endpoint: balance restored after mint failure",
|
||||
extra={"hashed_key": hashed_key, "restored_balance": balance},
|
||||
extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url},
|
||||
)
|
||||
|
||||
|
||||
@@ -282,7 +294,12 @@ async def refund_wallet_endpoint(
|
||||
)
|
||||
|
||||
bearer_value: str = authorization[7:]
|
||||
key: ApiKey = await validate_bearer_key(bearer_value, session)
|
||||
key: ApiKey | None = await _lookup_key_no_create(bearer_value, session)
|
||||
if key is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Key not found. Deposit first via /v1/wallet/create before requesting a refund.",
|
||||
)
|
||||
|
||||
if key.total_balance <= 0:
|
||||
if cached := await _refund_cache_get(bearer_value):
|
||||
@@ -373,11 +390,11 @@ async def refund_wallet_endpoint(
|
||||
|
||||
except HTTPException:
|
||||
# Minting failed — restore the debited balance
|
||||
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved)
|
||||
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
|
||||
raise
|
||||
except Exception as e:
|
||||
# Minting failed — restore the debited balance
|
||||
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved)
|
||||
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
|
||||
error_msg = str(e)
|
||||
if (
|
||||
"mint" in error_msg.lower()
|
||||
|
||||
+43
-1
@@ -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")
|
||||
|
||||
@@ -48,6 +48,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."""
|
||||
|
||||
@@ -524,7 +535,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", "")
|
||||
@@ -546,9 +558,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,
|
||||
@@ -560,6 +573,7 @@ class BaseUpstreamProvider:
|
||||
"body_preview": body_preview,
|
||||
"body_read_error": body_read_error,
|
||||
"method": request.method,
|
||||
"json_normalized": not is_json_body,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -587,17 +601,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(
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -372,17 +372,17 @@ async def test_refund_rejects_concurrent_topup_on_same_key(
|
||||
topup_amount_sat = 500
|
||||
topup_token = await testmint_wallet.mint_tokens(topup_amount_sat)
|
||||
|
||||
validate_called = asyncio.Event()
|
||||
key_looked_up = asyncio.Event()
|
||||
allow_refund_to_continue = asyncio.Event()
|
||||
original_validate_bearer_key = balance_module.validate_bearer_key
|
||||
original_lookup = balance_module._lookup_key_no_create
|
||||
delayed_once = False
|
||||
|
||||
async def delayed_validate_bearer_key(*args: Any, **kwargs: Any) -> ApiKey:
|
||||
async def delayed_lookup_key_no_create(*args: Any, **kwargs: Any) -> ApiKey | None:
|
||||
nonlocal delayed_once
|
||||
key = await original_validate_bearer_key(*args, **kwargs)
|
||||
if not delayed_once:
|
||||
key = await original_lookup(*args, **kwargs)
|
||||
if not delayed_once and key is not None:
|
||||
delayed_once = True
|
||||
validate_called.set()
|
||||
key_looked_up.set()
|
||||
await allow_refund_to_continue.wait()
|
||||
return key
|
||||
|
||||
@@ -390,7 +390,7 @@ async def test_refund_rejects_concurrent_topup_on_same_key(
|
||||
return await authenticated_client.post("/v1/wallet/refund")
|
||||
|
||||
async def issue_topup() -> Any:
|
||||
await validate_called.wait()
|
||||
await key_looked_up.wait()
|
||||
try:
|
||||
return await authenticated_client.post(
|
||||
"/v1/wallet/topup", params={"cashu_token": topup_token}
|
||||
@@ -399,7 +399,7 @@ async def test_refund_rejects_concurrent_topup_on_same_key(
|
||||
allow_refund_to_continue.set()
|
||||
|
||||
with patch(
|
||||
"routstr.balance.validate_bearer_key", new=delayed_validate_bearer_key
|
||||
"routstr.balance._lookup_key_no_create", new=delayed_lookup_key_no_create
|
||||
):
|
||||
refund_response, topup_response = await asyncio.gather(
|
||||
issue_refund(), issue_topup()
|
||||
|
||||
@@ -168,12 +168,12 @@ async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> No
|
||||
refund_token = "cashuArefund_apikey_token"
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(return_value=_update_result(1))
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()) as mock_store,
|
||||
@@ -203,12 +203,12 @@ async def test_apikey_refund_logs_token() -> None:
|
||||
refund_token = "cashuAlogged_token"
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(return_value=_update_result(1))
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
|
||||
@@ -232,12 +232,12 @@ async def test_apikey_refund_log_includes_path() -> None:
|
||||
refund_token = "cashuApath_token"
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(return_value=_update_result(1))
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
|
||||
@@ -269,6 +269,7 @@ async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None:
|
||||
key = _make_api_key(balance=5000, refund_currency="sat")
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
# Debit returns rowcount=0 → balance changed concurrently
|
||||
session.exec = AsyncMock(return_value=_update_result(0))
|
||||
session.commit = AsyncMock()
|
||||
@@ -276,7 +277,6 @@ async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None:
|
||||
mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted")
|
||||
|
||||
with (
|
||||
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", mock_send_token),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
|
||||
@@ -333,11 +333,11 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
|
||||
|
||||
# First exec call = debit (succeeds), second = restore
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)])
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", AsyncMock(side_effect=Exception("mint down"))),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
|
||||
@@ -355,3 +355,49 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
|
||||
assert exc_info.value.status_code == 503
|
||||
# Verify two exec calls: debit + restore
|
||||
assert session.exec.await_count == 2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# no-create guarantee: fresh Cashu/unknown sk- tokens must not create API keys
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_fresh_cashu_bearer_returns_401() -> None:
|
||||
"""Fresh Cashu token not in DB must get 401, never create a new ApiKey."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer cashuAfresh_never_deposited_token",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
session.get.assert_awaited_once()
|
||||
# No add/commit → no key was persisted
|
||||
session.add.assert_not_called()
|
||||
session.commit.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_unknown_sk_bearer_returns_401() -> None:
|
||||
"""Unknown sk- key not in DB must get 401."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-unknownhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
session.get.assert_awaited_once()
|
||||
|
||||
@@ -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")
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -6,6 +6,9 @@ const nextConfig: NextConfig = {
|
||||
images: {
|
||||
unoptimized: true,
|
||||
},
|
||||
turbopack: {
|
||||
root: __dirname,
|
||||
},
|
||||
};
|
||||
|
||||
export default nextConfig;
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
{
|
||||
"name": "routstr-service",
|
||||
"packageManager": "pnpm@10.15.0",
|
||||
"version": "0.1.0",
|
||||
"private": true,
|
||||
"scripts": {
|
||||
@@ -88,5 +89,11 @@
|
||||
"prettier-plugin-tailwindcss": "^0.7.2",
|
||||
"tailwindcss": "^4.2.0",
|
||||
"typescript": "^5.9.3"
|
||||
},
|
||||
"pnpm": {
|
||||
"onlyBuiltDependencies": [
|
||||
"sharp",
|
||||
"unrs-resolver"
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user