diff --git a/.github/workflows/test-litellm.yml b/.github/workflows/test-litellm.yml index 6583844a5f9..a2b9e6c7c34 100644 --- a/.github/workflows/test-litellm.yml +++ b/.github/workflows/test-litellm.yml @@ -29,7 +29,11 @@ jobs: run: | poetry install --with dev,proxy-dev --extras proxy poetry run pip install pytest-xdist - + - name: Setup litellm-enterprise as local package + run: | + cd enterprise + python -m pip install -e . + cd .. - name: Run tests run: | poetry run pytest tests/litellm -x -vv -n 4 \ No newline at end of file diff --git a/enterprise/proxy/enterprise_routes.py b/enterprise/litellm_enterprise/proxy/enterprise_routes.py similarity index 100% rename from enterprise/proxy/enterprise_routes.py rename to enterprise/litellm_enterprise/proxy/enterprise_routes.py diff --git a/enterprise/proxy/readme.md b/enterprise/litellm_enterprise/proxy/readme.md similarity index 100% rename from enterprise/proxy/readme.md rename to enterprise/litellm_enterprise/proxy/readme.md diff --git a/enterprise/proxy/utils.py b/enterprise/litellm_enterprise/proxy/utils.py similarity index 65% rename from enterprise/proxy/utils.py rename to enterprise/litellm_enterprise/proxy/utils.py index c3396964ee6..227ea0a9ff0 100644 --- a/enterprise/proxy/utils.py +++ b/enterprise/litellm_enterprise/proxy/utils.py @@ -1,4 +1,5 @@ -from typing import Union, Optional +from typing import Optional, Union + from litellm.secret_managers.main import str_to_bool @@ -6,14 +7,19 @@ def _should_block_robots(): """ Returns True if the robots.txt file should block web crawlers - Controlled by - + Controlled by + ```yaml general_settings: block_robots: true ``` """ - from litellm.proxy.proxy_server import general_settings, premium_user, CommonProxyErrors + from litellm.proxy.proxy_server import ( + CommonProxyErrors, + general_settings, + premium_user, + ) + _block_robots: Union[bool, str] = general_settings.get("block_robots", False) block_robots: Optional[bool] = None if isinstance(_block_robots, bool): @@ -22,6 +28,8 @@ def _should_block_robots(): block_robots = str_to_bool(_block_robots) if block_robots is True: if premium_user is not True: - raise ValueError(f"Blocking web crawlers is an enterprise feature. {CommonProxyErrors.not_premium_user.value}") + raise ValueError( + f"Blocking web crawlers is an enterprise feature. {CommonProxyErrors.not_premium_user.value}" + ) return True return False diff --git a/enterprise/proxy/vector_stores/endpoints.py b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py similarity index 100% rename from enterprise/proxy/vector_stores/endpoints.py rename to enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b38e4298558..5844f333fd1 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -28,7 +28,7 @@ from typing import ( from litellm.constants import ( DEFAULT_MAX_RECURSE_DEPTH, DEFAULT_SLACK_ALERTING_THRESHOLD, - LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS + LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS, ) from litellm.types.utils import ( ModelResponse, @@ -383,15 +383,21 @@ enterprise_router = APIRouter() try: # when using litellm cli import litellm.proxy.enterprise as enterprise - from enterprise.proxy.enterprise_routes import router as enterprise_router except Exception: # when using litellm docker image try: import enterprise # type: ignore - from enterprise.proxy.enterprise_routes import router as enterprise_router except Exception: pass +################### +# Import enterprise routes +try: + from litellm_enterprise.proxy.enterprise_routes import router as enterprise_router +except ImportError: + enterprise_router = APIRouter() +################### + server_root_path = os.getenv("SERVER_ROOT_PATH", "") _license_check = LicenseCheck() premium_user: bool = _license_check.is_premium() @@ -3849,11 +3855,11 @@ async def embeddings( # noqa: PLR0915 if llm_model_list is not None and data["model"] in router_model_names: for m in llm_model_list: if m["model_name"] == data["model"]: - if (m["litellm_params"]["model"] in litellm.open_ai_embedding_models - or any( - m["litellm_params"]["model"].startswith(provider) - for provider in LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS - ) + if m["litellm_params"][ + "model" + ] in litellm.open_ai_embedding_models or any( + m["litellm_params"]["model"].startswith(provider) + for provider in LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS ): pass else: diff --git a/tests/litellm/enterprise/test_enterprise_routes.py b/tests/litellm/enterprise/test_enterprise_routes.py index 61cea9c4c20..857513391ab 100644 --- a/tests/litellm/enterprise/test_enterprise_routes.py +++ b/tests/litellm/enterprise/test_enterprise_routes.py @@ -10,7 +10,7 @@ sys.path.insert( 0, os.path.abspath("../../..") ) # Adds the parent directory to the system path -from enterprise.proxy.enterprise_routes import router +from litellm_enterprise.proxy.enterprise_routes import router @pytest.fixture @@ -21,7 +21,8 @@ def client(): def test_robots_when_blocked(client): """Test get_robots returns block instructions when _should_block_robots returns True""" with mock.patch( - "enterprise.proxy.enterprise_routes._should_block_robots", return_value=True + "litellm_enterprise.proxy.enterprise_routes._should_block_robots", + return_value=True, ): response = client.get("/robots.txt") print("got response", response) @@ -34,7 +35,8 @@ def test_robots_when_blocked(client): def test_robots_when_not_blocked(client): """Test get_robots returns 404 when _should_block_robots returns False""" with mock.patch( - "enterprise.proxy.enterprise_routes._should_block_robots", return_value=False + "litellm_enterprise.proxy.enterprise_routes._should_block_robots", + return_value=False, ): response = client.get("/robots.txt") assert response.status_code == 404