mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): handle root path for typesafe passthrough
This commit is contained in:
parent
cede93e826
commit
f1b587e480
4 changed files with 46 additions and 20 deletions
|
|
@ -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 ""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue