diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 5194f62cf78..7acd16b915d 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 "" diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 2e7f9c4a41c..539ece6c3cc 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index c0e36e6e172..218db1455b8 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -41,7 +41,6 @@ from typing import ( Protocol, TypeAlias, TypeVar, - Union, cast, overload, ) 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 89feb2b6426..9ca01770276 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 @@ -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