diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e3ce9bcd850..6b5439b5554 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -134,7 +134,7 @@ from litellm.proxy.spend_tracking.spend_counter_batch import ( from litellm.proxy.utils import ( PrismaClient, ProxyLogging, - normalize_route_for_root_path, + strip_server_root_path, ) from litellm.repositories.table_repositories import TeamMembershipRepository from litellm.router_utils.common_utils import resolve_model_group_alias @@ -860,12 +860,11 @@ 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 + normalized_route: Final = strip_server_root_path(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 "" diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e0a4184291e..ce5da29ce82 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -102,7 +102,7 @@ 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.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 @@ -3145,12 +3145,10 @@ class InitPassThroughEndpointHelpers: """ Normalize an incoming route to the bare path stored in the registry. - 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 store root-stripped paths. Callers may pass routes from + ``get_request_route()`` or prefixed ``request.url.path`` values. """ - normalized_route: Final = normalize_route_for_root_path(route) - return normalized_route if normalized_route is not None else route + return strip_server_root_path(route) @staticmethod def is_registered_pass_through_route(route: str) -> bool: @@ -3167,11 +3165,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 = strip_server_root_path(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) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index ea294b76e92..62bc8435064 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -8330,6 +8330,14 @@ 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 diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 3469df082e0..7b12f6dde81 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -3601,8 +3601,30 @@ def test_mapped_pass_through_routes_with_server_root_path(): ) assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/litellm/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 + # get_request_route() supplies bare paths after stripping root_path. + assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/vertex_ai/v1/projects/foo") is True + + +@pytest.mark.parametrize("server_root_path", ["", "/", "/api/v1", "/api/v1/"]) +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 + 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.asyncio