diff --git a/tests/test_litellm/proxy/auth/test_search_tool_access.py b/tests/test_litellm/proxy/auth/test_search_tool_access.py new file mode 100644 index 00000000000..1559672dcb8 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_search_tool_access.py @@ -0,0 +1,413 @@ +""" +Tests for search tool access control. + +Covers: +- _can_object_call_search_tools() with least-privilege semantics +- search_tool_access_check() for key and team level permissions +- get_permitted_search_tool_names() helper +- ProxyErrorTypes for search tool access denied +- Ensures vector store access check semantics are unchanged +""" + +from typing import List, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.proxy._types import ( + LiteLLM_ObjectPermissionTable, + ProxyErrorTypes, + ProxyException, + UserAPIKeyAuth, +) +from litellm.proxy.auth.auth_checks import ( + _can_object_call_search_tools, + _can_object_call_vector_stores, + get_permitted_search_tool_names, + search_tool_access_check, +) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_object_permission( + search_tools: Optional[List[str]] = None, + vector_stores: Optional[List[str]] = None, +) -> LiteLLM_ObjectPermissionTable: + """Create a minimal object permission for testing.""" + return LiteLLM_ObjectPermissionTable( + object_permission_id="test-perm-id", + search_tools=search_tools, + vector_stores=vector_stores if vector_stores is not None else [], + ) + + +# =========================================================================== +# _can_object_call_search_tools — least-privilege semantics +# =========================================================================== + + +class TestCanObjectCallSearchTools: + """should enforce least-privilege semantics for search tools.""" + + def test_should_allow_when_permissions_are_none(self): + """None object_permissions → allow (no permission record).""" + result = _can_object_call_search_tools( + object_type="key", + search_tool_name="any-tool", + object_permissions=None, + ) + assert result is True + + def test_should_allow_when_search_tools_field_is_none(self): + """search_tools=None → allow (field not configured).""" + perm = _make_object_permission(search_tools=None) + result = _can_object_call_search_tools( + object_type="key", + search_tool_name="any-tool", + object_permissions=perm, + ) + assert result is True + + def test_should_deny_when_search_tools_is_empty_list(self): + """search_tools=[] → DENY (principle of least privilege).""" + perm = _make_object_permission(search_tools=[]) + with pytest.raises(ProxyException) as exc_info: + _can_object_call_search_tools( + object_type="key", + search_tool_name="any-tool", + object_permissions=perm, + ) + assert exc_info.value.type == ProxyErrorTypes.key_search_tool_access_denied + + def test_should_allow_tool_in_allowed_list(self): + """Requesting a tool that is in the allowed list → allow.""" + perm = _make_object_permission(search_tools=["tool-a", "tool-b"]) + result = _can_object_call_search_tools( + object_type="key", + search_tool_name="tool-a", + object_permissions=perm, + ) + assert result is True + + def test_should_deny_tool_not_in_allowed_list(self): + """Requesting a tool NOT in the allowed list → deny.""" + perm = _make_object_permission(search_tools=["tool-a"]) + with pytest.raises(ProxyException) as exc_info: + _can_object_call_search_tools( + object_type="key", + search_tool_name="tool-b", + object_permissions=perm, + ) + assert exc_info.value.type == ProxyErrorTypes.key_search_tool_access_denied + + def test_should_use_team_error_type_for_team_object(self): + """Team denial uses team-specific error type.""" + perm = _make_object_permission(search_tools=[]) + with pytest.raises(ProxyException) as exc_info: + _can_object_call_search_tools( + object_type="team", + search_tool_name="any-tool", + object_permissions=perm, + ) + assert exc_info.value.type == ProxyErrorTypes.team_search_tool_access_denied + + def test_should_use_org_error_type_for_org_object(self): + """Org denial uses org-specific error type.""" + perm = _make_object_permission(search_tools=["other"]) + with pytest.raises(ProxyException) as exc_info: + _can_object_call_search_tools( + object_type="org", + search_tool_name="not-other", + object_permissions=perm, + ) + assert exc_info.value.type == ProxyErrorTypes.org_search_tool_access_denied + + def test_should_allow_all_tools_in_allowed_list(self): + """Multiple tools in allowed list all pass.""" + perm = _make_object_permission(search_tools=["a", "b", "c"]) + for name in ["a", "b", "c"]: + result = _can_object_call_search_tools( + object_type="key", + search_tool_name=name, + object_permissions=perm, + ) + assert result is True + with pytest.raises(ProxyException): + _can_object_call_search_tools( + object_type="key", + search_tool_name="d", + object_permissions=perm, + ) + + +# =========================================================================== +# search_tool_access_check — async DB lookups +# =========================================================================== + + +@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): + result = await search_tool_access_check( + search_tool_name="any-tool", + valid_token=None, + ) + assert result is True + + +@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): + result = await search_tool_access_check( + search_tool_name="any-tool", + valid_token=None, + ) + assert result is True + + +@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): + result = await search_tool_access_check( + search_tool_name="any-tool", + valid_token=token, + ) + assert result is True + + +@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): + result = await search_tool_access_check( + search_tool_name="my-tool", + valid_token=token, + ) + assert result is True + + +@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 pytest.raises(ProxyException) as exc_info: + await search_tool_access_check( + search_tool_name="any-tool", + valid_token=token, + ) + assert exc_info.value.type == ProxyErrorTypes.key_search_tool_access_denied + + +@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 pytest.raises(ProxyException) as exc_info: + await search_tool_access_check( + search_tool_name="any-tool", + valid_token=token, + ) + assert exc_info.value.type == ProxyErrorTypes.team_search_tool_access_denied + + +@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": + return key_perm + elif where["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 pytest.raises(ProxyException) as exc_info: + await search_tool_access_check( + search_tool_name="the-tool", + valid_token=token, + ) + assert exc_info.value.type == ProxyErrorTypes.team_search_tool_access_denied + + +@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": + return key_perm + elif where["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): + result = await search_tool_access_check( + search_tool_name="tool-x", + valid_token=token, + ) + assert result is True + + +# =========================================================================== +# get_permitted_search_tool_names +# =========================================================================== + + +class TestGetPermittedSearchToolNames: + def test_should_return_none_for_none_permissions(self): + assert get_permitted_search_tool_names(None) is None + + def test_should_return_none_when_search_tools_is_none(self): + perm = _make_object_permission(search_tools=None) + assert get_permitted_search_tool_names(perm) is None + + def test_should_return_empty_list_for_empty_search_tools(self): + perm = _make_object_permission(search_tools=[]) + assert get_permitted_search_tool_names(perm) == [] + + def test_should_return_specific_tools(self): + perm = _make_object_permission(search_tools=["a", "b"]) + assert get_permitted_search_tool_names(perm) == ["a", "b"] + + +# =========================================================================== +# ProxyErrorTypes classmethod +# =========================================================================== + + +class TestSearchToolErrorTypes: + def test_should_return_key_error_type(self): + assert ( + ProxyErrorTypes.get_search_tool_access_error_type_for_object("key") + == ProxyErrorTypes.key_search_tool_access_denied + ) + + def test_should_return_team_error_type(self): + assert ( + ProxyErrorTypes.get_search_tool_access_error_type_for_object("team") + == ProxyErrorTypes.team_search_tool_access_denied + ) + + def test_should_return_org_error_type(self): + assert ( + ProxyErrorTypes.get_search_tool_access_error_type_for_object("org") + == ProxyErrorTypes.org_search_tool_access_denied + ) + + +# =========================================================================== +# Regression: vector store semantics unchanged (empty = allow all) +# =========================================================================== + + +class TestVectorStoreAccessNotBroken: + """should preserve existing vector store semantics: empty list = allow ALL.""" + + def test_should_allow_all_when_vector_stores_is_empty(self): + """Vector stores: empty list = access to ALL (existing behavior).""" + 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 + + def test_should_allow_when_vector_stores_is_none(self): + perm = MagicMock() + perm.vector_stores = None + result = _can_object_call_vector_stores( + object_type="key", + vector_store_ids_to_run=["store-1"], + object_permissions=perm, + ) + assert result is True + + def test_should_deny_unlisted_vector_store(self): + perm = MagicMock() + perm.vector_stores = ["store-1"] + with pytest.raises(ProxyException) as exc_info: + _can_object_call_vector_stores( + object_type="key", + vector_store_ids_to_run=["store-99"], + object_permissions=perm, + ) + assert exc_info.value.type == ProxyErrorTypes.key_vector_store_access_denied