diff --git a/litellm/constants.py b/litellm/constants.py index d58fc8a6318..db122cbb37c 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -123,6 +123,9 @@ MAX_BASE64_LENGTH_STDOUT_LOG: Final = get_env_int("MAX_BASE64_LENGTH_STDOUT_LOG" # When true, adds detailed per-phase timing breakdown headers to responses. # Headers: x-litellm-timing-{pre-processing,llm-api,post-processing,message-copy}-ms LITELLM_DETAILED_TIMING: Final = os.getenv("LITELLM_DETAILED_TIMING", "false").lower() == "true" +LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST: Final = ( + os.getenv("LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST", "false").lower() == "true" +) # Model cost map validation constants MODEL_COST_MAP_MIN_MODEL_COUNT: Final = int( diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index cb6d47ca4c8..e53645f1774 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -24,7 +24,8 @@ from starlette.routing import BaseRoute, Match from starlette.types import ASGIApp, Lifespan, Receive, Scope, Send from litellm._logging import verbose_proxy_logger -from litellm.proxy.route_priority import hot_routes_first +from litellm.constants import LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST +from litellm.proxy.route_priority import configured_pass_through_routes_first, hot_routes_first if TYPE_CHECKING: from fastapi import APIRouter, FastAPI @@ -442,6 +443,20 @@ def _in_registry_order( ) +def _route_table( + app: "FastAPI", + routes: Sequence[BaseRoute], + lazy_routes: Mapping[str, tuple[BaseRoute, ...]], + features: tuple[LazyFeature, ...], +) -> list[BaseRoute]: # mutable-ok: assigned to Router.routes, a list + ordered: Final = _in_registry_order(routes, lazy_routes, features, _lazy_slots(app)) + if LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST: + return hot_routes_first(ordered) + return hot_routes_first( + configured_pass_through_routes_first(ordered, tuple(chain.from_iterable(lazy_routes.values()))) + ) + + async def _force_load(app: "FastAPI", feat: LazyFeature, features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> bool: """Import + register a lazy feature exactly once per (app, module). Shared by the middleware and the /lazy/warm endpoint.""" @@ -488,8 +503,8 @@ def _register_feature(app: "FastAPI", feat: LazyFeature, module: object, feature {**_lazy_routes(app), feat.module_path: tuple(app.router.routes[before:])} ) app.state.lazy_routes = lazy_routes # rebind-ok: the app owns the record of which routes each feature added - app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table - _in_registry_order(app.router.routes, lazy_routes, features, _lazy_slots(app)) + app.router.routes[:] = _route_table( # rebind-ok: the app owns its route table + app, app.router.routes, lazy_routes, features ) _lazy_loaded(app).add(feat.module_path) app.openapi_schema = None @@ -553,8 +568,8 @@ def _restore_registry_order(app: "FastAPI", features: tuple[LazyFeature, ...]) - still_routed: Final = MappingProxyType( {module_path: tuple(r for r in routes if id(r) in present) for module_path, routes in _lazy_routes(app).items()} ) - app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table - _in_registry_order(app.router.routes, still_routed, features, _lazy_slots(app)) + app.router.routes[:] = _route_table( # rebind-ok: the app owns its route table + app, app.router.routes, still_routed, features ) app.openapi_schema = None diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index c3c1032a3f0..17daec7949a 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -2045,8 +2045,10 @@ def request_dispatched_to_pass_through_endpoint(request: Request | None) -> bool Reads the marker set by ``create_pass_through_route`` off the dispatched endpoint (``request.scope["endpoint"]``). Because routing has already run by the time auth dependencies execute, this reflects the handler that actually serves the request: - a custom path colliding with a built-in route resolves to the built-in handler, - which carries no marker, so model-access checks are never wrongly skipped. + an exact-path collision or a built-in that is not a ``//{name:path}`` + catch-all still resolves to the built-in handler, which carries no marker. A + configured route shadowed by such a catch-all is dispatched ahead of it unless + ``LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST`` is set. """ if request is None: return False @@ -2081,10 +2083,12 @@ def get_model_from_request( ``model`` field there names an upstream model, not a LiteLLM-managed one, and enforcing key/team model allowlists against it would reject valid requests. The check reads the FastAPI-resolved endpoint (``request.scope["endpoint"]``), not the - request path, so a custom path that collides with a built-in route never - suppresses model-access checks: on a collision the built-in handler is dispatched - and does not carry the marker. Built-in provider passthrough routes - (``/vertex_ai``, ``/gemini``, ...) are separate handlers and keep model enforcement. + request path, so only a dispatched configured handler suppresses model-access + checks: exact-path collisions and built-ins that are not ``//{name:path}`` + catch-alls still resolve to the built-in handler, which carries no marker. A + configured route shadowed by such a catch-all is dispatched ahead of it unless + ``LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST`` is set. Built-in provider passthrough + routes (``/vertex_ai``, ``/gemini``, ...) are separate handlers and keep model enforcement. """ if request_dispatched_to_pass_through_endpoint(request): return None diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index f5ba7f9e877..19be340a57d 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -8,7 +8,7 @@ from base64 import b64encode from collections.abc import AsyncGenerator, AsyncIterator, Callable, Iterable, Mapping, Sequence from dataclasses import dataclass from datetime import datetime -from itertools import count, groupby +from itertools import chain, count, groupby from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, TypedDict, cast from urllib.parse import urlencode, urlparse @@ -42,6 +42,7 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.constants import ( + LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST, MAXIMUM_TRACEBACK_LINES_TO_LOG, PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS, REDACTED_BY_LITELLM, @@ -112,6 +113,7 @@ from litellm.proxy.litellm_pre_call_utils import ( _strip_client_pricing_overrides, # pyright: ignore[reportPrivateUsage] # sanitize before trusted hooks add guardrail costs ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError +from litellm.proxy.route_priority import configured_pass_through_routes_first, shadowed_by_builtin_claim 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 @@ -3013,6 +3015,15 @@ class SafeRouteAdder: app.router.routes[:] = _placed_ahead( # rebind-ok: the app owns its route table app.router.routes, app.router.routes[-1], shadowed[0] ) + lazy_routes: Final[Mapping[str, tuple[BaseRoute, ...]] | None] = getattr( + getattr(app, "state", None), "lazy_routes", None + ) + if lazy_routes is not None and not LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST: + builtin: Final = tuple(chain.from_iterable(lazy_routes.values())) + if shadowed_by_builtin_claim(path, builtin): + app.router.routes[:] = configured_pass_through_routes_first( # rebind-ok: the app owns its route table + app.router.routes, builtin + ) verbose_proxy_logger.debug( "Successfully added route: %s with methods %s", path, diff --git a/litellm/proxy/route_priority.py b/litellm/proxy/route_priority.py index 77815b678a6..3eb7d85e395 100644 --- a/litellm/proxy/route_priority.py +++ b/litellm/proxy/route_priority.py @@ -1,10 +1,14 @@ """Starlette matches routes in registration order, so the routes that take the most traffic go first.""" -from collections.abc import Sequence +import re +from collections.abc import Collection, Sequence +from itertools import chain from typing import Final from starlette.routing import BaseRoute, Route +from litellm.types.passthrough_endpoints.pass_through_endpoints import LITELLM_PASS_THROUGH_ENDPOINT_MARKER + HOT_ROUTE_PATHS: Final[frozenset[str]] = frozenset( ( "/health/liveliness", @@ -22,3 +26,61 @@ def _is_hot(route: BaseRoute) -> bool: def hot_routes_first(routes: Sequence[BaseRoute]) -> list[BaseRoute]: # mutable-ok: assigned to Router.routes, a list return sorted(routes, key=lambda route: not _is_hot(route)) + + +BUILTIN_PREFIX_CLAIM: Final[re.Pattern[str]] = re.compile(r"^(/[^{}]+/)\{[A-Za-z_][A-Za-z0-9_]*:path\}$") + + +def _claim_prefix(route: BaseRoute) -> str | None: + match: Final = BUILTIN_PREFIX_CLAIM.match(getattr(route, "path", "")) + return match.group(1) if match else None + + +def _prefix_candidates(path: str) -> tuple[str, ...]: + return tuple(path[: k + 1] for k in range(1, len(path)) if path[k] == "/") + + +def shadowed_by_builtin_claim(path: str, builtin_routes: Collection[BaseRoute]) -> bool: + candidates: Final = frozenset(_prefix_candidates(path)) + return any((prefix := _claim_prefix(route)) is not None and prefix in candidates for route in builtin_routes) + + +def configured_pass_through_routes_first( + routes: Sequence[BaseRoute], builtin_routes: Collection[BaseRoute] +) -> list[BaseRoute]: # mutable-ok: assigned to Router.routes, a list + builtin_ids: Final = frozenset(id(route) for route in builtin_routes) + + def is_configured(route: BaseRoute) -> bool: + return isinstance(route, Route) and getattr(route.endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, False) is True + + def claim_prefix(index: int) -> str | None: + if id(routes[index]) not in builtin_ids: + return None + return _claim_prefix(routes[index]) + + claimed: Final = tuple( + (index, prefix) for index in range(len(routes) - 1, -1, -1) if (prefix := claim_prefix(index)) is not None + ) + earliest_claim: Final = {prefix: index for index, prefix in claimed} + + def hoist_target(index: int) -> int | None: + configured_path: Final = getattr(routes[index], "path", "") + hits: Final = tuple( + earliest_claim[candidate] + for candidate in _prefix_candidates(configured_path) + if candidate in earliest_claim + ) + target: Final = min(hits) if hits else len(routes) + return target if target < index else None + + targets: Final = {index: hoist_target(index) for index in range(len(routes)) if is_configured(routes[index])} + moved: Final = frozenset(index for index, target in targets.items() if target is not None) + return list( + chain.from_iterable( + ( + *tuple(routes[j] for j in targets if targets[j] == index), + *(() if index in moved else (route,)), + ) + for index, route in enumerate(routes) + ) + ) diff --git a/tests/documentation_tests/test_env_keys.py b/tests/documentation_tests/test_env_keys.py index 0713abc86d0..2a7da58f39c 100644 --- a/tests/documentation_tests/test_env_keys.py +++ b/tests/documentation_tests/test_env_keys.py @@ -31,6 +31,7 @@ EXCLUDED_GUARD_ONLY_VARS = { EXCLUDED_ROLLOUT_FLAGS = { "LITELLM_USE_RUST_OCR", "LITELLM_RUST", + "LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST", } # Internal infrastructure tuning parameters for streaming/queue management diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 31321254d94..1c46fb42b31 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -53,9 +53,9 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.types import utils as types_utils from litellm.types.passthrough_endpoints.pass_through_endpoints import ( - EndpointType, LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, + EndpointType, ) MESSAGE_START_SSE_FRAME = b'event: message_start\ndata: {"type": "message_start"}\n\n' @@ -784,6 +784,155 @@ def test_add_subpath_route(): assert callable(call_args["endpoint"]) +def test_add_subpath_route_on_a_loaded_lazy_app_outranks_the_builtin_prefix_route(): + """ + A DB pass-through endpoint created after a lazy feature loaded must still + dispatch ahead of that feature's overlapping built-in prefix route. + """ + from fastapi import FastAPI + from starlette.routing import Match + + from litellm.types.passthrough_endpoints.pass_through_endpoints import ( + LITELLM_PASS_THROUGH_ENDPOINT_MARKER, + ) + + saved_registry: Final = dict(_registered_pass_through_routes) + try: + app = FastAPI() + + async def builtin(): + return {"handler": "builtin"} + + app.add_api_route("/typesafe/{endpoint:path}", builtin, methods=["GET"]) + builtin_route = next( + route for route in app.router.routes if getattr(route, "path", "") == "/typesafe/{endpoint:path}" + ) + app.state.lazy_routes = MappingProxyType({"tests.fake_lazy_feature": (builtin_route,)}) + + InitPassThroughEndpointHelpers.add_subpath_route( + app=app, + path="/typesafe", + target="http://example.com", + custom_headers=None, + forward_headers=None, + merge_query_params=None, + dependencies=None, + cost_per_request=None, + endpoint_id="test-runtime-endpoint-id", + methods=["GET"], + ) + + scope = { + "type": "http", + "method": "GET", + "path": "/typesafe/health", + "root_path": "", + "headers": [], + "query_string": b"", + } + matched = next(route for route in app.router.routes if route.matches(dict(scope))[0] is Match.FULL) + assert getattr(matched.endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, False) is True, ( + "the first full match under /typesafe must be the configured pass-through route" + ) + finally: + _registered_pass_through_routes.clear() + _registered_pass_through_routes.update(saved_registry) + + +def test_add_subpath_route_respects_the_builtin_routes_first_flag(monkeypatch: pytest.MonkeyPatch): + """ + With LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST set, a runtime-added configured + route keeps losing to the loaded feature's overlapping built-in prefix route. + """ + from fastapi import FastAPI + from starlette.routing import Match + + import litellm.proxy.pass_through_endpoints.pass_through_endpoints as pass_through_endpoints_module + + monkeypatch.setattr(pass_through_endpoints_module, "LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST", True) + + saved_registry: Final = dict(_registered_pass_through_routes) + try: + app = FastAPI() + + async def builtin(): + return {"handler": "builtin"} + + app.add_api_route("/typesafe/{endpoint:path}", builtin, methods=["GET"]) + builtin_route = next( + route for route in app.router.routes if getattr(route, "path", "") == "/typesafe/{endpoint:path}" + ) + app.state.lazy_routes = MappingProxyType({"tests.fake_lazy_feature": (builtin_route,)}) + + InitPassThroughEndpointHelpers.add_subpath_route( + app=app, + path="/typesafe", + target="http://example.com", + custom_headers=None, + forward_headers=None, + merge_query_params=None, + dependencies=None, + cost_per_request=None, + endpoint_id="test-runtime-endpoint-id-flag", + methods=["GET"], + ) + + scope = { + "type": "http", + "method": "GET", + "path": "/typesafe/health", + "root_path": "", + "headers": [], + "query_string": b"", + } + matched = next(route for route in app.router.routes if route.matches(dict(scope))[0] is Match.FULL) + assert matched.endpoint is builtin, "the flag keeps the built-in prefix route first" + finally: + _registered_pass_through_routes.clear() + _registered_pass_through_routes.update(saved_registry) + + +def test_add_subpath_route_for_an_unrelated_prefix_leaves_the_route_order_untouched(): + """ + A runtime-added configured route under a prefix no loaded feature claims must + not pay for or trigger the full route-table reorder. + """ + from fastapi import FastAPI + + saved_registry: Final = dict(_registered_pass_through_routes) + try: + app = FastAPI() + + async def builtin(): + return {"handler": "builtin"} + + app.add_api_route("/typesafe/{endpoint:path}", builtin, methods=["GET"]) + builtin_route = next( + route for route in app.router.routes if getattr(route, "path", "") == "/typesafe/{endpoint:path}" + ) + app.state.lazy_routes = MappingProxyType({"tests.fake_lazy_feature": (builtin_route,)}) + before = list(app.router.routes) + + InitPassThroughEndpointHelpers.add_subpath_route( + app=app, + path="/other", + target="http://example.com", + custom_headers=None, + forward_headers=None, + merge_query_params=None, + dependencies=None, + cost_per_request=None, + endpoint_id="test-runtime-endpoint-id-unrelated", + methods=["GET"], + ) + + assert list(app.router.routes)[:-1] == before, "an unrelated add must not reorder existing routes" + assert len(app.router.routes) == len(before) + 1 + finally: + _registered_pass_through_routes.clear() + _registered_pass_through_routes.update(saved_registry) + + @pytest.mark.asyncio async def test_pass_through_handler_rejects_unregistered_method(): """ @@ -1694,7 +1843,9 @@ async def test_pass_through_request_streamed_response_is_owned_by_the_caller(): cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler))) mock_proxy_logging = MagicMock() - mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type, endpoint_type: data) + mock_proxy_logging.pre_call_hook = AsyncMock( + side_effect=lambda user_api_key_dict, data, call_type, endpoint_type: data + ) mock_proxy_logging.post_call_failure_hook = AsyncMock() mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) mock_proxy_logging.get_proxy_hook = MagicMock(return_value=MagicMock()) @@ -2742,12 +2893,8 @@ async def test_pass_through_request_follows_redirect_to_final_response(httpx_tra mock_user_api_key_dict = MagicMock() with respx.mock(assert_all_called=True) as upstream: - upstream.get("https://upstream.test/redirect/1").respond( - 302, headers={"Location": "/get"} - ) - upstream.get("https://upstream.test/get").respond( - 200, json={"url": "https://upstream.test/get"} - ) + upstream.get("https://upstream.test/redirect/1").respond(302, headers={"Location": "/get"}) + upstream.get("https://upstream.test/get").respond(200, json={"url": "https://upstream.test/get"}) response = await pass_through_request( request=mock_request, @@ -7989,9 +8136,7 @@ async def test_a_deleted_db_pass_through_stops_serving_on_the_next_db_sync(tmp_p @pytest.mark.asyncio -async def test_config_pass_through_reads_its_custom_key_header_when_the_db_holds_pass_throughs( - tmp_path, monkeypatch -): +async def test_config_pass_through_reads_its_custom_key_header_when_the_db_holds_pass_throughs(tmp_path, monkeypatch): proxy: Final = await _boot_db_backed_proxy( tmp_path, monkeypatch, diff --git a/tests/unit/proxy/test__lazy_features.py b/tests/unit/proxy/test__lazy_features.py index d2bb8244c2f..564957991e7 100644 --- a/tests/unit/proxy/test__lazy_features.py +++ b/tests/unit/proxy/test__lazy_features.py @@ -15,7 +15,9 @@ from litellm.proxy._lazy_features import ( attach_lazy_features, lazy_tag_to_prefix, loaded_lazy_modules, + reserve_lazy_slot, ) +from litellm.types.passthrough_endpoints.pass_through_endpoints import LITELLM_PASS_THROUGH_ENDPOINT_MARKER FLAG: Final = "LITELLM_DISABLE_LAZY_ROUTES" WARMUP_PATH: Final = "/lazy/warm/{name}" @@ -119,6 +121,157 @@ def test_flag_lets_a_route_added_during_startup_beat_an_overlapping_feature_rout ) +def test_configured_pass_through_added_after_attach_beats_a_reserved_slot_feature( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv(FLAG, raising=False) + zeta: Final = _feature_module(monkeypatch, "zeta", "/zeta/{endpoint:path}") + features: Final = (LazyFeature(name="zeta", module_path=zeta.module_path, path_prefixes=("/zeta",)),) + + async def early() -> dict[str, str]: + return {"feature": "early"} + + async def late() -> dict[str, str]: + return {"feature": "late"} + + async def configured() -> dict[str, str]: + return {"feature": "configured"} + + setattr(configured, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True) + + app: Final = FastAPI() + app.add_api_route("/early", early, methods=["GET"]) + reserve_lazy_slot(app, "zeta", features=features) + app.add_api_route("/late", late, methods=["GET"]) + attach_lazy_features(app, features) + app.add_api_route("/zeta/{subpath:path}", configured, methods=["GET"]) + + with TestClient(app) as client: + assert client.get("/zeta/health").json() == {"feature": "configured"}, ( + "a deployment's own pass-through endpoint must outrank the built-in provider route" + ) + + +def test_flag_configured_pass_through_added_during_startup_beats_a_reserved_slot_feature( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv(FLAG, "true") + features: Final = (_feature_module(monkeypatch, "zeta", "/zeta/{endpoint:path}"),) + + async def early() -> dict[str, str]: + return {"feature": "early"} + + async def late() -> dict[str, str]: + return {"feature": "late"} + + async def configured() -> dict[str, str]: + return {"feature": "configured"} + + setattr(configured, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True) + + @asynccontextmanager + async def adds_a_pass_through(app_: FastAPI) -> AsyncGenerator[None]: + app_.add_api_route("/zeta/{subpath:path}", configured, methods=["GET"]) + yield + + app: Final = FastAPI(lifespan=adds_a_pass_through) + app.add_api_route("/early", early, methods=["GET"]) + reserve_lazy_slot(app, "zeta", features=features) + app.add_api_route("/late", late, methods=["GET"]) + attach_lazy_features(app, features) + + with TestClient(app) as client: + assert client.get("/zeta/health").json() == {"feature": "configured"}, ( + "a deployment's own pass-through endpoint must outrank the built-in provider route" + ) + + +def test_builtin_routes_first_flag_keeps_a_reserved_slot_feature_ahead_of_a_configured_route( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm.proxy._lazy_features as lazy_features + + monkeypatch.delenv(FLAG, raising=False) + monkeypatch.setattr(lazy_features, "LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST", True) + zeta: Final = _feature_module(monkeypatch, "zeta", "/zeta/{endpoint:path}") + features: Final = (LazyFeature(name="zeta", module_path=zeta.module_path, path_prefixes=("/zeta",)),) + + async def early() -> dict[str, str]: + return {"feature": "early"} + + async def configured() -> dict[str, str]: + return {"feature": "configured"} + + setattr(configured, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True) + + app: Final = FastAPI() + app.add_api_route("/early", early, methods=["GET"]) + reserve_lazy_slot(app, "zeta", features=features) + attach_lazy_features(app, features) + app.add_api_route("/zeta/{subpath:path}", configured, methods=["GET"]) + + with TestClient(app) as client: + assert client.get("/zeta/health").json() == {"feature": "zeta"}, ( + "LITELLM_BUILTIN_PASS_THROUGH_ROUTES_FIRST restores the old built-in-first order" + ) + + +def test_late_route_without_the_marker_still_loses_to_a_reserved_slot_feature( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv(FLAG, raising=False) + zeta: Final = _feature_module(monkeypatch, "zeta", "/zeta/{endpoint:path}") + features: Final = (LazyFeature(name="zeta", module_path=zeta.module_path, path_prefixes=("/zeta",)),) + + async def early() -> dict[str, str]: + return {"feature": "early"} + + async def late() -> dict[str, str]: + return {"feature": "late"} + + app: Final = FastAPI() + app.add_api_route("/early", early, methods=["GET"]) + reserve_lazy_slot(app, "zeta", features=features) + app.add_api_route("/late", late, methods=["GET"]) + attach_lazy_features(app, features) + app.add_api_route("/zeta/{subpath:path}", late, methods=["GET"]) + + with TestClient(app) as client: + assert client.get("/zeta/health").json() == {"feature": "zeta"}, ( + "only configured pass-through routes are promoted; ordinary late routes keep losing" + ) + + +def test_configured_pass_through_under_a_different_segment_is_not_moved( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv(FLAG, raising=False) + zeta: Final = _feature_module(monkeypatch, "zeta", "/zeta/{endpoint:path}") + features: Final = (LazyFeature(name="zeta", module_path=zeta.module_path, path_prefixes=("/zeta",)),) + + async def configured() -> dict[str, str]: + return {"feature": "configured"} + + setattr(configured, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True) + + app: Final = FastAPI() + reserve_lazy_slot(app, "zeta", features=features) + attach_lazy_features(app, features) + app.add_api_route("/omega/{subpath:path}", configured, methods=["GET"]) + + with TestClient(app) as client: + assert client.get("/zeta/health").json() == {"feature": "zeta"} + + routes: Final = app.router.routes + omega_index: Final = next( + i for i, route in enumerate(routes) if getattr(route, "path", "") == "/omega/{subpath:path}" + ) + zeta_index: Final = next( + i for i, route in enumerate(routes) if getattr(route, "path", "") == "/zeta/{endpoint:path}" + ) + assert omega_index > zeta_index, "an unrelated configured route keeps its position after the feature routes" + + def test_flag_does_not_bring_back_a_feature_route_removed_during_startup(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv(FLAG, "true") features: Final = ( diff --git a/tests/unit/proxy/test_route_priority.py b/tests/unit/proxy/test_route_priority.py index dfdc816f4b4..702a955446b 100644 --- a/tests/unit/proxy/test_route_priority.py +++ b/tests/unit/proxy/test_route_priority.py @@ -1,13 +1,20 @@ import sys +from collections.abc import Callable, Sequence from types import ModuleType import httpx import pytest from fastapi import APIRouter, FastAPI from fastapi.testclient import TestClient -from starlette.routing import Match +from starlette.routing import BaseRoute, Match, Route -from litellm.proxy.route_priority import HOT_ROUTE_PATHS, hot_routes_first +from litellm.proxy.route_priority import ( + HOT_ROUTE_PATHS, + configured_pass_through_routes_first, + hot_routes_first, + shadowed_by_builtin_claim, +) +from litellm.types.passthrough_endpoints.pass_through_endpoints import LITELLM_PASS_THROUGH_ENDPOINT_MARKER FILLER_COUNT = 300 @@ -22,6 +29,34 @@ def _routes_scanned_before_dispatch(app: FastAPI, method: str, path: str) -> int raise AssertionError(f"{method} {path} has no route") +def _dispatch(routes: Sequence[BaseRoute], method: str, path: str) -> Callable[..., object]: + """Endpoint of the first route a real request to method+path would hit.""" + scope = {"type": "http", "method": method, "path": path, "root_path": "", "headers": [], "query_string": b""} + for route in routes: + match, _ = route.matches(dict(scope)) + if match == Match.FULL: + return route.endpoint + raise AssertionError(f"{method} {path} has no route") + + +def _endpoint(name: str) -> Callable[..., object]: + async def handler() -> dict[str, str]: + return {"handler": name} + + handler.__name__ = name + return handler + + +def _configured(path: str, name: str) -> Route: + endpoint = _endpoint(name) + setattr(endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True) + return Route(path, endpoint, methods=["GET"]) + + +def _builtin(path: str, name: str) -> Route: + return Route(path, _endpoint(name), methods=["GET"]) + + def _hot_router() -> APIRouter: router = APIRouter() @@ -169,3 +204,57 @@ def test_proxy_app_dispatches_liveness_and_chat_completions_before_the_rest(): assert _routes_scanned_before_dispatch(app, "GET", "/health/liveness") <= hot_count assert _routes_scanned_before_dispatch(app, "POST", "/v1/chat/completions") <= hot_count assert _routes_scanned_before_dispatch(app, "POST", "/chat/completions") <= hot_count + + +def test_deeper_builtin_prefix_claim_still_beats_a_configured_catch_all(): + builtin = _builtin("/zeta/inner/{endpoint:path}", "builtin") + configured = _configured("/zeta/{subpath:path}", "configured") + + reordered = configured_pass_through_routes_first([builtin, configured], [builtin]) + + assert _dispatch(reordered, "GET", "/zeta/inner/x") is builtin.endpoint + + +def test_an_exact_builtin_still_beats_an_identical_configured_path(): + builtin = _builtin("/zeta/same", "builtin") + configured = _configured("/zeta/same", "configured") + + reordered = configured_pass_through_routes_first([builtin, configured], [builtin]) + + assert _dispatch(reordered, "GET", "/zeta/same") is builtin.endpoint + + +def test_an_exact_configured_route_beats_a_builtin_prefix_claim(): + builtin = _builtin("/zeta/{endpoint:path}", "builtin") + configured = _configured("/zeta/v1/thing", "configured") + + reordered = configured_pass_through_routes_first([builtin, configured], [builtin]) + + assert _dispatch(reordered, "GET", "/zeta/v1/thing") is configured.endpoint + + +def test_a_builtin_with_a_trailing_literal_is_not_a_prefix_claim(): + builtin = _builtin("/zeta/{id:path}/search", "builtin") + configured = _configured("/zeta/{subpath:path}", "configured") + + reordered = configured_pass_through_routes_first([builtin, configured], [builtin]) + + assert _dispatch(reordered, "GET", "/zeta/a/search") is builtin.endpoint + + +def test_shadowed_by_builtin_claim_is_true_under_a_builtin_prefix_claim(): + builtin = _builtin("/typesafe/{endpoint:path}", "builtin") + + assert shadowed_by_builtin_claim("/typesafe/{subpath:path}", (builtin,)) is True + + +def test_shadowed_by_builtin_claim_is_false_under_a_non_claim_builtin(): + builtin = _builtin("/zeta/{id:path}/search", "builtin") + + assert shadowed_by_builtin_claim("/zeta/{subpath:path}", (builtin,)) is False + + +def test_shadowed_by_builtin_claim_is_false_for_an_unrelated_prefix(): + builtin = _builtin("/typesafe/{endpoint:path}", "builtin") + + assert shadowed_by_builtin_claim("/other/{subpath:path}", (builtin,)) is False