Merge pull request #686 from Routstr/fix/arbitrary-upstream-path-proxy

Fix/arbitrary upstream path proxy
This commit is contained in:
9qeklajc
2026-08-24 01:13:11 +02:00
committed by GitHub
5 changed files with 484 additions and 23 deletions
+11
View File
@@ -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",
+149 -19
View File
@@ -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", "")
@@ -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
+202
View File
@@ -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
@@ -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()