diff --git a/routstr/core/main.py b/routstr/core/main.py index 60913447..535480a2 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -30,6 +30,7 @@ from .db import create_session, init_db, run_migrations from .exceptions import general_exception_handler, http_exception_handler from .logging import get_logger, setup_logging from .middleware import LoggingMiddleware +from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401 from .settings import SettingsService from .settings import settings as global_settings from .version import __version__ diff --git a/routstr/core/not_found.py b/routstr/core/not_found.py new file mode 100644 index 00000000..3105ae26 --- /dev/null +++ b/routstr/core/not_found.py @@ -0,0 +1,55 @@ +"""Shared 404 handler used by the proxy catch-all and tests.""" + +from __future__ import annotations + +from pathlib import Path + +from fastapi import Request +from fastapi.responses import HTMLResponse, JSONResponse, Response + +_NOT_FOUND_HTML_FILE = Path(__file__).parent.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 response. + + HTML 404 page only for GET requests from browsers (Accept: text/html). + All POST requests and API clients receive a JSON 404. + """ + accept = request.headers.get("accept", "").lower() + prefers_html = ( + request.method == "GET" + and "text/html" in accept + and "application/json" not in accept + ) + request_id = getattr(request.state, "request_id", "unknown") + + if prefers_html 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, + }, + ) + + +async def not_found_catch_all(request: Request, path: str) -> Response: + """ASGI handler form of :func:`build_not_found_response`.""" + return build_not_found_response(request, path) diff --git a/routstr/proxy.py b/routstr/proxy.py index ad301a45..01bb753b 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -1,9 +1,8 @@ import json -from pathlib import Path from typing import Any from fastapi import APIRouter, Depends, HTTPException, Request -from fastapi.responses import HTMLResponse, JSONResponse, Response, StreamingResponse +from fastapi.responses import Response, StreamingResponse from sqlmodel import select from .algorithm import create_model_mappings @@ -18,6 +17,7 @@ from .core.db import ( get_session, ) from .core.exceptions import UpstreamError +from .core.not_found import build_not_found_response from .core.settings import settings from .payment.helpers import ( calculate_discounted_max_cost, @@ -151,50 +151,30 @@ 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, - }, - ) +_API_PATH_PREFIXES = ( + "v1/", + "responses", + "chat/", + "completions", + "models", + "embeddings", + "audio/", + "images/", + "moderations", + "providers", +) @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) + # GET requests must hit a known API prefix; otherwise return a 404 (HTML + # for browsers, JSON for API clients). POST requests are always forwarded + # so that OpenAI-style endpoints work with or without the `v1/` prefix + # (e.g. `/chat/completions` as well as `/v1/chat/completions`). + if request.method == "GET" and not path.startswith(_API_PATH_PREFIXES): + return build_not_found_response(request, path) headers = dict(request.headers) diff --git a/tests/unit/test_proxy_not_found.py b/tests/unit/test_proxy_not_found.py index 54edd29c..f2d20c38 100644 --- a/tests/unit/test_proxy_not_found.py +++ b/tests/unit/test_proxy_not_found.py @@ -1,4 +1,4 @@ -"""Tests for the built-in 404 handler in routstr.proxy.""" +"""Tests for the app-level 404 handler in routstr.core.main.""" from __future__ import annotations @@ -6,18 +6,22 @@ import pytest from fastapi import FastAPI from fastapi.testclient import TestClient -from routstr import proxy -from routstr.proxy import proxy_router +from routstr.core import main as core_main def _make_app() -> FastAPI: app = FastAPI() - app.include_router(proxy_router) + app.add_api_route( + "/{path:path}", + core_main.not_found_catch_all, + methods=["GET", "POST"], + include_in_schema=False, + ) return app @pytest.mark.skipif( - proxy._NOT_FOUND_HTML is None, + core_main._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: @@ -41,24 +45,26 @@ def test_unknown_path_returns_json_404_for_api_client() -> None: assert "/some/random/page" in payload["error"]["message"] -def test_root_path_returns_404_for_proxy_router() -> None: +def test_root_path_returns_404() -> 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) + from routstr.core import not_found as nf + + monkeypatch.setattr(nf, "_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") + + +def test_post_unknown_path_returns_json_even_for_browser() -> None: + client = TestClient(_make_app()) + response = client.post("/some/random/page", headers={"accept": "text/html"}) + assert response.status_code == 404 + assert response.headers["content-type"].startswith("application/json") + payload = response.json() + assert payload["error"]["type"] == "not_found"