diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index bf99347ef6..045d2fd5f1 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2155,10 +2155,6 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase): rotation_interval: Optional[str] = None # How often to rotate (e.g., "30d", "90d") last_rotation_at: Optional[datetime] = None # When this key was last rotated key_rotation_at: Optional[datetime] = None # When this key should next be rotated - router_settings: Optional[ - Dict - ] = None # Router settings for this key (Key > Team > Global precedence) - model_config = ConfigDict(protected_namespaces=()) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 6e55cb2adf..769c250d9f 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -616,26 +616,6 @@ class ProxyBaseLLMRequestProcessing: user_api_key_dict=user_api_key_dict, data=self.data, call_type=route_type # type: ignore ) - # Apply hierarchical router_settings (Key > Team > Global) - if llm_router is not None and proxy_config is not None: - from litellm.proxy.proxy_server import prisma_client - - router_settings = await proxy_config._get_hierarchical_router_settings( - user_api_key_dict=user_api_key_dict, - prisma_client=prisma_client, - ) - - # If router_settings found (from key, team, or global), apply them - # This ensures key/team settings override global settings - if router_settings is not None and router_settings: - # Get model_list from current router - model_list = llm_router.get_model_list() - if model_list is not None: - # Create user_config with model_list and router_settings - # This creates a per-request router with the hierarchical settings - user_config = {"model_list": model_list, **router_settings} - self.data["user_config"] = user_config - if "messages" in self.data and self.data["messages"]: logging_obj.update_messages(self.data["messages"]) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f12d4d6ab4..2343fe8c35 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -3380,105 +3380,6 @@ class ProxyConfig: decrypted_variables[k] = decrypted_value return decrypted_variables - async def _get_hierarchical_router_settings( - self, - user_api_key_dict: Optional["UserAPIKeyAuth"], - prisma_client: Optional[PrismaClient], - ) -> Optional[dict]: - """ - Get router_settings in priority order: Key > Team > Global - - Returns: - dict: Combined router_settings, or None if no settings found - """ - if prisma_client is None: - return None - - import json - - import yaml - - # 1. Try key-level router_settings - if user_api_key_dict is not None: - # Check if router_settings is available on the key object - key_router_settings_value = getattr( - user_api_key_dict, "router_settings", None - ) - if key_router_settings_value is not None: - key_router_settings = None - if isinstance(key_router_settings_value, str): - try: - key_router_settings = yaml.safe_load(key_router_settings_value) - except (yaml.YAMLError, json.JSONDecodeError): - try: - key_router_settings = json.loads(key_router_settings_value) - except json.JSONDecodeError: - pass - elif isinstance(key_router_settings_value, dict): - key_router_settings = key_router_settings_value - - # If key has router_settings (non-empty dict), use it - if ( - key_router_settings is not None - and isinstance(key_router_settings, dict) - and key_router_settings - ): - return key_router_settings - - # 2. Try team-level router_settings - if user_api_key_dict is not None and user_api_key_dict.team_id is not None: - try: - team_obj = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": user_api_key_dict.team_id} - ) - if team_obj is not None: - team_router_settings_value = getattr( - team_obj, "router_settings", None - ) - if team_router_settings_value is not None: - team_router_settings = None - if isinstance(team_router_settings_value, str): - try: - team_router_settings = yaml.safe_load( - team_router_settings_value - ) - except (yaml.YAMLError, json.JSONDecodeError): - try: - team_router_settings = json.loads( - team_router_settings_value - ) - except json.JSONDecodeError: - pass - elif isinstance(team_router_settings_value, dict): - team_router_settings = team_router_settings_value - - # If team has router_settings (non-empty dict), use it - if ( - team_router_settings is not None - and isinstance(team_router_settings, dict) - and team_router_settings - ): - return team_router_settings - except Exception: - # If team lookup fails, continue to global settings - pass - - # 3. Try global router_settings - try: - db_router_settings = await prisma_client.db.litellm_config.find_first( - where={"param_name": "router_settings"} - ) - if ( - db_router_settings is not None - and isinstance(db_router_settings.param_value, dict) - and db_router_settings.param_value - ): - return db_router_settings.param_value - except Exception: - pass - - return None - async def _add_router_settings_from_db_config( self, config_data: dict, diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 6edcdab15c..69cf8240c6 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -77,84 +77,6 @@ class TestProxyBaseLLMRequestProcessing: pytest.fail("litellm_call_id is not a valid UUID") assert data_passed["litellm_call_id"] == returned_data["litellm_call_id"] - @pytest.mark.asyncio - async def test_should_apply_hierarchical_router_settings_to_user_config( - self, monkeypatch - ): - processing_obj = ProxyBaseLLMRequestProcessing(data={}) - mock_request = MagicMock(spec=Request) - mock_request.headers = {} - - async def mock_add_litellm_data_to_request(*args, **kwargs): - return {} - - async def mock_common_processing_pre_call_logic( - user_api_key_dict, data, call_type - ): - data_copy = copy.deepcopy(data) - return data_copy - - mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) - mock_proxy_logging_obj.pre_call_hook = AsyncMock( - side_effect=mock_common_processing_pre_call_logic - ) - monkeypatch.setattr( - litellm.proxy.common_request_processing, - "add_litellm_data_to_request", - mock_add_litellm_data_to_request, - ) - - mock_general_settings = {} - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_proxy_config = MagicMock(spec=ProxyConfig) - - mock_router_settings = { - "routing_strategy": "least-busy", - "timeout": 30.0, - "num_retries": 3, - } - mock_proxy_config._get_hierarchical_router_settings = AsyncMock( - return_value=mock_router_settings - ) - - mock_model_list = [ - {"model_name": "gpt-3.5-turbo", "litellm_params": {"model": "gpt-3.5-turbo"}}, - {"model_name": "gpt-4", "litellm_params": {"model": "gpt-4"}}, - ] - mock_llm_router = MagicMock() - mock_llm_router.get_model_list = MagicMock(return_value=mock_model_list) - - mock_prisma_client = MagicMock() - monkeypatch.setattr( - "litellm.proxy.proxy_server.prisma_client", - mock_prisma_client, - ) - - route_type = "acompletion" - - returned_data, logging_obj = await processing_obj.common_processing_pre_call_logic( - request=mock_request, - general_settings=mock_general_settings, - user_api_key_dict=mock_user_api_key_dict, - proxy_logging_obj=mock_proxy_logging_obj, - proxy_config=mock_proxy_config, - route_type=route_type, - llm_router=mock_llm_router, - ) - - mock_proxy_config._get_hierarchical_router_settings.assert_called_once_with( - user_api_key_dict=mock_user_api_key_dict, - prisma_client=mock_prisma_client, - ) - mock_llm_router.get_model_list.assert_called_once() - - assert "user_config" in returned_data - user_config = returned_data["user_config"] - assert user_config["model_list"] == mock_model_list - assert user_config["routing_strategy"] == "least-busy" - assert user_config["timeout"] == 30.0 - assert user_config["num_retries"] == 3 - @pytest.mark.asyncio async def test_stream_timeout_header_processing(self): """ diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 18d3257c9c..acd9909039 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3203,2196 +3203,3 @@ def test_deep_merge_dicts_skips_none_and_empty_lists(monkeypatch): assert result["general_settings"]["nested"]["key1"] == "updated_value1" assert result["general_settings"]["nested"]["key2"] == "value2" assert result["general_settings"]["nested"]["key3"] == "value3" - - -@pytest.mark.asyncio -async def test_get_hierarchical_router_settings(): - """ - Test _get_hierarchical_router_settings method's priority order: Key > Team > Global - """ - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.proxy_server import ProxyConfig - - proxy_config = ProxyConfig() - - # Test Case 1: Returns None when prisma_client is None - result = await proxy_config._get_hierarchical_router_settings( - user_api_key_dict=None, - prisma_client=None, - ) - assert result is None - - # Test Case 2: Returns key-level router_settings when available (as dict) - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_user_api_key_dict.router_settings = {"routing_strategy": "key-level", "timeout": 10} - mock_user_api_key_dict.team_id = None - - mock_prisma_client = MagicMock() - - result = await proxy_config._get_hierarchical_router_settings( - user_api_key_dict=mock_user_api_key_dict, - prisma_client=mock_prisma_client, - ) - assert result == {"routing_strategy": "key-level", "timeout": 10} - - # Test Case 3: Returns key-level router_settings when available (as YAML string) - mock_user_api_key_dict.router_settings = "routing_strategy: key-yaml\ntimeout: 20" - result = await proxy_config._get_hierarchical_router_settings( - user_api_key_dict=mock_user_api_key_dict, - prisma_client=mock_prisma_client, - ) - assert result == {"routing_strategy": "key-yaml", "timeout": 20} - - # Test Case 4: Falls back to team-level router_settings when key-level is not available - mock_user_api_key_dict.router_settings = None - mock_user_api_key_dict.team_id = "team-123" - - mock_team_obj = MagicMock() - mock_team_obj.router_settings = {"routing_strategy": "team-level", "timeout": 30} - - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_team_obj - ) - - result = await proxy_config._get_hierarchical_router_settings( - user_api_key_dict=mock_user_api_key_dict, - prisma_client=mock_prisma_client, - ) - assert result == {"routing_strategy": "team-level", "timeout": 30} - mock_prisma_client.db.litellm_teamtable.find_unique.assert_called_once_with( - where={"team_id": "team-123"} - ) - - # Test Case 5: Falls back to global router_settings when neither key nor team settings are available - mock_user_api_key_dict.router_settings = None - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) - - mock_db_config = MagicMock() - mock_db_config.param_value = {"routing_strategy": "global-level", "timeout": 40} - - mock_prisma_client.db.litellm_config.find_first = AsyncMock( - return_value=mock_db_config - ) - - result = await proxy_config._get_hierarchical_router_settings( - user_api_key_dict=mock_user_api_key_dict, - prisma_client=mock_prisma_client, - ) - assert result == {"routing_strategy": "global-level", "timeout": 40} - mock_prisma_client.db.litellm_config.find_first.assert_called_once_with( - where={"param_name": "router_settings"} - ) - - # Test Case 6: Returns None when no settings are found - mock_user_api_key_dict.router_settings = None - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) - mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) - - result = await proxy_config._get_hierarchical_router_settings( - user_api_key_dict=mock_user_api_key_dict, - prisma_client=mock_prisma_client, - ) - assert result is None - - -@pytest.mark.asyncio -async def test_model_info_v2_pagination_basic(monkeypatch): - """ - Test basic pagination functionality for /v2/model/info endpoint. - Tests multiple pages with different page sizes. - """ - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth - - # Create 75 mock models for testing pagination - mock_models = [ - { - "model_name": f"model-{i}", - "litellm_params": {"model": f"gpt-{i}"}, - "model_info": {"id": f"model-{i}"}, - } - for i in range(1, 76) # 75 models total - ] - - # Mock llm_router - mock_router = MagicMock() - mock_router.model_list = mock_models - - # Mock prisma_client - mock_prisma_client = MagicMock() - - # Mock proxy_config.get_config - mock_get_config = AsyncMock(return_value={}) - - # Mock user authentication - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_user_api_key_dict.user_id = "test-user" - mock_user_api_key_dict.api_key = "test-key" - mock_user_api_key_dict.team_models = [] - mock_user_api_key_dict.models = [] - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - - # Override auth dependency - original_overrides = app.dependency_overrides.copy() - app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict - - client = TestClient(app) - try: - # Test page 1 with size 25 (should return models 1-25) - response = client.get("/v2/model/info", params={"page": 1, "size": 25}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 75 - assert data["current_page"] == 1 - assert data["size"] == 25 - assert data["total_pages"] == 3 # ceil(75/25) = 3 - assert len(data["data"]) == 25 - assert data["data"][0]["model_name"] == "model-1" - assert data["data"][24]["model_name"] == "model-25" - - # Test page 2 with size 25 (should return models 26-50) - response = client.get("/v2/model/info", params={"page": 2, "size": 25}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 75 - assert data["current_page"] == 2 - assert data["size"] == 25 - assert data["total_pages"] == 3 - assert len(data["data"]) == 25 - assert data["data"][0]["model_name"] == "model-26" - assert data["data"][24]["model_name"] == "model-50" - - # Test page 3 with size 25 (should return models 51-75) - response = client.get("/v2/model/info", params={"page": 3, "size": 25}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 75 - assert data["current_page"] == 3 - assert data["size"] == 25 - assert data["total_pages"] == 3 - assert len(data["data"]) == 25 - assert data["data"][0]["model_name"] == "model-51" - assert data["data"][24]["model_name"] == "model-75" - - # Test different page size (size 10) - response = client.get("/v2/model/info", params={"page": 1, "size": 10}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 75 - assert data["current_page"] == 1 - assert data["size"] == 10 - assert data["total_pages"] == 8 # ceil(75/10) = 8 - assert len(data["data"]) == 10 - - finally: - app.dependency_overrides = original_overrides - - -@pytest.mark.asyncio -async def test_model_info_v2_pagination_edge_cases(monkeypatch): - """ - Test edge cases for pagination in /v2/model/info endpoint. - Tests empty results, last page with partial results, and boundary conditions. - """ - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth - - # Mock prisma_client - mock_prisma_client = MagicMock() - - # Mock user authentication - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_user_api_key_dict.user_id = "test-user" - mock_user_api_key_dict.api_key = "test-key" - mock_user_api_key_dict.team_models = [] - mock_user_api_key_dict.models = [] - - # Mock proxy_config.get_config - mock_get_config = AsyncMock(return_value={}) - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - - # Override auth dependency - original_overrides = app.dependency_overrides.copy() - app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict - - client = TestClient(app) - try: - # Test Case 1: Empty model list (no models configured) - mock_router_empty = MagicMock() - mock_router_empty.model_list = [] - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router_empty) - - response = client.get("/v2/model/info", params={"page": 1, "size": 25}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 0 - assert data["current_page"] == 1 - assert data["size"] == 25 - assert data["total_pages"] == 0 - assert len(data["data"]) == 0 - - # Test Case 2: Last page with partial results (23 models, page size 10) - mock_models_partial = [ - { - "model_name": f"model-{i}", - "litellm_params": {"model": f"gpt-{i}"}, - "model_info": {"id": f"model-{i}"}, - } - for i in range(1, 24) # 23 models total - ] - mock_router_partial = MagicMock() - mock_router_partial.model_list = mock_models_partial - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router_partial) - - # Page 1 should have 10 models - response = client.get("/v2/model/info", params={"page": 1, "size": 10}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 23 - assert data["current_page"] == 1 - assert data["total_pages"] == 3 # ceil(23/10) = 3 - assert len(data["data"]) == 10 - - # Page 2 should have 10 models - response = client.get("/v2/model/info", params={"page": 2, "size": 10}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 23 - assert data["current_page"] == 2 - assert data["total_pages"] == 3 - assert len(data["data"]) == 10 - - # Page 3 (last page) should have only 3 models - response = client.get("/v2/model/info", params={"page": 3, "size": 10}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 23 - assert data["current_page"] == 3 - assert data["total_pages"] == 3 - assert len(data["data"]) == 3 - assert data["data"][0]["model_name"] == "model-21" - assert data["data"][2]["model_name"] == "model-23" - - # Test Case 3: Page beyond available pages (should return empty data) - response = client.get("/v2/model/info", params={"page": 4, "size": 10}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 23 - assert data["current_page"] == 4 - assert data["total_pages"] == 3 - assert len(data["data"]) == 0 # No data for page beyond total_pages - - # Test Case 4: Single model with page size 1 - mock_models_single = [ - { - "model_name": "single-model", - "litellm_params": {"model": "gpt-4"}, - "model_info": {"id": "single-model"}, - } - ] - mock_router_single = MagicMock() - mock_router_single.model_list = mock_models_single - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router_single) - - response = client.get("/v2/model/info", params={"page": 1, "size": 1}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 1 - assert data["current_page"] == 1 - assert data["total_pages"] == 1 - assert len(data["data"]) == 1 - assert data["data"][0]["model_name"] == "single-model" - - finally: - app.dependency_overrides = original_overrides - - -@pytest.mark.asyncio -async def test_model_info_v2_search_config_models(monkeypatch): - """ - Test search parameter for config models (models from config.yaml). - Config models don't have db_model=True in model_info. - """ - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth - - # Create mock config models (no db_model flag or db_model=False) - mock_config_models = [ - { - "model_name": "gpt-4-turbo", - "litellm_params": {"model": "gpt-4-turbo"}, - "model_info": {"id": "gpt-4-turbo"}, # No db_model flag = config model - }, - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": {"id": "gpt-3.5-turbo", "db_model": False}, # Explicitly config model - }, - { - "model_name": "claude-3-opus", - "litellm_params": {"model": "claude-3-opus"}, - "model_info": {"id": "claude-3-opus"}, # No db_model flag = config model - }, - { - "model_name": "gemini-pro", - "litellm_params": {"model": "gemini-pro"}, - "model_info": {"id": "gemini-pro"}, # No db_model flag = config model - }, - ] - - # Mock llm_router - mock_router = MagicMock() - mock_router.model_list = mock_config_models - - # Mock prisma_client - mock_prisma_client = MagicMock() - - # Mock proxy_config.get_config - mock_get_config = AsyncMock(return_value={}) - - # Mock user authentication - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_user_api_key_dict.user_id = "test-user" - mock_user_api_key_dict.api_key = "test-key" - mock_user_api_key_dict.team_models = [] - mock_user_api_key_dict.models = [] - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - - # Override auth dependency - original_overrides = app.dependency_overrides.copy() - app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict - - client = TestClient(app) - try: - # Test search for "gpt" - should return gpt-4-turbo and gpt-3.5-turbo - response = client.get("/v2/model/info", params={"search": "gpt"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 2 # Only config models matching search - assert len(data["data"]) == 2 - model_names = [m["model_name"] for m in data["data"]] - assert "gpt-4-turbo" in model_names - assert "gpt-3.5-turbo" in model_names - assert "claude-3-opus" not in model_names - assert "gemini-pro" not in model_names - - # Test search for "claude" - should return claude-3-opus - response = client.get("/v2/model/info", params={"search": "claude"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 1 - assert len(data["data"]) == 1 - assert data["data"][0]["model_name"] == "claude-3-opus" - - # Test case-insensitive search - response = client.get("/v2/model/info", params={"search": "GPT"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 2 - assert len(data["data"]) == 2 - - # Test partial match - response = client.get("/v2/model/info", params={"search": "turbo"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 2 - assert len(data["data"]) == 2 - model_names = [m["model_name"] for m in data["data"]] - assert "gpt-4-turbo" in model_names - assert "gpt-3.5-turbo" in model_names - - # Test search with no matches - response = client.get("/v2/model/info", params={"search": "nonexistent"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 0 - assert len(data["data"]) == 0 - - finally: - app.dependency_overrides = original_overrides - - -@pytest.mark.asyncio -async def test_model_info_v2_search_db_models(monkeypatch): - """ - Test search parameter for db models (models from database). - DB models have db_model=True and id in model_info. - """ - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth - - # Create mock db models (db_model=True with id) - mock_db_models_in_router = [ - { - "model_name": "db-gpt-4", - "litellm_params": {"model": "gpt-4"}, - "model_info": {"id": "db-model-1", "db_model": True}, # DB model - }, - { - "model_name": "db-claude-3", - "litellm_params": {"model": "claude-3"}, - "model_info": {"id": "db-model-2", "db_model": True}, # DB model - }, - ] - - # Mock llm_router - mock_router = MagicMock() - mock_router.model_list = mock_db_models_in_router - - # Mock prisma_client with database query methods - mock_db_models_from_db = [ - MagicMock( - model_id="db-model-3", - model_name="db-gemini-pro", - litellm_params='{"model": "gemini-pro"}', - model_info='{"id": "db-model-3", "db_model": true}', - ), - MagicMock( - model_id="db-model-4", - model_name="db-gpt-3.5", - litellm_params='{"model": "gpt-3.5-turbo"}', - model_info='{"id": "db-model-4", "db_model": true}', - ), - ] - - # Mock the database count and find_many methods dynamically based on search - async def mock_db_count_func(*args, **kwargs): - where_condition = kwargs.get("where", {}) - search_term = where_condition.get("model_name", {}).get("contains", "") - excluded_ids = where_condition.get("model_id", {}).get("not", {}).get("in", []) - - # Count models matching search term but not in excluded_ids - count = 0 - for model in mock_db_models_from_db: - if search_term.lower() in model.model_name.lower(): - if model.model_id not in excluded_ids: - count += 1 - return count - - async def mock_db_find_many_func(*args, **kwargs): - where_condition = kwargs.get("where", {}) - search_term = where_condition.get("model_name", {}).get("contains", "") - excluded_ids = where_condition.get("model_id", {}).get("not", {}).get("in", []) - take = kwargs.get("take", 10) - - # Return models matching search term but not in excluded_ids - result = [] - for model in mock_db_models_from_db: - if search_term.lower() in model.model_name.lower(): - if model.model_id not in excluded_ids: - result.append(model) - if len(result) >= take: - break - return result - - mock_db_count = AsyncMock(side_effect=mock_db_count_func) - mock_db_find_many = AsyncMock(side_effect=mock_db_find_many_func) - - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_proxymodeltable.count = mock_db_count - mock_prisma_client.db.litellm_proxymodeltable.find_many = mock_db_find_many - - # Mock proxy_config.decrypt_model_list_from_db to return router-format models - def mock_decrypt_models(db_models_list): - result = [] - for db_model in db_models_list: - result.append( - { - "model_name": db_model.model_name, - "litellm_params": {"model": db_model.model_name.replace("db-", "")}, - "model_info": {"id": db_model.model_id, "db_model": True}, - } - ) - return result - - # Mock proxy_config.get_config - mock_get_config = AsyncMock(return_value={}) - - # Mock user authentication - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_user_api_key_dict.user_id = "test-user" - mock_user_api_key_dict.api_key = "test-key" - mock_user_api_key_dict.team_models = [] - mock_user_api_key_dict.models = [] - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - monkeypatch.setattr(proxy_config, "decrypt_model_list_from_db", mock_decrypt_models) - - # Override auth dependency - original_overrides = app.dependency_overrides.copy() - app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict - - client = TestClient(app) - try: - # Test search for "gpt" - should return db-gpt-4 from router and db-gpt-3.5 from db - response = client.get("/v2/model/info", params={"search": "gpt"}) - assert response.status_code == 200 - data = response.json() - # Should have db-gpt-4 from router + db-gpt-3.5 from db = 2 total - assert data["total_count"] == 2 - assert len(data["data"]) == 2 - model_names = [m["model_name"] for m in data["data"]] - assert "db-gpt-4" in model_names - assert "db-gpt-3.5" in model_names - - # Verify database was queried - mock_db_count.assert_called() - # Verify the where condition excludes models already in router - call_args = mock_db_count.call_args - assert call_args is not None - where_condition = call_args[1]["where"] - assert "model_name" in where_condition - assert where_condition["model_name"]["contains"] == "gpt" - assert where_condition["model_name"]["mode"] == "insensitive" - - # Test search for "claude" - should return db-claude-3 from router only - response = client.get("/v2/model/info", params={"search": "claude"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 1 - assert len(data["data"]) == 1 - assert data["data"][0]["model_name"] == "db-claude-3" - - # Test search for "gemini" - should return db-gemini-pro from db only - response = client.get("/v2/model/info", params={"search": "gemini"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 1 - assert len(data["data"]) == 1 - assert data["data"][0]["model_name"] == "db-gemini-pro" - - # Test case-insensitive search - response = client.get("/v2/model/info", params={"search": "GPT"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 2 - - finally: - app.dependency_overrides = original_overrides - - -@pytest.mark.asyncio -async def test_model_info_v2_filter_by_model_id(monkeypatch): - """ - Test modelId parameter for filtering by specific model ID. - Tests that modelId searches in router config first, then database. - """ - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth - - # Create mock config models - mock_config_models = [ - { - "model_name": "gpt-4-turbo", - "litellm_params": {"model": "gpt-4-turbo"}, - "model_info": {"id": "config-model-1"}, - }, - { - "model_name": "claude-3-opus", - "litellm_params": {"model": "claude-3-opus"}, - "model_info": {"id": "config-model-2"}, - }, - ] - - # Mock llm_router with get_model_info method - mock_router = MagicMock() - mock_router.model_list = mock_config_models - mock_router.get_model_info = MagicMock( - side_effect=lambda id: next( - (m for m in mock_config_models if m["model_info"]["id"] == id), None - ) - ) - - # Mock prisma_client for database queries - mock_prisma_client = MagicMock() - mock_db_table = MagicMock() - mock_prisma_client.db.litellm_proxymodeltable = mock_db_table - - # Mock database model - mock_db_model = MagicMock() - mock_db_model.model_id = "db-model-1" - mock_db_model.model_name = "db-gpt-3.5" - mock_db_model.litellm_params = '{"model": "gpt-3.5-turbo"}' - mock_db_model.model_info = '{"id": "db-model-1", "db_model": true}' - - # Mock find_unique to return db model when searching for db-model-1 - async def mock_find_unique(where): - if where.get("model_id") == "db-model-1": - return mock_db_model - return None - - mock_db_table.find_unique = AsyncMock(side_effect=mock_find_unique) - - # Mock proxy_config.decrypt_model_list_from_db - def mock_decrypt_models(db_models_list): - if db_models_list: - return [ - { - "model_name": db_models_list[0].model_name, - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": {"id": db_models_list[0].model_id, "db_model": True}, - } - ] - return [] - - # Mock proxy_config.get_config - mock_get_config = AsyncMock(return_value={}) - - # Mock user authentication - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_user_api_key_dict.user_id = "test-user" - mock_user_api_key_dict.api_key = "test-key" - mock_user_api_key_dict.team_models = [] - mock_user_api_key_dict.models = [] - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - monkeypatch.setattr(proxy_config, "decrypt_model_list_from_db", mock_decrypt_models) - - # Override auth dependency - original_overrides = app.dependency_overrides.copy() - app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict - - client = TestClient(app) - try: - # Test Case 1: Filter by modelId that exists in config - response = client.get("/v2/model/info", params={"modelId": "config-model-1"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 1 - assert len(data["data"]) == 1 - assert data["data"][0]["model_info"]["id"] == "config-model-1" - assert data["data"][0]["model_name"] == "gpt-4-turbo" - # Verify router.get_model_info was called - mock_router.get_model_info.assert_called_with(id="config-model-1") - - # Test Case 2: Filter by modelId that exists in database (not in config) - response = client.get("/v2/model/info", params={"modelId": "db-model-1"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 1 - assert len(data["data"]) == 1 - assert data["data"][0]["model_info"]["id"] == "db-model-1" - assert data["data"][0]["model_name"] == "db-gpt-3.5" - # Verify database was queried - mock_db_table.find_unique.assert_called() - - # Test Case 3: Filter by modelId that doesn't exist - mock_db_table.find_unique = AsyncMock(return_value=None) - response = client.get("/v2/model/info", params={"modelId": "non-existent-model"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 0 - assert len(data["data"]) == 0 - - # Test Case 4: Filter by modelId with search parameter (should filter further) - response = client.get( - "/v2/model/info", params={"modelId": "config-model-1", "search": "claude"} - ) - assert response.status_code == 200 - data = response.json() - # config-model-1 is gpt-4-turbo, doesn't match "claude", so should return empty - assert data["total_count"] == 0 - assert len(data["data"]) == 0 - - finally: - app.dependency_overrides = original_overrides - - -@pytest.mark.asyncio -async def test_model_info_v2_filter_by_team_id(monkeypatch): - """ - Test teamId parameter for filtering models by team ID. - Tests that teamId filters models based on direct_access or access_via_team_ids. - """ - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy._types import LiteLLM_TeamTable, UserAPIKeyAuth - from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth - - # Create mock models with different access configurations - mock_models = [ - { - "model_name": "model-direct-access", - "litellm_params": {"model": "gpt-4"}, - "model_info": { - "id": "model-1", - "direct_access": True, # Should be included - }, - }, - { - "model_name": "model-team-access", - "litellm_params": {"model": "claude-3"}, - "model_info": { - "id": "model-2", - "direct_access": False, - "access_via_team_ids": ["team-123"], # Should be included - }, - }, - { - "model_name": "model-no-access", - "litellm_params": {"model": "gemini-pro"}, - "model_info": { - "id": "model-3", - "direct_access": False, - "access_via_team_ids": ["team-456"], # Should NOT be included - }, - }, - { - "model_name": "model-multiple-teams", - "litellm_params": {"model": "gpt-3.5"}, - "model_info": { - "id": "model-4", - "direct_access": False, - "access_via_team_ids": ["team-789", "team-123"], # Should be included - }, - }, - ] - - # Mock llm_router - mock_router = MagicMock() - mock_router.model_list = mock_models - - # Mock get_model_list to return models based on model_name filter - def mock_get_model_list(model_name=None, team_id=None): - if model_name: - return [m for m in mock_models if m["model_name"] == model_name] - return mock_models - - mock_router.get_model_list = MagicMock(side_effect=mock_get_model_list) - - # Mock team database object - team has access to specific models - mock_team_db_object = MagicMock() - mock_team_db_object.model_dump.return_value = { - "team_id": "team-123", - "models": ["model-direct-access", "model-team-access", "model-multiple-teams"], # Specific models - } - - # Mock prisma_client - mock_prisma_client = MagicMock() - mock_team_table = MagicMock() - mock_prisma_client.db.litellm_teamtable = mock_team_table - mock_team_table.find_unique = AsyncMock(return_value=mock_team_db_object) - - # Mock LiteLLM_TeamTable - team has access to specific models - mock_team_object = LiteLLM_TeamTable( - team_id="team-123", - models=["model-direct-access", "model-team-access", "model-multiple-teams"], - ) - - # Mock proxy_config.get_config - mock_get_config = AsyncMock(return_value={}) - - # Mock user authentication - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_user_api_key_dict.user_id = "test-user" - mock_user_api_key_dict.api_key = "test-key" - mock_user_api_key_dict.team_models = [] - mock_user_api_key_dict.models = [] - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - # Mock LiteLLM_TeamTable instantiation - monkeypatch.setattr( - "litellm.proxy.proxy_server.LiteLLM_TeamTable", - lambda **kwargs: mock_team_object, - ) - - # Override auth dependency - original_overrides = app.dependency_overrides.copy() - app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict - - client = TestClient(app) - try: - # Test Case 1: Filter by teamId - should return models with direct_access=True or team-123 in access_via_team_ids - response = client.get("/v2/model/info", params={"teamId": "team-123"}) - assert response.status_code == 200 - data = response.json() - # Should include: model-1 (direct_access), model-2 (team-123 in access_via_team_ids), model-4 (team-123 in access_via_team_ids) - # Should NOT include: model-3 (team-456 only) - assert data["total_count"] == 3 - assert len(data["data"]) == 3 - model_ids = [m["model_info"]["id"] for m in data["data"]] - assert "model-1" in model_ids # direct_access - assert "model-2" in model_ids # team-123 in access_via_team_ids - assert "model-4" in model_ids # team-123 in access_via_team_ids - assert "model-3" not in model_ids # Should be excluded - - # Test Case 2: Filter by teamId that doesn't exist - should return empty list - mock_team_table.find_unique = AsyncMock(return_value=None) - response = client.get("/v2/model/info", params={"teamId": "non-existent-team"}) - assert response.status_code == 200 - data = response.json() - assert data["total_count"] == 0 - assert len(data["data"]) == 0 - - # Test Case 3: Filter by different teamId - should only return models with that team in access_via_team_ids - mock_team_db_object_456 = MagicMock() - mock_team_db_object_456.model_dump.return_value = { - "team_id": "team-456", - "models": ["model-no-access"], # Team has access to model-no-access - } - mock_team_table.find_unique = AsyncMock(return_value=mock_team_db_object_456) - mock_team_object_456 = LiteLLM_TeamTable( - team_id="team-456", - models=["model-no-access"], - ) - monkeypatch.setattr( - "litellm.proxy.proxy_server.LiteLLM_TeamTable", - lambda **kwargs: mock_team_object_456, - ) - - response = client.get("/v2/model/info", params={"teamId": "team-456"}) - assert response.status_code == 200 - data = response.json() - # Should include: model-1 (direct_access), model-3 (team-456 in access_via_team_ids) - # Should NOT include: model-2 (team-123 only), model-4 (team-789 and team-123, but not team-456) - assert data["total_count"] >= 2 - model_ids = [m["model_info"]["id"] for m in data["data"]] - assert "model-1" in model_ids # direct_access - assert "model-3" in model_ids # team-456 in access_via_team_ids - - finally: - app.dependency_overrides = original_overrides - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "sort_by,sort_order,expected_order", - [ - # Test model_name sorting - ("model_name", "asc", ["a-model", "b-model", "z-model"]), - ("model_name", "desc", ["z-model", "b-model", "a-model"]), - # Test created_at sorting - ("created_at", "asc", ["old-model", "mid-model", "new-model"]), - ("created_at", "desc", ["new-model", "mid-model", "old-model"]), - # Test updated_at sorting - ("updated_at", "asc", ["old-updated", "mid-updated", "new-updated"]), - ("updated_at", "desc", ["new-updated", "mid-updated", "old-updated"]), - # Test costs sorting - ("costs", "asc", ["low-cost", "mid-cost", "high-cost"]), - ("costs", "desc", ["high-cost", "mid-cost", "low-cost"]), - # Test status sorting (False/config models come before True/db models in asc) - ("status", "asc", ["config-model-1", "config-model-2", "db-model"]), - ("status", "desc", ["db-model", "config-model-1", "config-model-2"]), - ], -) -async def test_model_info_v2_sorting(monkeypatch, sort_by, sort_order, expected_order): - """ - Test sorting functionality for /v2/model/info endpoint. - Tests all sortBy fields (model_name, created_at, updated_at, costs, status) - with both asc and desc sort orders. - """ - from datetime import datetime, timedelta - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth - - # Create base time for date comparisons - base_time = datetime(2024, 1, 1, 12, 0, 0) - - # Create mock models with different values for each sort field - mock_models = [] - - if sort_by == "model_name": - # Models with different names - mock_models = [ - { - "model_name": "z-model", - "litellm_params": {"model": "z-model"}, - "model_info": {"id": "z-model"}, - }, - { - "model_name": "a-model", - "litellm_params": {"model": "a-model"}, - "model_info": {"id": "a-model"}, - }, - { - "model_name": "b-model", - "litellm_params": {"model": "b-model"}, - "model_info": {"id": "b-model"}, - }, - ] - elif sort_by == "created_at": - # Models with different created_at timestamps - mock_models = [ - { - "model_name": "new-model", - "litellm_params": {"model": "new-model"}, - "model_info": { - "id": "new-model", - "created_at": (base_time + timedelta(days=3)).isoformat(), - }, - }, - { - "model_name": "old-model", - "litellm_params": {"model": "old-model"}, - "model_info": { - "id": "old-model", - "created_at": (base_time - timedelta(days=3)).isoformat(), - }, - }, - { - "model_name": "mid-model", - "litellm_params": {"model": "mid-model"}, - "model_info": { - "id": "mid-model", - "created_at": base_time.isoformat(), - }, - }, - ] - elif sort_by == "updated_at": - # Models with different updated_at timestamps - mock_models = [ - { - "model_name": "new-updated", - "litellm_params": {"model": "new-updated"}, - "model_info": { - "id": "new-updated", - "updated_at": (base_time + timedelta(days=3)).isoformat(), - }, - }, - { - "model_name": "old-updated", - "litellm_params": {"model": "old-updated"}, - "model_info": { - "id": "old-updated", - "updated_at": (base_time - timedelta(days=3)).isoformat(), - }, - }, - { - "model_name": "mid-updated", - "litellm_params": {"model": "mid-updated"}, - "model_info": { - "id": "mid-updated", - "updated_at": base_time.isoformat(), - }, - }, - ] - elif sort_by == "costs": - # Models with different costs (input_cost + output_cost) - mock_models = [ - { - "model_name": "high-cost", - "litellm_params": {"model": "high-cost"}, - "model_info": { - "id": "high-cost", - "input_cost_per_token": 0.00005, - "output_cost_per_token": 0.00015, - }, - }, - { - "model_name": "low-cost", - "litellm_params": {"model": "low-cost"}, - "model_info": { - "id": "low-cost", - "input_cost_per_token": 0.00001, - "output_cost_per_token": 0.00003, - }, - }, - { - "model_name": "mid-cost", - "litellm_params": {"model": "mid-cost"}, - "model_info": { - "id": "mid-cost", - "input_cost_per_token": 0.00003, - "output_cost_per_token": 0.00007, - }, - }, - ] - elif sort_by == "status": - # Models with different db_model status (False = config, True = db) - mock_models = [ - { - "model_name": "db-model", - "litellm_params": {"model": "db-model"}, - "model_info": {"id": "db-model", "db_model": True}, - }, - { - "model_name": "config-model-1", - "litellm_params": {"model": "config-model-1"}, - "model_info": {"id": "config-model-1", "db_model": False}, - }, - { - "model_name": "config-model-2", - "litellm_params": {"model": "config-model-2"}, - "model_info": {"id": "config-model-2", "db_model": False}, - }, - ] - - # Mock llm_router - mock_router = MagicMock() - mock_router.model_list = mock_models - - # Mock prisma_client - mock_prisma_client = MagicMock() - - # Mock proxy_config.get_config - mock_get_config = AsyncMock(return_value={}) - - # Mock user authentication - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_user_api_key_dict.user_id = "test-user" - mock_user_api_key_dict.api_key = "test-key" - mock_user_api_key_dict.team_models = [] - mock_user_api_key_dict.models = [] - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - - # Override auth dependency - original_overrides = app.dependency_overrides.copy() - app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict - - client = TestClient(app) - try: - # Test sorting with specified sortBy and sortOrder - response = client.get( - "/v2/model/info", params={"sortBy": sort_by, "sortOrder": sort_order} - ) - assert response.status_code == 200 - data = response.json() - assert len(data["data"]) == len(expected_order) - - # Verify models are in expected order - actual_order = [m["model_name"] for m in data["data"]] - assert actual_order == expected_order, ( - f"Sorting failed for sortBy={sort_by}, sortOrder={sort_order}. " - f"Expected: {expected_order}, Got: {actual_order}" - ) - - finally: - app.dependency_overrides = original_overrides - - -@pytest.mark.asyncio -async def test_model_info_v2_sorting_invalid_sort_order(monkeypatch): - """ - Test that invalid sortOrder values return a 400 error. - """ - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth - - # Create mock models - mock_models = [ - { - "model_name": "test-model", - "litellm_params": {"model": "test-model"}, - "model_info": {"id": "test-model"}, - } - ] - - # Mock llm_router - mock_router = MagicMock() - mock_router.model_list = mock_models - - # Mock prisma_client - mock_prisma_client = MagicMock() - - # Mock proxy_config.get_config - mock_get_config = AsyncMock(return_value={}) - - # Mock user authentication - mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) - mock_user_api_key_dict.user_id = "test-user" - mock_user_api_key_dict.api_key = "test-key" - mock_user_api_key_dict.team_models = [] - mock_user_api_key_dict.models = [] - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - - # Override auth dependency - original_overrides = app.dependency_overrides.copy() - app.dependency_overrides[user_api_key_auth] = lambda: mock_user_api_key_dict - - client = TestClient(app) - try: - # Test invalid sortOrder - response = client.get( - "/v2/model/info", params={"sortBy": "model_name", "sortOrder": "invalid"} - ) - assert response.status_code == 400 - data = response.json() - assert "Invalid sortOrder" in data["detail"] - - finally: - app.dependency_overrides = original_overrides - - -@pytest.mark.asyncio -async def test_apply_search_filter_to_models(monkeypatch): - """ - Test the _apply_search_filter_to_models helper function. - Tests search filtering logic for config models, db models, and database queries. - """ - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy.proxy_server import _apply_search_filter_to_models, proxy_config - - # Create mock models with mix of config and db models - mock_models = [ - { - "model_name": "gpt-4-turbo", - "model_info": {"id": "gpt-4-turbo"}, # Config model - }, - { - "model_name": "db-gpt-3.5", - "model_info": {"id": "db-model-1", "db_model": True}, # DB model in router - }, - { - "model_name": "claude-3-opus", - "model_info": {"id": "claude-3-opus"}, # Config model - }, - ] - - # Mock prisma_client - mock_prisma_client = MagicMock() - mock_db_table = MagicMock() - mock_prisma_client.db.litellm_proxymodeltable = mock_db_table - - # Mock database models - mock_db_model_1 = MagicMock( - model_id="db-model-2", - model_name="db-gemini-pro", - litellm_params='{"model": "gemini-pro"}', - model_info='{"id": "db-model-2", "db_model": true}', - ) - - # Mock proxy_config.decrypt_model_list_from_db - mock_decrypt = MagicMock(return_value=[{"model_name": "db-gemini-pro", "model_info": {"id": "db-model-2", "db_model": True}}]) - - monkeypatch.setattr(proxy_config, "decrypt_model_list_from_db", mock_decrypt) - - # Test Case 1: No search term - should return all models unchanged - result_models, total_count = await _apply_search_filter_to_models( - all_models=mock_models.copy(), - search="", - page=1, - size=50, - prisma_client=mock_prisma_client, - proxy_config=proxy_config, - ) - assert result_models == mock_models - assert total_count is None - - # Test Case 2: Search for "gpt" - should filter router models and query DB - mock_db_table.count = AsyncMock(return_value=0) - mock_db_table.find_many = AsyncMock(return_value=[]) - - result_models, total_count = await _apply_search_filter_to_models( - all_models=mock_models.copy(), - search="gpt", - page=1, - size=50, - prisma_client=mock_prisma_client, - proxy_config=proxy_config, - ) - assert len(result_models) == 2 - model_names = [m["model_name"] for m in result_models] - assert "gpt-4-turbo" in model_names - assert "db-gpt-3.5" in model_names - assert "claude-3-opus" not in model_names - assert total_count == 2 # Only router models match - - # Test Case 3: Search with DB models matching - mock_db_table.count = AsyncMock(return_value=1) - mock_db_table.find_many = AsyncMock(return_value=[mock_db_model_1]) - - result_models, total_count = await _apply_search_filter_to_models( - all_models=mock_models.copy(), - search="gemini", - page=1, - size=50, - prisma_client=mock_prisma_client, - proxy_config=proxy_config, - ) - assert total_count == 1 # Router models (0) + DB models (1) - assert len(result_models) == 1 - assert result_models[0]["model_name"] == "db-gemini-pro" - - # Test Case 4: Case-insensitive search - # Reset mocks - no DB models should match "GPT" - mock_db_table.count = AsyncMock(return_value=0) - mock_db_table.find_many = AsyncMock(return_value=[]) - - result_models, total_count = await _apply_search_filter_to_models( - all_models=mock_models.copy(), - search="GPT", - page=1, - size=50, - prisma_client=mock_prisma_client, - proxy_config=proxy_config, - ) - assert len(result_models) == 2 - model_names = [m["model_name"] for m in result_models] - assert "gpt-4-turbo" in model_names - assert "db-gpt-3.5" in model_names - - # Test Case 5: Database query error - should fallback to router models count - mock_db_table.count = AsyncMock(side_effect=Exception("DB error")) - mock_db_table.find_many = AsyncMock(return_value=[]) - - result_models, total_count = await _apply_search_filter_to_models( - all_models=mock_models.copy(), - search="gpt", - page=1, - size=50, - prisma_client=mock_prisma_client, - proxy_config=proxy_config, - ) - # Should still return filtered router models - assert len(result_models) == 2 - assert total_count == 2 # Fallback to router models count - - -def test_paginate_models_response(): - """ - Test the _paginate_models_response helper function. - Tests pagination calculation and response formatting. - """ - from litellm.proxy.proxy_server import _paginate_models_response - - # Create mock models - mock_models = [ - {"model_name": f"model-{i}", "model_info": {"id": f"model-{i}"}} - for i in range(25) - ] - - # Test Case 1: Basic pagination - first page - result = _paginate_models_response( - all_models=mock_models, - page=1, - size=10, - total_count=None, - search=None, - ) - assert result["total_count"] == 25 - assert result["current_page"] == 1 - assert result["total_pages"] == 3 # ceil(25/10) = 3 - assert result["size"] == 10 - assert len(result["data"]) == 10 - assert result["data"][0]["model_name"] == "model-0" - - # Test Case 2: Second page - result = _paginate_models_response( - all_models=mock_models, - page=2, - size=10, - total_count=None, - search=None, - ) - assert result["current_page"] == 2 - assert len(result["data"]) == 10 - assert result["data"][0]["model_name"] == "model-10" - - # Test Case 3: Last page (partial) - result = _paginate_models_response( - all_models=mock_models, - page=3, - size=10, - total_count=None, - search=None, - ) - assert result["current_page"] == 3 - assert len(result["data"]) == 5 # Only 5 models left - assert result["data"][0]["model_name"] == "model-20" - - # Test Case 4: With explicit total_count (for search scenarios) - result = _paginate_models_response( - all_models=mock_models[:10], # Only 10 models in list - page=1, - size=10, - total_count=50, # But total_count says 50 - search="test", - ) - assert result["total_count"] == 50 - assert result["total_pages"] == 5 # ceil(50/10) = 5 - assert len(result["data"]) == 10 - - # Test Case 5: Empty models list - result = _paginate_models_response( - all_models=[], - page=1, - size=10, - total_count=0, - search=None, - ) - assert result["total_count"] == 0 - assert result["total_pages"] == 0 - assert len(result["data"]) == 0 - - # Test Case 6: Page beyond available data - result = _paginate_models_response( - all_models=mock_models[:10], - page=5, - size=10, - total_count=10, - search=None, - ) - assert result["current_page"] == 5 - assert len(result["data"]) == 0 # No data for page 5 - - -def test_enrich_model_info_with_litellm_data(): - """ - Test the _enrich_model_info_with_litellm_data helper function. - Tests model info enrichment, debug mode, and sensitive info removal. - """ - from unittest.mock import MagicMock, patch - - from litellm.proxy.proxy_server import _enrich_model_info_with_litellm_data - - # Test Case 1: Basic model enrichment without debug - model = { - "model_name": "test-model", - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": {"id": "test-model"}, - "api_key": "sk-secret-key", # Should be removed - } - - with patch("litellm.proxy.proxy_server.get_litellm_model_info") as mock_get_info, patch( - "litellm.proxy.proxy_server.remove_sensitive_info_from_deployment" - ) as mock_remove_sensitive: - mock_get_info.return_value = { - "input_cost_per_token": 0.001, - "output_cost_per_token": 0.002, - "max_tokens": 4096, - } - mock_remove_sensitive.return_value = { - "model_name": "test-model", - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": { - "id": "test-model", - "input_cost_per_token": 0.001, - "output_cost_per_token": 0.002, - "max_tokens": 4096, - }, - } - - result = _enrich_model_info_with_litellm_data(model=model, debug=False) - - # Verify get_litellm_model_info was called - mock_get_info.assert_called_once_with(model=model) - # Verify remove_sensitive_info_from_deployment was called - mock_remove_sensitive.assert_called_once() - # Verify result doesn't have api_key - assert "api_key" not in result - # Verify model_info was enriched - assert "input_cost_per_token" in result["model_info"] - - # Test Case 2: Model enrichment with debug mode - model_with_debug = { - "model_name": "test-model-debug", - "litellm_params": {"model": "gpt-4"}, - "model_info": {}, - } - - mock_router = MagicMock() - mock_client = MagicMock() - mock_router._get_client.return_value = mock_client - - with patch("litellm.proxy.proxy_server.get_litellm_model_info") as mock_get_info, patch( - "litellm.proxy.proxy_server.remove_sensitive_info_from_deployment" - ) as mock_remove_sensitive: - mock_get_info.return_value = {} - mock_remove_sensitive.return_value = { - "model_name": "test-model-debug", - "litellm_params": {"model": "gpt-4"}, - "model_info": {}, - "openai_client": str(mock_client), - } - - result = _enrich_model_info_with_litellm_data( - model=model_with_debug, debug=True, llm_router=mock_router - ) - - # Verify debug info was added - mock_remove_sensitive.assert_called_once() - call_args = mock_remove_sensitive.call_args[0][0] - assert "openai_client" in call_args - # Verify router._get_client was called for debug - mock_router._get_client.assert_called_once() - - # Test Case 3: Model with fallback to litellm.get_model_info - model_fallback = { - "model_name": "test-model-fallback", - "litellm_params": {"model": "claude-3-opus"}, - "model_info": {}, - } - - with patch("litellm.proxy.proxy_server.get_litellm_model_info") as mock_get_info, patch( - "litellm.get_model_info" - ) as mock_litellm_info, patch( - "litellm.proxy.proxy_server.remove_sensitive_info_from_deployment" - ) as mock_remove_sensitive: - # First call returns empty, triggering fallback - mock_get_info.return_value = {} - mock_litellm_info.return_value = { - "input_cost_per_token": 0.015, - "output_cost_per_token": 0.075, - "max_tokens": 200000, - } - mock_remove_sensitive.return_value = { - "model_name": "test-model-fallback", - "litellm_params": {"model": "claude-3-opus"}, - "model_info": { - "input_cost_per_token": 0.015, - "output_cost_per_token": 0.075, - "max_tokens": 200000, - }, - } - - result = _enrich_model_info_with_litellm_data(model=model_fallback, debug=False) - - # Verify fallback was attempted - mock_litellm_info.assert_called_once_with(model="claude-3-opus") - # Verify model_info was enriched with fallback data - call_args = mock_remove_sensitive.call_args[0][0] - assert call_args["model_info"]["input_cost_per_token"] == 0.015 - - # Test Case 4: Model with split model name fallback - model_split = { - "model_name": "test-model-split", - "litellm_params": {"model": "azure/gpt-4"}, - "model_info": {}, - } - - with patch("litellm.proxy.proxy_server.get_litellm_model_info") as mock_get_info, patch( - "litellm.get_model_info" - ) as mock_litellm_info, patch( - "litellm.proxy.proxy_server.remove_sensitive_info_from_deployment" - ) as mock_remove_sensitive: - # Both first and second pass return empty, triggering third pass - mock_get_info.return_value = {} - # Second pass (no split) - mock_litellm_info.side_effect = [ - {}, # First call returns empty - {"max_tokens": 8192}, # Third pass with split succeeds - ] - mock_remove_sensitive.return_value = { - "model_name": "test-model-split", - "litellm_params": {"model": "azure/gpt-4"}, - "model_info": {"max_tokens": 8192}, - } - - result = _enrich_model_info_with_litellm_data(model=model_split, debug=False) - - # Verify third pass was attempted with split model name - assert mock_litellm_info.call_count == 2 - # Check that second call used split model name - second_call = mock_litellm_info.call_args_list[1] - assert second_call[1]["model"] == "gpt-4" - assert second_call[1]["custom_llm_provider"] == "azure" - - # Test Case 5: Model with existing model_info (should preserve existing keys) - model_existing = { - "model_name": "test-model-existing", - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": {"id": "existing-id", "custom_key": "custom_value"}, - } - - with patch("litellm.proxy.proxy_server.get_litellm_model_info") as mock_get_info, patch( - "litellm.proxy.proxy_server.remove_sensitive_info_from_deployment" - ) as mock_remove_sensitive: - mock_get_info.return_value = { - "input_cost_per_token": 0.001, - "id": "new-id", # Should not override existing "id" - } - mock_remove_sensitive.return_value = { - "model_name": "test-model-existing", - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": { - "id": "existing-id", # Existing key preserved - "custom_key": "custom_value", # Existing key preserved - "input_cost_per_token": 0.001, # New key added - }, - } - - result = _enrich_model_info_with_litellm_data(model=model_existing, debug=False) - - # Verify existing keys are preserved - call_args = mock_remove_sensitive.call_args[0][0] - assert call_args["model_info"]["id"] == "existing-id" - assert call_args["model_info"]["custom_key"] == "custom_value" - assert call_args["model_info"]["input_cost_per_token"] == 0.001 - - -@pytest.mark.asyncio -async def test_model_list_scope_parameter_validation(monkeypatch): - """Test that invalid scope parameter raises HTTPException""" - from fastapi import HTTPException - - from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.proxy_server import model_list - - mock_user_api_key_dict = UserAPIKeyAuth( - user_id="test-user", - user_role=LitellmUserRoles.INTERNAL_USER, - api_key="test-key", - ) - - # Test invalid scope parameter - with pytest.raises(HTTPException) as exc_info: - await model_list( - user_api_key_dict=mock_user_api_key_dict, - scope="invalid_scope", - ) - - assert exc_info.value.status_code == 400 - assert "Invalid scope parameter" in exc_info.value.detail - assert "Only 'expand' is currently supported" in exc_info.value.detail - - -@pytest.mark.asyncio -async def test_model_list_scope_expand_proxy_admin(monkeypatch): - """Test that proxy admin with scope=expand returns all proxy models""" - from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.proxy_server import model_list - - # Mock user API key dict for proxy admin - mock_user_api_key_dict = UserAPIKeyAuth( - user_id="proxy-admin-user", - user_role=LitellmUserRoles.PROXY_ADMIN, - api_key="test-key", - ) - - # Mock llm_router with proxy models - mock_router = MagicMock() - mock_router.get_model_names.return_value = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"] - mock_router.get_model_access_groups.return_value = {} - - # Mock prisma_client - mock_prisma_client = MagicMock() - - # Mock user_api_key_cache - mock_user_api_key_cache = MagicMock() - - # Mock proxy_logging_obj - mock_proxy_logging_obj = MagicMock() - - # Mock get_complete_model_list - mock_all_models = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"] - - # Mock create_model_info_response - def mock_create_model_info_response(model_id, provider, include_metadata=False, fallback_type=None, llm_router=None): - return {"id": model_id, "object": "model"} - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache) - monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj) - monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr( - "litellm.proxy.auth.model_checks.get_complete_model_list", - lambda **kwargs: mock_all_models, - ) - monkeypatch.setattr( - "litellm.proxy.utils.create_model_info_response", - mock_create_model_info_response, - ) - - # Call model_list with scope=expand - result = await model_list( - user_api_key_dict=mock_user_api_key_dict, - scope="expand", - ) - - # Verify result contains all proxy models - assert result["object"] == "list" - assert len(result["data"]) == 3 - assert all(model["id"] in mock_all_models for model in result["data"]) - - # Verify router methods were called - mock_router.get_model_names.assert_called_once() - mock_router.get_model_access_groups.assert_called_once() - - -@pytest.mark.asyncio -async def test_model_list_scope_expand_org_admin(monkeypatch): - """Test that org admin with scope=expand returns all proxy models""" - from litellm.proxy._types import ( - LiteLLM_UserTable, - LitellmUserRoles, - UserAPIKeyAuth, - ) - from litellm.proxy.proxy_server import model_list - - # Mock user API key dict for org admin - mock_user_api_key_dict = UserAPIKeyAuth( - user_id="org-admin-user", - user_role=LitellmUserRoles.INTERNAL_USER, # Not proxy admin, but org admin - api_key="test-key", - ) - - # Mock user object with org admin membership - from datetime import datetime - - from litellm.proxy._types import LiteLLM_OrganizationMembershipTable - - mock_user_obj = LiteLLM_UserTable( - user_id="org-admin-user", - user_email="org-admin@example.com", - organization_memberships=[ - LiteLLM_OrganizationMembershipTable( - user_id="org-admin-user", - organization_id="org-123", - user_role=LitellmUserRoles.ORG_ADMIN.value, - spend=0.0, - created_at=datetime.now(), - updated_at=datetime.now(), - ) - ], - teams=[], - ) - - # Mock llm_router with proxy models - mock_router = MagicMock() - mock_router.get_model_names.return_value = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"] - mock_router.get_model_access_groups.return_value = {} - - # Mock prisma_client - mock_prisma_client = MagicMock() - - # Mock user_api_key_cache - mock_user_api_key_cache = MagicMock() - - # Mock proxy_logging_obj - mock_proxy_logging_obj = MagicMock() - - # Mock get_user_object to return user with org admin role - async def mock_get_user_object(*args, **kwargs): - return mock_user_obj - - # Mock get_complete_model_list - mock_all_models = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"] - - # Mock create_model_info_response - def mock_create_model_info_response(model_id, provider, include_metadata=False, fallback_type=None, llm_router=None): - return {"id": model_id, "object": "model"} - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache) - monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj) - monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr( - "litellm.proxy.auth.auth_checks.get_user_object", - mock_get_user_object, - ) - monkeypatch.setattr( - "litellm.proxy.auth.model_checks.get_complete_model_list", - lambda **kwargs: mock_all_models, - ) - monkeypatch.setattr( - "litellm.proxy.utils.create_model_info_response", - mock_create_model_info_response, - ) - - # Call model_list with scope=expand - result = await model_list( - user_api_key_dict=mock_user_api_key_dict, - scope="expand", - ) - - # Verify result contains all proxy models - assert result["object"] == "list" - assert len(result["data"]) == 3 - assert all(model["id"] in mock_all_models for model in result["data"]) - - # Verify router methods were called - mock_router.get_model_names.assert_called_once() - mock_router.get_model_access_groups.assert_called_once() - - -@pytest.mark.asyncio -async def test_model_list_scope_expand_team_admin(monkeypatch): - """Test that team admin with scope=expand returns all proxy models""" - from litellm.proxy._types import ( - LiteLLM_TeamTable, - LiteLLM_UserTable, - LitellmUserRoles, - UserAPIKeyAuth, - ) - from litellm.proxy.proxy_server import model_list - - # Mock user API key dict for team admin - mock_user_api_key_dict = UserAPIKeyAuth( - user_id="team-admin-user", - user_role=LitellmUserRoles.INTERNAL_USER, # Not proxy admin, but team admin - api_key="test-key", - ) - - # Mock team with user as admin - use dict structure that matches Prisma return - mock_team = MagicMock() - mock_team.model_dump.return_value = { - "team_id": "team-123", - "members_with_roles": [ - {"user_id": "team-admin-user", "role": "admin"} - ], - } - # Create team object from the dict (validator will convert members_with_roles to Member objects) - mock_team_obj = LiteLLM_TeamTable(**mock_team.model_dump()) - - # Mock user object with team membership - mock_user_obj = LiteLLM_UserTable( - user_id="team-admin-user", - user_email="team-admin@example.com", - organization_memberships=[], - teams=["team-123"], - ) - - # Mock llm_router with proxy models - mock_router = MagicMock() - mock_router.get_model_names.return_value = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"] - mock_router.get_model_access_groups.return_value = {} - - # Mock prisma_client - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock( - return_value=[mock_team] - ) - - # Mock user_api_key_cache - mock_user_api_key_cache = MagicMock() - - # Mock proxy_logging_obj - mock_proxy_logging_obj = MagicMock() - - # Mock get_user_object to return user with team membership - async def mock_get_user_object(*args, **kwargs): - return mock_user_obj - - # Mock get_complete_model_list - mock_all_models = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"] - - # Mock create_model_info_response - def mock_create_model_info_response(model_id, provider, include_metadata=False, fallback_type=None, llm_router=None): - return {"id": model_id, "object": "model"} - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache) - monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj) - monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr( - "litellm.proxy.auth.auth_checks.get_user_object", - mock_get_user_object, - ) - monkeypatch.setattr( - "litellm.proxy.auth.model_checks.get_complete_model_list", - lambda **kwargs: mock_all_models, - ) - monkeypatch.setattr( - "litellm.proxy.utils.create_model_info_response", - mock_create_model_info_response, - ) - - # Call model_list with scope=expand - result = await model_list( - user_api_key_dict=mock_user_api_key_dict, - scope="expand", - ) - - # Verify result contains all proxy models - assert result["object"] == "list" - assert len(result["data"]) == 3 - assert all(model["id"] in mock_all_models for model in result["data"]) - - # Verify router methods were called - mock_router.get_model_names.assert_called_once() - mock_router.get_model_access_groups.assert_called_once() - - -@pytest.mark.asyncio -async def test_model_list_scope_expand_normal_user(monkeypatch): - """Test that normal internal user with scope=expand returns only their models (not expanded)""" - from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.proxy_server import model_list - - # Mock user API key dict for normal internal user - mock_user_api_key_dict = UserAPIKeyAuth( - user_id="normal-user", - user_role=LitellmUserRoles.INTERNAL_USER, - api_key="test-key", - models=["gpt-3.5-turbo"], # User only has access to this model - ) - - # Mock user object without admin privileges - mock_user_obj = LiteLLM_UserTable( - user_id="normal-user", - user_email="normal@example.com", - organization_memberships=[], # No org admin - teams=[], # No teams - ) - - # Mock llm_router - mock_router = MagicMock() - mock_router.get_model_names.return_value = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"] - - # Mock prisma_client - mock_prisma_client = MagicMock() - - # Mock user_api_key_cache - mock_user_api_key_cache = MagicMock() - - # Mock proxy_logging_obj - mock_proxy_logging_obj = MagicMock() - - # Mock get_user_object to return user without admin privileges - async def mock_get_user_object(*args, **kwargs): - return mock_user_obj - - # Mock get_available_models_for_user to return only user's models - async def mock_get_available_models_for_user(*args, **kwargs): - return ["gpt-3.5-turbo"] # Only user's accessible models - - # Mock create_model_info_response - def mock_create_model_info_response(model_id, provider, include_metadata=False, fallback_type=None, llm_router=None): - return {"id": model_id, "object": "model"} - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache) - monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj) - monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr( - "litellm.proxy.auth.auth_checks.get_user_object", - mock_get_user_object, - ) - monkeypatch.setattr( - "litellm.proxy.utils.get_available_models_for_user", - mock_get_available_models_for_user, - ) - monkeypatch.setattr( - "litellm.proxy.utils.create_model_info_response", - mock_create_model_info_response, - ) - - # Call model_list with scope=expand - result = await model_list( - user_api_key_dict=mock_user_api_key_dict, - scope="expand", - ) - - # Verify result contains only user's models (not all proxy models) - assert result["object"] == "list" - assert len(result["data"]) == 1 - assert result["data"][0]["id"] == "gpt-3.5-turbo" - - # Verify router methods were NOT called (normal path, not expanded) - mock_router.get_model_names.assert_not_called() - mock_router.get_model_access_groups.assert_not_called() - - -@pytest.mark.asyncio -async def test_model_list_no_scope_parameter(monkeypatch): - """Test that model_list without scope parameter uses normal behavior""" - from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.proxy_server import model_list - - # Mock user API key dict - mock_user_api_key_dict = UserAPIKeyAuth( - user_id="test-user", - user_role=LitellmUserRoles.INTERNAL_USER, - api_key="test-key", - models=["gpt-3.5-turbo"], - ) - - # Mock llm_router - mock_router = MagicMock() - - # Mock prisma_client - mock_prisma_client = MagicMock() - - # Mock user_api_key_cache - mock_user_api_key_cache = MagicMock() - - # Mock proxy_logging_obj - mock_proxy_logging_obj = MagicMock() - - # Mock get_available_models_for_user - async def mock_get_available_models_for_user(*args, **kwargs): - return ["gpt-3.5-turbo"] - - # Mock create_model_info_response - def mock_create_model_info_response(model_id, provider, include_metadata=False, fallback_type=None, llm_router=None): - return {"id": model_id, "object": "model"} - - # Apply monkeypatches - monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache) - monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj) - monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) - monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None) - monkeypatch.setattr( - "litellm.proxy.utils.get_available_models_for_user", - mock_get_available_models_for_user, - ) - monkeypatch.setattr( - "litellm.proxy.utils.create_model_info_response", - mock_create_model_info_response, - ) - - # Call model_list without scope parameter - result = await model_list( - user_api_key_dict=mock_user_api_key_dict, - scope=None, - ) - - # Verify result uses normal behavior - assert result["object"] == "list" - assert len(result["data"]) == 1 - assert result["data"][0]["id"] == "gpt-3.5-turbo" - - # Verify router methods were NOT called (normal path) - mock_router.get_model_names.assert_not_called() - mock_router.get_model_access_groups.assert_not_called() - - -@pytest.mark.asyncio -async def test_update_general_settings_store_prompts_in_spend_logs(monkeypatch): - """ - Test that _update_general_settings correctly normalizes store_prompts_in_spend_logs - values (handles bool, string, None, and other types). - """ - from unittest.mock import patch - - from litellm.proxy.proxy_server import ProxyConfig - - proxy_config = ProxyConfig() - - # Test Case 1: None value - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": None} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is None - - # Test Case 2: bool True - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": True} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is True - - # Test Case 3: bool False - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": False} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is False - - # Test Case 4: string "true" (lowercase) - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": "true"} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is True - - # Test Case 5: string "True" (capitalized) - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": "True"} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is True - - # Test Case 6: string "TRUE" (uppercase) - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": "TRUE"} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is True - - # Test Case 7: string "false" (lowercase) - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": "false"} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is False - - # Test Case 8: string "False" (capitalized) - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": "False"} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is False - - # Test Case 9: string "FALSE" (uppercase) - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": "FALSE"} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is False - - # Test Case 10: other string value (should be False) - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": "invalid"} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is False - - # Test Case 11: integer 1 (should be True) - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": 1} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is True - - # Test Case 12: integer 0 (should be False) - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs): - await proxy_config._update_general_settings( - {"store_prompts_in_spend_logs": 0} - ) - assert mock_gs.get("store_prompts_in_spend_logs") is False - - -@pytest.mark.asyncio -async def test_update_general_settings_maximum_spend_logs_retention_period(monkeypatch): - """ - Test that _update_general_settings correctly handles maximum_spend_logs_retention_period - and reschedules cleanup job when value changes. - """ - from unittest.mock import AsyncMock, patch - - from litellm.proxy.proxy_server import ProxyConfig - - proxy_config = ProxyConfig() - - # Test Case 1: Setting a new value should reschedule cleanup job - mock_reschedule = AsyncMock() - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs), patch.object( - proxy_config, "_reschedule_spend_log_cleanup_job", mock_reschedule - ): - await proxy_config._update_general_settings( - {"maximum_spend_logs_retention_period": "7d"} - ) - assert mock_gs.get("maximum_spend_logs_retention_period") == "7d" - mock_reschedule.assert_called_once() - - # Test Case 2: Setting the same value should not reschedule - mock_reschedule.reset_mock() - mock_gs = {"maximum_spend_logs_retention_period": "7d"} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs), patch.object( - proxy_config, "_reschedule_spend_log_cleanup_job", mock_reschedule - ): - await proxy_config._update_general_settings( - {"maximum_spend_logs_retention_period": "7d"} - ) - assert mock_gs.get("maximum_spend_logs_retention_period") == "7d" - mock_reschedule.assert_not_called() - - # Test Case 3: Changing value should reschedule - mock_reschedule.reset_mock() - mock_gs = {"maximum_spend_logs_retention_period": "7d"} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs), patch.object( - proxy_config, "_reschedule_spend_log_cleanup_job", mock_reschedule - ): - await proxy_config._update_general_settings( - {"maximum_spend_logs_retention_period": "30d"} - ) - assert mock_gs.get("maximum_spend_logs_retention_period") == "30d" - mock_reschedule.assert_called_once() - - # Test Case 4: Setting to None should reschedule - mock_reschedule.reset_mock() - mock_gs = {"maximum_spend_logs_retention_period": "7d"} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs), patch.object( - proxy_config, "_reschedule_spend_log_cleanup_job", mock_reschedule - ): - await proxy_config._update_general_settings( - {"maximum_spend_logs_retention_period": None} - ) - assert mock_gs.get("maximum_spend_logs_retention_period") is None - mock_reschedule.assert_called_once() - - # Test Case 5: Changing from None to a value should reschedule - mock_reschedule.reset_mock() - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs), patch.object( - proxy_config, "_reschedule_spend_log_cleanup_job", mock_reschedule - ): - await proxy_config._update_general_settings( - {"maximum_spend_logs_retention_period": "24h"} - ) - assert mock_gs.get("maximum_spend_logs_retention_period") == "24h" - mock_reschedule.assert_called_once() - - # Test Case 6: Setting None when already None should not reschedule - mock_reschedule.reset_mock() - mock_gs = {} - with patch("litellm.proxy.proxy_server.general_settings", mock_gs), patch.object( - proxy_config, "_reschedule_spend_log_cleanup_job", mock_reschedule - ): - await proxy_config._update_general_settings( - {"maximum_spend_logs_retention_period": None} - ) - assert mock_gs.get("maximum_spend_logs_retention_period") is None - mock_reschedule.assert_not_called()