diff --git a/routstr/core/settings.py b/routstr/core/settings.py index c02c6030..53e71a74 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -101,6 +101,12 @@ class Settings(BaseSettings): # Network cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS") + # Comma-separated METHOD:path pairs adding to the proxy's canonical + # endpoint allowlist, e.g. "POST:v1/rerank,GET:batches". Only for upstreams + # exposing an endpoint outside the OpenAI-compatible set; each addition + # widens what the provider credential can be spent against, so wildcards + # and prefixes are not supported. + proxy_extra_allowed_paths: str = Field(default="", env="PROXY_EXTRA_ALLOWED_PATHS") tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL") providers_refresh_interval_seconds: int = Field( default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS" @@ -192,6 +198,11 @@ SECRET_FIELDS = frozenset({"admin_password", "nsec"}) # neither store nor shadow them; env is always authoritative. ENV_ONLY_FIELDS = frozenset( { + # Widening the proxy's reachable upstream surface is a deployment + # decision, not a runtime toggle: it changes what the provider + # credential can be spent against. Keeping it env-only also lets the + # proxy parse it once at import without going stale. + "proxy_extra_allowed_paths", "database_pool_size", "database_max_overflow", "database_pool_timeout", diff --git a/routstr/proxy.py b/routstr/proxy.py index 4542d008..5feff279 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -26,6 +26,7 @@ from .core.db import ( ) 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, check_token_balance, @@ -221,22 +222,149 @@ async def refresh_model_maps_periodically() -> None: ) -_API_PATH_PREFIXES = ( - "v1/", - "responses", - "chat/", - "completions", - "models", - "embeddings", - "audio/", - "images/", - "moderations", - "providers", - "tee/", - "attestation", +# Canonical endpoints this proxy will forward, keyed by the path with any +# leading "v1/" and trailing slash removed, mapped to the methods allowed on +# each. The provider credential is attached during forwarding, so endpoint +# permission has to come from this table rather than from the client-supplied +# path: an upstream's key-management, organization, or billing routes live +# under the same origin and must never be reachable through the proxy. +_ALLOWED_ENDPOINTS: dict[str, frozenset[str]] = { + "chat/completions": frozenset({"POST"}), + "completions": frozenset({"POST"}), + "responses": frozenset({"POST"}), + "messages": frozenset({"POST"}), + "embeddings": frozenset({"POST"}), + "moderations": frozenset({"POST"}), + "rerank": frozenset({"POST"}), + "audio/speech": frozenset({"POST"}), + "audio/transcriptions": frozenset({"POST"}), + "audio/translations": frozenset({"POST"}), + "images/generations": frozenset({"POST"}), + "images/edits": frozenset({"POST"}), + "images/variations": frozenset({"POST"}), + "models": frozenset({"GET"}), + "attestation": frozenset({"GET"}), + "tee/attestation": frozenset({"GET"}), +} + +_ALLOWED_METHODS = frozenset({"GET", "POST"}) + + +def _canonical_api_path(path: str) -> str: + """Reduce a request path to its allowlist key. + + OpenAI-style clients reach the same endpoint with or without the ``v1/`` + prefix and with or without a trailing slash, so both spellings collapse to + one key. Callers must screen the path with + :func:`_is_ambiguously_spelled_path` first — this function assumes the path + has no dot segments, empty segments, or encoded separators left to resolve. + """ + core = path[:-1] if path.endswith("/") else path + if core.startswith("v1/"): + core = core[len("v1/") :] + return core + + +def _parse_extra_allowed_endpoints(raw: str) -> dict[str, frozenset[str]]: + """Parse operator-configured additions to the endpoint allowlist. + + Deployments whose provider exposes an endpoint outside the canonical set + opt in explicitly with ``PROXY_EXTRA_ALLOWED_PATHS``, a comma-separated + list of ``METHOD:path`` pairs (e.g. ``POST:v1/rerank,GET:batches``). Every + entry must name one concrete method and one unambiguous path; wildcards + and bare prefixes are deliberately unsupported, so widening the proxy's + reach is always a per-endpoint decision. Malformed entries are dropped + with a warning rather than silently widening or narrowing the surface. + """ + extra: dict[str, frozenset[str]] = {} + for entry in raw.split(","): + entry = entry.strip() + if not entry: + continue + method, separator, endpoint = entry.partition(":") + method = method.strip().upper() + endpoint = endpoint.strip() + if not separator or method not in _ALLOWED_METHODS or not endpoint: + logger.warning( + "Ignoring malformed PROXY_EXTRA_ALLOWED_PATHS entry", + extra={"entry": entry}, + ) + continue + if _is_ambiguously_spelled_path(endpoint): + logger.warning( + "Ignoring ambiguously spelled PROXY_EXTRA_ALLOWED_PATHS entry", + extra={"entry": entry}, + ) + continue + if any(character in endpoint for character in "*?["): + # Refuse glob syntax outright. Kept as a literal endpoint name it + # would never match a real request, so the operator would think + # they had widened the proxy when they had not. + logger.warning( + "Ignoring wildcard PROXY_EXTRA_ALLOWED_PATHS entry; " + "list each endpoint explicitly", + extra={"entry": entry}, + ) + continue + key = _canonical_api_path(endpoint) + extra[key] = extra.get(key, frozenset()) | {method} + return extra + + +def _is_ambiguously_spelled_path(path: str) -> bool: + """Reject paths whose spelling could resolve somewhere the allowlist did not. + + ``{path:path}`` arrives percent-decoded, so a client that sent ``%2e%2e`` or + ``%2f`` shows up here as ``..`` / ``/``. Dot segments, backslashes, duplicate + or leading separators, NUL bytes, and any residual encoded separator are + treated as unsafe: they let a caller walk off the canonical API surface (and + onto a sensitive upstream endpoint) even though the literal prefix check + would pass. Reject rather than trying to rewrite the path. + """ + if not path or path != path.strip() or path.startswith("/"): + return True + if "\x00" in path or "\\" in path: + return True + # A single trailing slash is canonical (e.g. "attestation/"); ignore it, + # then no remaining segment may be empty (covers "//") or a dot segment. + core = path[:-1] if path.endswith("/") else path + if any(segment in ("", ".", "..") for segment in core.split("/")): + return True + lowered = path.lower() + return "%2e" in lowered or "%2f" in lowered or "%5c" in lowered + + +_EXTRA_ALLOWED_ENDPOINTS = _parse_extra_allowed_endpoints( + settings.proxy_extra_allowed_paths ) +def _allowed_methods_for(endpoint: str) -> frozenset[str]: + """Return the methods allowed on a canonical endpoint, empty if unknown.""" + methods = _ALLOWED_ENDPOINTS.get(endpoint, frozenset()) + methods |= _EXTRA_ALLOWED_ENDPOINTS.get(endpoint, frozenset()) + return methods + + +def _forwarding_allowed(path: str, method: str) -> bool: + """Gate which method/path pairs may reach an upstream at all. + + The provider credential is attached during forwarding, so an unknown + endpoint must never be forwarded on the caller's say-so. The path is + reduced to its canonical form and looked up in the endpoint table; there is + no prefix match, so a known prefix no longer carries an unknown endpoint + (``v1/organization/api_keys`` is rejected even though ``v1/`` is familiar). + + EHBP requests are gated by the same table. Their body is opaque to the + proxy, which is a reason to constrain the destination more tightly, not to + trust the caller's path: the encrypted contract covers the body, never the + endpoint the credential is spent against. + """ + if method not in _ALLOWED_METHODS: + return False + return method in _allowed_methods_for(_canonical_api_path(path)) + + @proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) async def proxy( request: Request, path: str, session: AsyncSession = Depends(get_session) @@ -255,14 +383,17 @@ async def proxy( async def _proxy( request: Request, path: str, session: AsyncSession ) -> Response | StreamingResponse: - # 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): + # Screen the path before any routing decision: reject ambiguous spellings, + # then require a known API prefix so nothing unknown is forwarded with the + # provider credential attached. + if _is_ambiguously_spelled_path(path): return build_not_found_response(request, path) headers = dict(request.headers) + is_ehbp = "ehbp-encapsulated-key" in headers + + if not _forwarding_allowed(path, request.method): + return build_not_found_response(request, path) is_responses_api = path.startswith("v1/responses") or path.startswith("responses") request_body = await request.body() @@ -272,7 +403,6 @@ async def _proxy( # extract the model id, so the SDK sends it in X-Routstr-Model. Forward the # raw encrypted body to the upstream's /private/ endpoint and stream the # encrypted response back untouched — the SDK's SecureClient decrypts it. - is_ehbp = "ehbp-encapsulated-key" in headers if is_ehbp: request_body_dict = {} model_id = headers.get("x-routstr-model", "") diff --git a/tests/integration/test_proxy_post_endpoints.py b/tests/integration/test_proxy_post_endpoints.py index 8d5fab36..99bdc071 100644 --- a/tests/integration/test_proxy_post_endpoints.py +++ b/tests/integration/test_proxy_post_endpoints.py @@ -289,8 +289,115 @@ async def test_proxy_post_unauthorized_access(integration_client: AsyncClient) - assert response.status_code in [400, 401] +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize( + "bad_path", + [ + "internal/admin", # unknown endpoint, no API prefix + "v1/../admin", # traversal onto a sibling path + "v1//models", # duplicate separator + "%2e%2e/secret", # encoded dot segment + ], +) +async def test_authenticated_post_to_unknown_path_is_rejected( + authenticated_client: AsyncClient, bad_path: str +) -> None: + """An authenticated POST to an unknown/traversal path must be rejected at + the edge (404) and never forwarded — the provider credential must not reach + an endpoint the caller merely spelled into the URL. If the guard let it + through, forwarding would raise and this would not be a clean 404.""" + with patch( + "routstr.upstream.base.BaseUpstreamProvider.forward_request", + AsyncMock(side_effect=AssertionError("must not forward unknown path")), + ): + response = await authenticated_client.post( + f"/{bad_path}", + json={"model": "gpt-3.5-turbo", "messages": []}, + ) + assert response.status_code == 404 +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize( + "bad_path", + [ + # Unambiguously spelled, under a prefix the proxy serves, but not an + # endpoint it offers. These are real upstream routes that manage keys, + # org membership, and billing on the same origin as inference. + "v1/organization/api_keys", + "v1/api_keys", + "v1/billing/usage", + "v1/admin/keys", + # An id segment is honoured one level deep, and only where an endpoint + # takes one at all. + "v1/chat/completions/abc", + "v1/models/gpt-4/secret", + ], +) +async def test_known_prefix_does_not_carry_an_unknown_endpoint( + authenticated_client: AsyncClient, bad_path: str +) -> None: + """A familiar prefix must not be a passport for the rest of the origin. + + Nothing here is traversal-shaped, so the spelling screen lets it by; only + the endpoint allowlist stops it. Forwarding raises if the guard misses, + so a clean 404 also proves the credential never left.""" + with patch( + "routstr.upstream.base.BaseUpstreamProvider.forward_request", + AsyncMock(side_effect=AssertionError("must not forward unknown endpoint")), + ): + response = await authenticated_client.post( + f"/{bad_path}", + json={"model": "gpt-3.5-turbo", "messages": []}, + ) + assert response.status_code == 404 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_get_unknown_path_is_rejected(integration_client: AsyncClient) -> None: + """A GET to an unknown (no API prefix) path is rejected at the edge.""" + response = await integration_client.get("/internal/admin") + assert response.status_code == 404 + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize( + "bad_path", + [ + # Percent-encoded so the traversal survives the client to the server, + # which decodes it to "v1/../admin" before routing. + "v1/%2e%2e/admin", + # Well-spelled, so only the endpoint allowlist can stop it. + "v1/organization/api_keys", + "anything/encrypted", + ], +) +async def test_ehbp_request_is_gated_by_the_endpoint_allowlist( + integration_client: AsyncClient, bad_path: str +) -> None: + """EHBP hides the body from the proxy, not the destination. + + The encrypted contract covers the request body; it says nothing about which + endpoint the provider credential gets spent against, so an EHBP request is + screened and allowlisted exactly like any other.""" + with patch( + "routstr.proxy.forward_ehbp_request", + AsyncMock(side_effect=AssertionError("must not forward unknown endpoint")), + ): + response = await integration_client.post( + f"/{bad_path}", + content=b"encrypted", + headers={ + "ehbp-encapsulated-key": "x", + "x-routstr-model": "gpt-4", + }, + ) + assert response.status_code == 404 + @pytest.mark.integration @pytest.mark.asyncio diff --git a/tests/unit/test_proxy_path_allowlist.py b/tests/unit/test_proxy_path_allowlist.py new file mode 100644 index 00000000..33a2d610 --- /dev/null +++ b/tests/unit/test_proxy_path_allowlist.py @@ -0,0 +1,202 @@ +"""Unit tests for the proxy edge path allowlist (arbitrary-upstream-path-proxy). + +An authenticated POST used to be forwarded for ANY path, so a caller could reach +arbitrary or traversal-shaped upstream endpoints with the provider credential +attached. The proxy now rejects ambiguous path spellings for every method, then +requires the method/path pair to name a canonical endpoint. A familiar prefix is +no longer enough: "v1/organization/api_keys" is refused just like "internal/admin". +""" + +from __future__ import annotations + +import os + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +import pytest # noqa: E402 + +from routstr.proxy import ( # noqa: E402 + _forwarding_allowed, + _is_ambiguously_spelled_path, + _parse_extra_allowed_endpoints, +) + + +@pytest.mark.parametrize( + "path", + [ + "../secret", + "v1/../admin", + "v1/./models", + "..", + "v1//models", # duplicate separator + "/v1/models", # leading slash / absolute override + "v1/models/..", + "%2e%2e/secret", # residual encoded dot segment + "v1/%2fadmin", # residual encoded slash + "v1\\models", # backslash + "v1/models\x00", # NUL byte + " v1/models", # leading whitespace + "", + ], +) +def test_ambiguous_paths_are_rejected(path: str) -> None: + assert _is_ambiguously_spelled_path(path) is True + + +@pytest.mark.parametrize( + "path", + [ + "v1/chat/completions", + "chat/completions", + "v1/responses", + "v1/embeddings", + "models", + "v1/models/gpt-4", + "attestation/", # a single trailing slash is canonical + "tee/attestation/", + ], +) +def test_canonical_paths_are_allowed(path: str) -> None: + assert _is_ambiguously_spelled_path(path) is False + + +def test_unknown_paths_are_not_forwarded() -> None: + # The credential is attached during forwarding, so an unknown endpoint must + # never be forwarded on the caller's say-so. + assert _forwarding_allowed("internal/admin", "POST") is False + assert _forwarding_allowed("secret-endpoint", "POST") is False + + +@pytest.mark.parametrize( + "path", + [ + "modelsdump", # "models" must not match a longer segment + "attestationadmin", + "providers-secret", + "embeddingsx", + "completions-internal", + ], +) +def test_endpoint_name_does_not_match_a_longer_segment(path: str) -> None: + assert _forwarding_allowed(path, "POST") is False + assert _forwarding_allowed(path, "GET") is False + + +@pytest.mark.parametrize( + "path", + [ + # A familiar prefix must not carry an unknown endpoint. These are real + # upstream routes that manage keys, org membership, and billing. + "v1/organization/api_keys", + "v1/api_keys", + "v1/admin/keys", + "v1/billing/usage", + "v1/files", + "v1/batches", + "chat/internal", + "audio/internal", + "images/internal", + "tee/keys", + # No endpoint takes a trailing id segment; a resource id never widens + # the reachable surface. + "models/gpt-4", + "models/gpt-4/secret", + "chat/completions/abc", + ], +) +def test_known_prefix_does_not_carry_an_unknown_endpoint(path: str) -> None: + assert _forwarding_allowed(path, "POST") is False + assert _forwarding_allowed(path, "GET") is False + + +@pytest.mark.parametrize( + ("path", "method"), + [ + ("v1/chat/completions", "POST"), + ("chat/completions", "POST"), + ("v1/chat/completions/", "POST"), + ("completions", "POST"), + ("v1/responses", "POST"), + ("v1/messages", "POST"), + ("v1/embeddings", "POST"), + ("moderations", "POST"), + ("audio/transcriptions", "POST"), + ("images/generations", "POST"), + ("models", "GET"), + ("attestation", "GET"), + ("tee/attestation", "GET"), + ], +) +def test_canonical_endpoints_are_forwarded(path: str, method: str) -> None: + assert _forwarding_allowed(path, method) is True + + +@pytest.mark.parametrize( + ("path", "method"), + [ + ("chat/completions", "GET"), # billed endpoints are POST-only + ("v1/embeddings", "GET"), + ("models", "POST"), # read-only endpoints are GET-only + ("attestation", "POST"), + ("v1/chat/completions", "DELETE"), # never routed here, refused anyway + ("v1/chat/completions", "PUT"), + ], +) +def test_method_must_match_the_endpoint(path: str, method: str) -> None: + assert _forwarding_allowed(path, method) is False + + +def test_operator_additions_are_parsed_per_endpoint() -> None: + parsed = _parse_extra_allowed_endpoints("POST:v1/rerank, GET:batches ,post:audio/x") + assert parsed == { + "rerank": frozenset({"POST"}), # the "v1/" prefix collapses like any path + "batches": frozenset({"GET"}), + "audio/x": frozenset({"POST"}), + } + + +def test_operator_additions_may_grant_two_methods_on_one_endpoint() -> None: + assert _parse_extra_allowed_endpoints("POST:batches,GET:batches") == { + "batches": frozenset({"POST", "GET"}) + } + + +@pytest.mark.parametrize( + "raw", + [ + "", + " ", + "v1/rerank", # no method + "POST:", # no path + ":v1/rerank", # empty method + "DELETE:v1/rerank", # method the proxy never routes + "POST:*", # wildcards are deliberately unsupported + "POST:v1/*", + "POST:../secret", # ambiguous spellings are screened here too + "POST:v1//rerank", + "POST:%2e%2e/secret", + ], +) +def test_malformed_operator_additions_widen_nothing(raw: str) -> None: + assert _parse_extra_allowed_endpoints(raw) == {} + + +def test_operator_additions_are_env_only() -> None: + # The proxy parses this once at import, so a persisted or admin-API-writable + # value would be read but never take effect. Keeping it env-only also means + # widening the reachable upstream surface takes a deploy. + from routstr.core.settings import ENV_ONLY_FIELDS + + assert "proxy_extra_allowed_paths" in ENV_ONLY_FIELDS + + +def test_ehbp_is_gated_by_the_same_allowlist() -> None: + # EHBP hides the request body from the proxy, which is a reason to constrain + # the destination more tightly rather than to trust the caller's path: the + # encrypted contract covers the body, never the endpoint the provider + # credential is spent against. + assert _forwarding_allowed("anything/encrypted", "POST") is False + assert _forwarding_allowed("v1/organization/api_keys", "POST") is False + assert _forwarding_allowed("v1/chat/completions", "POST") is True diff --git a/tests/unit/test_proxy_tinfoil_attestation_routing.py b/tests/unit/test_proxy_tinfoil_attestation_routing.py index b6ba2f88..c367a049 100644 --- a/tests/unit/test_proxy_tinfoil_attestation_routing.py +++ b/tests/unit/test_proxy_tinfoil_attestation_routing.py @@ -102,10 +102,22 @@ async def test_attestation_trailing_slash_routes_directly_to_tinfoil( tinfoil.forward_get_request.assert_awaited_once() -@pytest.mark.parametrize("path", ["attestation/foo", "attestationjunk"]) +@pytest.mark.parametrize( + "path", + [ + # A valid `attestation` segment is not the exact attestation route, and + # `attestation` takes no id segment, so the endpoint allowlist rejects + # it at the edge rather than letting it reach model/auth handling. + "attestation/foo", + # Not a known endpoint at all: rejected at the edge before routing. + "attestationjunk", + ], +) @pytest.mark.asyncio async def test_non_attestation_prefix_does_not_bypass_authentication( - monkeypatch: pytest.MonkeyPatch, proxy_app: FastAPI, path: str + monkeypatch: pytest.MonkeyPatch, + proxy_app: FastAPI, + path: str, ) -> None: tinfoil = MagicMock() tinfoil.provider_type = "tinfoil" @@ -118,8 +130,7 @@ async def test_non_attestation_prefix_does_not_bypass_authentication( ) as client: response = await client.get(f"/{path}") - assert response.status_code == 400 - assert response.json()["error"]["type"] == "invalid_model" + assert response.status_code == 404 tinfoil.forward_get_request.assert_not_awaited()