fix: extract normalize_route_for_root_path to deduplicate root-path stripping

This commit is contained in:
joereyna 2026-03-12 07:55:22 -07:00
parent 791e598ad5
commit 938452cc59
3 changed files with 18 additions and 27 deletions

View file

@ -50,7 +50,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body, _safe_get_request_headers,
populate_request_with_path_params)
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
from litellm.proxy.utils import PrismaClient, ProxyLogging, get_server_root_path
from litellm.proxy.utils import PrismaClient, ProxyLogging, normalize_route_for_root_path
from litellm.secret_managers.main import get_secret_bool
from litellm.types.services import ServiceTypes
@ -386,18 +386,10 @@ async def check_api_key_for_custom_headers_or_pass_through_endpoints(
api_key: str,
) -> Union[UserAPIKeyAuth, str]:
is_mapped_pass_through_route: bool = False
root_path = get_server_root_path()
if root_path and root_path != "/":
if route.startswith(root_path):
normalized_route = route[len(root_path):]
if normalized_route: # guard against route == root_path exactly
for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value: # type: ignore
if normalized_route.startswith(mapped_route):
is_mapped_pass_through_route = True
break
else:
normalized_route = normalize_route_for_root_path(route)
if normalized_route is not None:
for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value: # type: ignore
if route.startswith(mapped_route):
if normalized_route.startswith(mapped_route):
is_mapped_pass_through_route = True
break
if is_mapped_pass_through_route:

View file

@ -54,7 +54,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
)
from litellm.proxy.utils import get_server_root_path
from litellm.proxy.utils import get_server_root_path, normalize_route_for_root_path
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
@ -2058,21 +2058,10 @@ class InitPassThroughEndpointHelpers:
bool: True if route is a registered pass-through endpoint, False otherwise
"""
## CHECK IF MAPPED PASS THROUGH ENDPOINT
# When SERVER_ROOT_PATH is set, all valid routes carry that prefix.
# Strip it before comparing against mapped routes; if the route does not
# carry the prefix, it cannot be a mapped pass-through route.
root_path = get_server_root_path()
if root_path and root_path != "/":
if route.startswith(root_path):
normalized_route = route[len(root_path):]
if normalized_route: # guard against route == root_path exactly
for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value:
if normalized_route.startswith(mapped_route):
return True
# Route lacks expected prefix (or is exactly root_path) — not a mapped pass-through route
else:
normalized_route = normalize_route_for_root_path(route)
if normalized_route is not None:
for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value:
if route.startswith(mapped_route):
if normalized_route.startswith(mapped_route):
return True
# Fast path: check if any registered route key contains this path

View file

@ -5145,6 +5145,16 @@ def get_server_root_path() -> str:
return os.getenv("SERVER_ROOT_PATH", "")
def normalize_route_for_root_path(route: str) -> Optional[str]:
"""Strip SERVER_ROOT_PATH prefix. Returns de-prefixed route, or None if route is not under root path."""
root_path = get_server_root_path()
if root_path and root_path != "/":
if route.startswith(root_path + "/"):
return route[len(root_path):]
return None
return route
def get_prisma_client_or_throw(message: str):
from litellm.proxy.proxy_server import prisma_client