mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): avoid double stripping passthrough roots
This commit is contained in:
parent
f1b587e480
commit
9243b381da
4 changed files with 78 additions and 52 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue