diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260319000000_add_search_tools_to_object_permissions/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260319000000_add_search_tools_to_object_permissions/migration.sql index 1b33edc3a8b..b01d8c106c1 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260319000000_add_search_tools_to_object_permissions/migration.sql +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260319000000_add_search_tools_to_object_permissions/migration.sql @@ -1,2 +1,2 @@ -- AlterTable -ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "search_tools" TEXT[] DEFAULT ARRAY[]::TEXT[]; +ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "search_tools" TEXT[]; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 084b1bdce6d..4fbbca07793 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -269,7 +269,7 @@ model LiteLLM_ObjectPermissionTable { mcp_access_groups String[] @default([]) mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]} vector_stores String[] @default([]) - search_tools String[] @default([]) // Search tool names allowed for this key/team/org. Empty = no access (least privilege). + search_tools String[] // Search tool names allowed for this key/team/org. NULL = all access, empty = no access. agents String[] @default([]) agent_access_groups String[] @default([]) models String[] @default([]) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 95f22bb4f3b..06c6ddd714b 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3612,11 +3612,16 @@ async def search_tool_access_check( Checks if the key/team has access to a specific search tool. Uses valid_token.object_permission_id (key level) and - valid_token.team_object_permission_id (team level) to look up permissions. + valid_token.team_object_permission_id (team level) to look up permissions + via the cached get_object_permission helper. Raises ProxyException if access is denied. """ - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) if prisma_client is None: verbose_proxy_logger.debug( @@ -3626,10 +3631,12 @@ async def search_tool_access_check( # Check key-level permissions if valid_token is not None and valid_token.object_permission_id is not None: - key_object_permission = ( - await prisma_client.db.litellm_objectpermissiontable.find_unique( - where={"object_permission_id": valid_token.object_permission_id}, - ) + key_object_permission = await get_object_permission( + object_permission_id=valid_token.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=getattr(valid_token, "parent_otel_span", None), + proxy_logging_obj=proxy_logging_obj, ) if key_object_permission is not None: _can_object_call_search_tools( @@ -3645,10 +3652,12 @@ async def search_tool_access_check( else None ) if team_object_permission_id is not None: - team_object_permission = ( - await prisma_client.db.litellm_objectpermissiontable.find_unique( - where={"object_permission_id": team_object_permission_id}, - ) + team_object_permission = await get_object_permission( + object_permission_id=team_object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=getattr(valid_token, "parent_otel_span", None), + proxy_logging_obj=proxy_logging_obj, ) if team_object_permission is not None: _can_object_call_search_tools( diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index e56de5d6feb..414ae23dad5 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -269,7 +269,7 @@ model LiteLLM_ObjectPermissionTable { mcp_access_groups String[] @default([]) mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]} vector_stores String[] @default([]) - search_tools String[] @default([]) // Search tool names allowed for this key/team/org. Empty = no access (least privilege). + search_tools String[] // Search tool names allowed for this key/team/org. NULL = all access, empty = no access. agents String[] @default([]) agent_access_groups String[] @default([]) models String[] @default([]) diff --git a/litellm/proxy/search_endpoints/endpoints.py b/litellm/proxy/search_endpoints/endpoints.py index 22f777fc929..b67c9fe2443 100644 --- a/litellm/proxy/search_endpoints/endpoints.py +++ b/litellm/proxy/search_endpoints/endpoints.py @@ -22,7 +22,12 @@ async def _get_allowed_search_tool_names( None → no restriction (all tools accessible) list → only those tool names are accessible (may be empty = none) """ - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.auth.auth_checks import get_object_permission + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) if prisma_client is None: return None @@ -30,25 +35,27 @@ async def _get_allowed_search_tool_names( key_allowed: Optional[List[str]] = None team_allowed: Optional[List[str]] = None - # Key-level permissions + # Key-level permissions (via cached helper) if user_api_key_dict.object_permission_id is not None: - key_perm = ( - await prisma_client.db.litellm_objectpermissiontable.find_unique( - where={ - "object_permission_id": user_api_key_dict.object_permission_id - }, - ) + key_perm = await get_object_permission( + object_permission_id=user_api_key_dict.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=getattr(user_api_key_dict, "parent_otel_span", None), + proxy_logging_obj=proxy_logging_obj, ) if key_perm is not None: key_allowed = key_perm.search_tools # None means no restriction - # Team-level permissions + # Team-level permissions (via cached helper) team_perm_id = getattr(user_api_key_dict, "team_object_permission_id", None) if team_perm_id is not None: - team_perm = ( - await prisma_client.db.litellm_objectpermissiontable.find_unique( - where={"object_permission_id": team_perm_id}, - ) + team_perm = await get_object_permission( + object_permission_id=team_perm_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=getattr(user_api_key_dict, "parent_otel_span", None), + proxy_logging_obj=proxy_logging_obj, ) if team_perm is not None: team_allowed = team_perm.search_tools # None means no restriction @@ -226,6 +233,14 @@ async def search( search_tool_name=resolved_search_tool_name, valid_token=user_api_key_dict, ) + else: + raise HTTPException( + status_code=400, + detail={ + "error": "search_tool_name is required. Provide it in the URL path " + "(/v1/search/{search_tool_name}) or in the request body." + }, + ) # Process request using ProxyBaseLLMRequestProcessing processor = ProxyBaseLLMRequestProcessing(data=data) diff --git a/schema.prisma b/schema.prisma index d5abf7ef382..7b782ada24b 100644 --- a/schema.prisma +++ b/schema.prisma @@ -269,7 +269,7 @@ model LiteLLM_ObjectPermissionTable { mcp_access_groups String[] @default([]) mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]} vector_stores String[] @default([]) - search_tools String[] @default([]) // Search tool names allowed for this key/team/org. Empty = no access (least privilege). + search_tools String[] // Search tool names allowed for this key/team/org. NULL = all access, empty = no access. agents String[] @default([]) agent_access_groups String[] @default([]) models String[] @default([]) diff --git a/tests/test_litellm/proxy/auth/test_search_tool_access.py b/tests/test_litellm/proxy/auth/test_search_tool_access.py index 1559672dcb8..06797760c8d 100644 --- a/tests/test_litellm/proxy/auth/test_search_tool_access.py +++ b/tests/test_litellm/proxy/auth/test_search_tool_access.py @@ -13,6 +13,7 @@ from typing import List, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest +from starlette import status from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, @@ -149,10 +150,14 @@ class TestCanObjectCallSearchTools: # =========================================================================== +_PROXY_SERVER = "litellm.proxy.proxy_server" +_GET_OBJ_PERM = "litellm.proxy.auth.auth_checks.get_object_permission" + + @pytest.mark.asyncio async def test_should_allow_when_no_prisma_client(): """No prisma client → allow.""" - with patch("litellm.proxy.proxy_server.prisma_client", None): + with patch(f"{_PROXY_SERVER}.prisma_client", None): result = await search_tool_access_check( search_tool_name="any-tool", valid_token=None, @@ -163,8 +168,9 @@ async def test_should_allow_when_no_prisma_client(): @pytest.mark.asyncio async def test_should_allow_when_no_token(): """No valid_token → allow.""" - mock_prisma = MagicMock() - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + with patch(f"{_PROXY_SERVER}.prisma_client", MagicMock()), \ + patch(f"{_PROXY_SERVER}.proxy_logging_obj", MagicMock()), \ + patch(f"{_PROXY_SERVER}.user_api_key_cache", MagicMock()): result = await search_tool_access_check( search_tool_name="any-tool", valid_token=None, @@ -175,12 +181,13 @@ async def test_should_allow_when_no_token(): @pytest.mark.asyncio async def test_should_allow_when_no_permission_ids(): """Token with no object_permission_id or team_object_permission_id → allow.""" - mock_prisma = MagicMock() token = UserAPIKeyAuth( object_permission_id=None, team_object_permission_id=None, ) - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + with patch(f"{_PROXY_SERVER}.prisma_client", MagicMock()), \ + patch(f"{_PROXY_SERVER}.proxy_logging_obj", MagicMock()), \ + patch(f"{_PROXY_SERVER}.user_api_key_cache", MagicMock()): result = await search_tool_access_check( search_tool_name="any-tool", valid_token=token, @@ -191,18 +198,17 @@ async def test_should_allow_when_no_permission_ids(): @pytest.mark.asyncio async def test_should_allow_key_with_matching_permission(): """Key with object_permission that includes the tool → allow.""" - mock_prisma = MagicMock() mock_perm = MagicMock() mock_perm.search_tools = ["my-tool"] - mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=mock_perm - ) token = UserAPIKeyAuth( object_permission_id="key-perm-id", team_object_permission_id=None, ) - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + with patch(f"{_PROXY_SERVER}.prisma_client", MagicMock()), \ + patch(f"{_PROXY_SERVER}.proxy_logging_obj", MagicMock()), \ + patch(f"{_PROXY_SERVER}.user_api_key_cache", MagicMock()), \ + patch(_GET_OBJ_PERM, AsyncMock(return_value=mock_perm)): result = await search_tool_access_check( search_tool_name="my-tool", valid_token=token, @@ -213,18 +219,17 @@ async def test_should_allow_key_with_matching_permission(): @pytest.mark.asyncio async def test_should_deny_key_with_empty_search_tools(): """Key with empty search_tools → deny.""" - mock_prisma = MagicMock() mock_perm = MagicMock() mock_perm.search_tools = [] - mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=mock_perm - ) token = UserAPIKeyAuth( object_permission_id="key-perm-id", team_object_permission_id=None, ) - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + with patch(f"{_PROXY_SERVER}.prisma_client", MagicMock()), \ + patch(f"{_PROXY_SERVER}.proxy_logging_obj", MagicMock()), \ + patch(f"{_PROXY_SERVER}.user_api_key_cache", MagicMock()), \ + patch(_GET_OBJ_PERM, AsyncMock(return_value=mock_perm)): with pytest.raises(ProxyException) as exc_info: await search_tool_access_check( search_tool_name="any-tool", @@ -236,18 +241,17 @@ async def test_should_deny_key_with_empty_search_tools(): @pytest.mark.asyncio async def test_should_deny_team_with_empty_search_tools(): """Team with empty search_tools → deny.""" - mock_prisma = MagicMock() mock_team_perm = MagicMock() mock_team_perm.search_tools = [] - mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=mock_team_perm - ) token = UserAPIKeyAuth( object_permission_id=None, team_object_permission_id="team-perm-id", ) - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + with patch(f"{_PROXY_SERVER}.prisma_client", MagicMock()), \ + patch(f"{_PROXY_SERVER}.proxy_logging_obj", MagicMock()), \ + patch(f"{_PROXY_SERVER}.user_api_key_cache", MagicMock()), \ + patch(_GET_OBJ_PERM, AsyncMock(return_value=mock_team_perm)): with pytest.raises(ProxyException) as exc_info: await search_tool_access_check( search_tool_name="any-tool", @@ -259,30 +263,27 @@ async def test_should_deny_team_with_empty_search_tools(): @pytest.mark.asyncio async def test_should_deny_when_key_allows_but_team_denies(): """Key allows, team denies → deny (team check second).""" - mock_prisma = MagicMock() - key_perm = MagicMock() key_perm.search_tools = ["the-tool"] team_perm = MagicMock() team_perm.search_tools = [] - async def mock_find(where): - if where["object_permission_id"] == "key-perm-id": + async def mock_get_perm(object_permission_id, **kwargs): + if object_permission_id == "key-perm-id": return key_perm - elif where["object_permission_id"] == "team-perm-id": + elif object_permission_id == "team-perm-id": return team_perm return None - mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock( - side_effect=mock_find - ) - token = UserAPIKeyAuth( object_permission_id="key-perm-id", team_object_permission_id="team-perm-id", ) - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + with patch(f"{_PROXY_SERVER}.prisma_client", MagicMock()), \ + patch(f"{_PROXY_SERVER}.proxy_logging_obj", MagicMock()), \ + patch(f"{_PROXY_SERVER}.user_api_key_cache", MagicMock()), \ + patch(_GET_OBJ_PERM, AsyncMock(side_effect=mock_get_perm)): with pytest.raises(ProxyException) as exc_info: await search_tool_access_check( search_tool_name="the-tool", @@ -294,30 +295,27 @@ async def test_should_deny_when_key_allows_but_team_denies(): @pytest.mark.asyncio async def test_should_allow_when_both_key_and_team_allow(): """Both key and team allow → allow.""" - mock_prisma = MagicMock() - key_perm = MagicMock() key_perm.search_tools = ["tool-x"] team_perm = MagicMock() team_perm.search_tools = ["tool-x", "tool-y"] - async def mock_find(where): - if where["object_permission_id"] == "key-perm-id": + async def mock_get_perm(object_permission_id, **kwargs): + if object_permission_id == "key-perm-id": return key_perm - elif where["object_permission_id"] == "team-perm-id": + elif object_permission_id == "team-perm-id": return team_perm return None - mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock( - side_effect=mock_find - ) - token = UserAPIKeyAuth( object_permission_id="key-perm-id", team_object_permission_id="team-perm-id", ) - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + with patch(f"{_PROXY_SERVER}.prisma_client", MagicMock()), \ + patch(f"{_PROXY_SERVER}.proxy_logging_obj", MagicMock()), \ + patch(f"{_PROXY_SERVER}.user_api_key_cache", MagicMock()), \ + patch(_GET_OBJ_PERM, AsyncMock(side_effect=mock_get_perm)): result = await search_tool_access_check( search_tool_name="tool-x", valid_token=token, @@ -378,18 +376,18 @@ class TestSearchToolErrorTypes: class TestVectorStoreAccessNotBroken: - """should preserve existing vector store semantics: empty list = allow ALL.""" + """Vector store access follows same semantics: None = all access, [] = no access.""" - def test_should_allow_all_when_vector_stores_is_empty(self): - """Vector stores: empty list = access to ALL (existing behavior).""" + def test_should_deny_all_when_vector_stores_is_empty(self): + """Vector stores: empty list = no access (matches search_tools semantics).""" perm = MagicMock() perm.vector_stores = [] - result = _can_object_call_vector_stores( - object_type="key", - vector_store_ids_to_run=["store-1"], - object_permissions=perm, - ) - assert result is True + with pytest.raises(ProxyException): + _can_object_call_vector_stores( + object_type="key", + vector_store_ids_to_run=["store-1"], + object_permissions=perm, + ) def test_should_allow_when_vector_stores_is_none(self): perm = MagicMock()