diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index d58bfbbdb5..9f03457522 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -257,6 +257,13 @@ class LazyFeatureMiddleware: self.app = app self._fastapi_app = fastapi_app self._features = features + # SERVER_ROOT_PATH is a process-startup env var, cache the normalized + # form once instead of recomputing per request. Lazy import to avoid + # pulling proxy.utils into this module's import graph at startup + # (proxy_server imports both). + from litellm.proxy.utils import get_server_root_path + + self._root_path = get_server_root_path().rstrip("/") # Loaded set / per-feature locks live on app.state so the warm endpoint # and the middleware share them — preventing duplicate registrations # when both paths fire for the same feature. @@ -274,6 +281,15 @@ class LazyFeatureMiddleware: self._features ): path = scope.get("path", "") + # Strip SERVER_ROOT_PATH so prefix matching works under a server + # root path. Without this, requests like /api/v1/policies/... never + # match the registered prefixes (/policies/...) and lazy features + # stay unloaded — every endpoint under them returns 404. The + # `+ "/"` boundary prevents false-positive matches (e.g. /apiv2 + # against root /api). If the path doesn't start with the prefix + # (e.g. a reverse proxy already stripped it), we leave it alone. + if self._root_path and path.startswith(self._root_path + "/"): + path = path[len(self._root_path) :] for feat in self._features: if feat.module_path in self._loaded: continue diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 859594f7a0..34b133c583 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -6378,6 +6378,84 @@ class TestLazyFeatureMiddleware: ) assert loads == ["json"] + @pytest.mark.asyncio + @pytest.mark.parametrize( + "server_root_path,request_path,should_load,case", + [ + # SERVER_ROOT_PATH set: incoming path includes prefix → strip and match. + ("/api/v1", "/api/v1/dummy/x", True, "root_path strip + match"), + # Trailing-slash env var must be normalized. + ("/api/v1/", "/api/v1/dummy/x", True, "trailing-slash env normalization"), + # Reverse proxy already stripped the prefix → original path still matches. + ("/api/v1", "/dummy/x", True, "pre-stripped path still loads"), + # No SERVER_ROOT_PATH set → unchanged behavior. + ("", "/dummy/x", True, "no root path"), + # SERVER_ROOT_PATH=/ must be a no-op (not strip every leading slash). + ("/", "/dummy/x", True, "root_path='/' is no-op"), + # Boundary check: /apiv2 must not match root /api. + ("/api", "/apiv2/foo", False, "boundary check prevents false match"), + # Genuine non-match under root_path. + ("/api/v1", "/api/v1/unrelated", False, "unrelated path under root"), + ], + ) + async def test_root_path_handling( + self, monkeypatch, server_root_path, request_path, should_load, case + ): + """ + The middleware must strip SERVER_ROOT_PATH before prefix-matching so + lazy features load under deployments that set a server root path, + while handling boundary, trailing-slash, and reverse-proxy edge cases + correctly. + """ + from fastapi import FastAPI + + from litellm.proxy._lazy_features import ( + LazyFeature, + LazyFeatureMiddleware, + ) + + monkeypatch.setenv("SERVER_ROOT_PATH", server_root_path) + + loads = [] + + def fake_register(app, module): + loads.append(getattr(module, "__name__", "?")) + + feat = LazyFeature( + name=f"dummy_{case}", + module_path="json", + path_prefixes=("/dummy",), + register_fn=fake_register, + ) + + async def downstream(scope, receive, send): + await send({"type": "http.response.start", "status": 200, "headers": []}) + await send({"type": "http.response.body", "body": b""}) + + target_app = FastAPI() + mw = LazyFeatureMiddleware(downstream, fastapi_app=target_app, features=(feat,)) + + async def receive(): + return {"type": "http.request", "body": b"", "more_body": False} + + async def send(message): + pass + + await mw( + { + "type": "http", + "path": request_path, + "method": "GET", + "headers": [], + }, + receive, + send, + ) + if should_load: + assert loads == ["json"], f"{case}: expected feature to load" + else: + assert loads == [], f"{case}: feature must not load" + @pytest.mark.asyncio async def test_concurrent_first_requests_only_register_once(self): """