fix(proxy): avoid double stripping passthrough roots

This commit is contained in:
vgvr0 2026-09-29 20:38:39 +02:00
parent f1b587e480
commit 9243b381da
4 changed files with 78 additions and 52 deletions

View file

@ -134,7 +134,6 @@ from litellm.proxy.spend_tracking.spend_counter_batch import (
from litellm.proxy.utils import (
PrismaClient,
ProxyLogging,
strip_server_root_path,
)
from litellm.repositories.table_repositories import TeamMembershipRepository
from litellm.router_utils.common_utils import resolve_model_group_alias
@ -860,7 +859,10 @@ async def check_api_key_for_custom_headers_or_pass_through_endpoints(
api_key: str,
) -> UserAPIKeyAuth | str:
is_mapped_pass_through_route: bool = False
normalized_route: Final = strip_server_root_path(route)
# ``route`` is obtained from get_request_route() by the auth flow and is
# already root-relative. Do not strip SERVER_ROOT_PATH a second time: the
# deployment root may itself be a mapped route name (e.g. /typesafe).
normalized_route: Final = route
for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value:
if normalized_route == mapped_route or normalized_route.startswith(mapped_route + "/"):
is_mapped_pass_through_route = True

View file

@ -102,7 +102,6 @@ from litellm.proxy.litellm_pre_call_utils import (
_get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.utils import strip_server_root_path
from litellm.repositories.team_repository import TeamRepository
from litellm.secret_managers.main import get_secret_str
from litellm.types import utils as types_utils
@ -3143,12 +3142,14 @@ class InitPassThroughEndpointHelpers:
@staticmethod
def _route_for_registry_lookup(route: str) -> str:
"""
Normalize an incoming route to the bare path stored in the registry.
Return the root-relative route used as a registry key.
Registry keys store root-stripped paths. Callers may pass routes from
``get_request_route()`` or prefixed ``request.url.path`` values.
Registry keys and all callers of this helper use routes from
``get_request_route()``. They are already root-stripped, so stripping
again would make a deployment root that overlaps a route name ambiguous
(for example root ``/typesafe`` and route ``/typesafe/decisions``).
"""
return strip_server_root_path(route)
return route
@staticmethod
def is_registered_pass_through_route(route: str) -> bool:
@ -3165,7 +3166,7 @@ class InitPassThroughEndpointHelpers:
bool: True if route is a registered pass-through endpoint, False otherwise
"""
## CHECK IF MAPPED PASS THROUGH ENDPOINT
normalized_route: Final = strip_server_root_path(route)
normalized_route: Final = route
for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value:
if normalized_route == mapped_route or normalized_route.startswith(mapped_route + "/"):
return True

View file

@ -41,7 +41,6 @@ from typing import (
Protocol,
TypeAlias,
TypeVar,
Union,
cast,
overload,
)
@ -8330,14 +8329,6 @@ def normalize_route_for_root_path(route: str) -> str | None:
return route
def strip_server_root_path(route: str) -> str:
"""Return a route with the SERVER_ROOT_PATH prefix removed when present."""
root_path: Final = get_server_root_path().rstrip("/")
if root_path and route.startswith(root_path + "/"):
return route[len(root_path) :]
return route
def get_prisma_client_or_throw(message: str):
from litellm.proxy.proxy_server import prisma_client

View file

@ -15,8 +15,10 @@ from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import HTTPException, Request, Response, UploadFile
import respx
from fastapi import FastAPI, HTTPException, Request, Response, UploadFile
from fastapi.responses import StreamingResponse
from fastapi.testclient import TestClient
from pydantic import TypeAdapter, ValidationError
from starlette.datastructures import FormData, Headers, QueryParams
from starlette.datastructures import UploadFile as StarletteUploadFile
@ -3475,7 +3477,7 @@ def test_is_registered_pass_through_route_with_custom_root():
}
with patch("litellm.proxy.utils.get_server_root_path", return_value="/proxy"):
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/proxy/api/endpoint") is True
# Callers pass get_request_route() output, which is already root-relative.
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/api/endpoint") is True
with patch("litellm.proxy.utils.get_server_root_path", return_value="/"):
@ -3488,8 +3490,7 @@ def test_is_registered_pass_through_route_with_custom_root():
def test_get_registered_pass_through_route_with_custom_root():
"""
get_registered_pass_through_route matches bare registry paths against
bare or SERVER_ROOT_PATH-prefixed incoming routes.
get_registered_pass_through_route matches root-relative registry paths.
"""
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
InitPassThroughEndpointHelpers,
@ -3511,13 +3512,7 @@ def test_get_registered_pass_through_route_with_custom_root():
_registered_pass_through_routes[route_key] = target_config
with patch("litellm.proxy.utils.get_server_root_path", return_value="/litellm"):
# Prefixed incoming route
result = InitPassThroughEndpointHelpers.get_registered_pass_through_route("/litellm/chat/completions")
assert result is not None
assert result["target"] == "http://api.example.com/v1/chat/completions"
assert result["headers"]["Authorization"] == "Bearer token123"
# Bare incoming route (get_request_route convention)
# get_request_route() supplies root-relative routes.
result = InitPassThroughEndpointHelpers.get_registered_pass_through_route("/chat/completions")
assert result is not None
assert result["target"] == "http://api.example.com/v1/chat/completions"
@ -3538,14 +3533,8 @@ def test_get_registered_pass_through_route_with_custom_root():
("", "exact", "/ml", True),
("", "exact", "/ml/extra", False),
("/llmproxy", "subpath", "/ml/api/v1/time-series-forecast/predict", True),
(
"/llmproxy",
"subpath",
"/llmproxy/ml/api/v1/time-series-forecast/predict",
True,
),
("/llmproxy", "subpath", "/ml/api/v1/time-series-forecast/predict", True),
("/llmproxy", "exact", "/ml", True),
("/llmproxy", "exact", "/llmproxy/ml", True),
("/llmproxy", "subpath", "/other/api", False),
],
)
@ -3553,8 +3542,8 @@ def test_db_registered_pass_through_route_bare_path_convention(
server_root_path, route_type, incoming_route, should_match
):
"""
Regression: #28547 / SERVER_ROOT_PATH — registry stores bare /ml paths;
get_request_route() supplies bare paths; prefixed url.path must still match.
Regression: #28547 / SERVER_ROOT_PATH — registry stores bare /ml paths and
get_request_route() supplies root-relative paths for registry lookup.
"""
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
InitPassThroughEndpointHelpers,
@ -3594,18 +3583,12 @@ def test_mapped_pass_through_routes_with_server_root_path():
)
with patch("litellm.proxy.utils.get_server_root_path", return_value="/litellm"):
# prefixed route should match mapped routes like /vertex_ai
assert (
InitPassThroughEndpointHelpers.is_registered_pass_through_route("/litellm/vertex_ai/v1/projects/foo")
is True
)
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/litellm/bedrock/model/invoke") is True
# get_request_route() supplies bare paths after stripping root_path.
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/vertex_ai/v1/projects/foo") is True
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/bedrock/model/invoke") is True
@pytest.mark.parametrize("server_root_path", ["", "/", "/api/v1", "/api/v1/"])
@pytest.mark.parametrize("server_root_path", ["", "/", "/api/v1", "/api/v1/", "/typesafe"])
def test_typesafe_mapped_pass_through_route_with_server_root_path(server_root_path):
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
InitPassThroughEndpointHelpers,
@ -3615,18 +3598,67 @@ def test_typesafe_mapped_pass_through_route_with_server_root_path(server_root_pa
"litellm.proxy.utils.get_server_root_path", return_value=server_root_path
):
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/typesafe/decisions") is True
if server_root_path not in ("", "/"):
assert (
InitPassThroughEndpointHelpers.is_registered_pass_through_route(
f"{server_root_path.rstrip('/')}/typesafe/decisions"
)
is True
)
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/anthropic/v1/messages") is True
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/typesafeevil/decisions") is False
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/not-a-passthrough/decisions") is False
@pytest.mark.parametrize(
"root_path,request_path",
[
("/typesafe", "/typesafe/typesafe/decisions"),
("/api/v1", "/api/v1/typesafe/decisions"),
],
)
def test_typesafe_passthrough_dispatches_after_root_path_normalization(
monkeypatch: pytest.MonkeyPatch, root_path: str, request_path: str
):
"""A root-relative TypeSafe route passes auth, registry validation, and dispatch."""
from litellm.proxy.auth.auth_utils import get_request_route
from litellm.proxy.auth.user_api_key_auth import (
check_api_key_for_custom_headers_or_pass_through_endpoints,
user_api_key_auth,
)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
authenticated_routes: list[str] = []
async def authenticated_passthrough_request(request: Request) -> UserAPIKeyAuth:
route = get_request_route(request)
authenticated_routes.append(route)
assert await check_api_key_for_custom_headers_or_pass_through_endpoints(
request=request,
route=route,
pass_through_endpoints=None,
api_key="sk-test",
) == "sk-test"
return UserAPIKeyAuth(api_key="sk-test")
app = FastAPI(root_path=root_path)
app.add_api_route(
"/typesafe/{endpoint:path}",
create_pass_through_route(
endpoint="typesafe",
target="https://typesafe-upstream.test/decisions",
custom_headers={"Authorization": "Bearer upstream-key"},
custom_llm_provider="typesafe",
),
methods=["POST"],
)
app.dependency_overrides[user_api_key_auth] = authenticated_passthrough_request
with respx.mock(assert_all_called=True) as upstream:
upstream_request = upstream.post("https://typesafe-upstream.test/decisions").mock(
return_value=httpx.Response(200, json={"reached": "upstream"})
)
response = TestClient(app).post(request_path, json={"input": "test"})
assert response.status_code == 200
assert response.json() == {"reached": "upstream"}
assert authenticated_routes == ["/typesafe/decisions"]
assert upstream_request.called
@pytest.mark.asyncio
async def test_multipart_passthrough_preserves_boundary():
"""