This commit is contained in:
vgvr0 2026-09-30 16:54:07 -04:00 • committed by GitHub
commit bbc8f67d98
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 99 additions and 48 deletions

View file

@ -137,7 +137,6 @@ from litellm.proxy.spend_tracking.spend_counter_batch import (
from litellm.proxy.utils import (
PrismaClient,
ProxyLogging,
normalize_route_for_root_path,
)
from litellm.repositories.table_repositories import TeamMembershipRepository
from litellm.router_utils.common_utils import resolve_model_group_alias
@ -865,12 +864,14 @@ 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 = normalize_route_for_root_path(route)
if normalized_route is not None:
for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value:
if normalized_route.startswith(mapped_route):
is_mapped_pass_through_route = True
break
# ``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
break
if is_mapped_pass_through_route:
if request.headers.get("litellm_user_api_key") is not None:
api_key = request.headers.get("litellm_user_api_key") or ""

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 normalize_route_for_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,14 +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 should pass routes from
``get_request_route()`` (already stripped); prefixed ``request.url.path``
values are stripped via ``normalize_route_for_root_path``.
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``).
"""
normalized_route: Final = normalize_route_for_root_path(route)
return normalized_route if normalized_route is not None else route
return route
@staticmethod
def is_registered_pass_through_route(route: str) -> bool:
@ -3167,11 +3166,10 @@ class InitPassThroughEndpointHelpers:
bool: True if route is a registered pass-through endpoint, False otherwise
"""
## CHECK IF MAPPED PASS THROUGH ENDPOINT
normalized_route: Final = normalize_route_for_root_path(route)
if normalized_route is not None:
for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value:
if normalized_route.startswith(mapped_route):
return True
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
comparison_route: Final = InitPassThroughEndpointHelpers._route_for_registry_lookup(route)

View file

@ -41,7 +41,6 @@ from typing import (
Protocol,
TypeAlias,
TypeVar,
Union,
cast,
overload,
)

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,7 @@ 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", "exact", "/ml", True),
("/llmproxy", "exact", "/llmproxy/ml", True),
("/llmproxy", "subpath", "/other/api", False),
],
)
@ -3553,8 +3541,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,15 +3582,80 @@ 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
# bare route without prefix should not match when root is set
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/vertex_ai/v1/projects/foo") is False
@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,
)
with patch(
"litellm.proxy.utils.get_server_root_path", return_value=server_root_path
):
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/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