diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index a2c83295403..fae9ce96ce6 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -269,6 +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[]? // 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/_types.py b/litellm/proxy/_types.py index 7a376380521..928db3d20c7 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -857,6 +857,7 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase): mcp_access_groups: Optional[List[str]] = None mcp_tool_permissions: Optional[Dict[str, List[str]]] = None vector_stores: Optional[List[str]] = None + search_tools: Optional[List[str]] = None agents: Optional[List[str]] = None agent_access_groups: Optional[List[str]] = None models: Optional[List[str]] = None @@ -1876,6 +1877,7 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase): """ vector_stores: Optional[List[str]] = [] + search_tools: Optional[List[str]] = None # NULL = all access, [] = no access agents: Optional[List[str]] = [] agent_access_groups: Optional[List[str]] = [] @@ -3562,6 +3564,21 @@ class ProxyErrorTypes(str, enum.Enum): Organization does not have access to the vector store """ + key_search_tool_access_denied = "key_search_tool_access_denied" + """ + Key does not have access to the search tool + """ + + team_search_tool_access_denied = "team_search_tool_access_denied" + """ + Team does not have access to the search tool + """ + + org_search_tool_access_denied = "org_search_tool_access_denied" + """ + Organization does not have access to the search tool + """ + team_member_already_in_team = "team_member_already_in_team" """ Team member is already in team @@ -3604,6 +3621,23 @@ class ProxyErrorTypes(str, enum.Enum): elif object_type == "org": return cls.org_vector_store_access_denied + @classmethod + def get_search_tool_access_error_type_for_object( + cls, object_type: Literal["key", "team", "org"] + ) -> "ProxyErrorTypes": + """ + Get the search tool access error type for object_type + """ + if object_type == "key": + return cls.key_search_tool_access_denied + elif object_type == "team": + return cls.team_search_tool_access_denied + elif object_type == "org": + return cls.org_search_tool_access_denied + raise ValueError( + f"Unknown object_type '{object_type}' for search tool access error" + ) + DB_CONNECTION_ERROR_TYPES = ( httpx.ConnectError, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 1aa14fff574..8b67a59e34d 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3574,3 +3574,104 @@ def _can_object_call_vector_stores( ) return True + + +def _can_object_call_search_tools( + object_type: Literal["key", "team", "org"], + search_tool_name: str, + object_permissions: Optional[LiteLLM_ObjectPermissionTable], +) -> bool: + """ + Raises ProxyException if the object (key, team, org) cannot access the specific search tool. + + Key difference from vector stores: follows principle of least privilege. + - object_permissions is None → allow (no permission record) + - search_tools is None → allow (field not configured, no restriction) + - search_tools == [] → DENY ALL (empty list = no access granted) + - search_tool_name in search_tools → allow + - search_tool_name not in search_tools → deny + """ + if object_permissions is None: + return True + + if object_permissions.search_tools is None: + return True + + # Empty list = no access (principle of least privilege) + # Non-empty list = only listed tools are accessible + if search_tool_name not in object_permissions.search_tools: + raise ProxyException( + message=f"User not allowed to access search tool '{search_tool_name}'. Allowed search tools: {object_permissions.search_tools}", + type=ProxyErrorTypes.get_search_tool_access_error_type_for_object( + object_type + ), + param="search_tool", + code=status.HTTP_401_UNAUTHORIZED, + ) + + return True + + +async def search_tool_access_check( + search_tool_name: str, + valid_token: Optional[UserAPIKeyAuth], +): + """ + 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 + via the cached get_object_permission helper. + + Raises ProxyException if access is denied. + """ + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + if prisma_client is None: + verbose_proxy_logger.debug( + "Prisma client not found, skipping search tool access check" + ) + return True + + # Check key-level permissions + if valid_token is not None and valid_token.object_permission_id is not None: + 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( + object_type="key", + search_tool_name=search_tool_name, + object_permissions=key_object_permission, + ) + + # Check team-level permissions + team_object_permission_id = ( + getattr(valid_token, "team_object_permission_id", None) + if valid_token + else None + ) + if team_object_permission_id is not None: + 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( + object_type="team", + search_tool_name=search_tool_name, + object_permissions=team_object_permission, + ) + + return True diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 46be6b31e1f..26e429fdbf7 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -269,6 +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[]? // 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 8bed5b54075..b67c9fe2443 100644 --- a/litellm/proxy/search_endpoints/endpoints.py +++ b/litellm/proxy/search_endpoints/endpoints.py @@ -12,6 +12,67 @@ from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessin router = APIRouter() +async def _get_allowed_search_tool_names( + user_api_key_dict: UserAPIKeyAuth, +) -> Optional[List[str]]: + """ + Compute the intersection of key-level and team-level search tool permissions. + + Returns: + None → no restriction (all tools accessible) + list → only those tool names are accessible (may be empty = none) + """ + 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 + + key_allowed: Optional[List[str]] = None + team_allowed: Optional[List[str]] = None + + # Key-level permissions (via cached helper) + if user_api_key_dict.object_permission_id is not None: + 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 (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 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 + + # Combine: both None → None (no restriction) + # One set → use that set + # Both set → intersection + if key_allowed is None and team_allowed is None: + return None + if key_allowed is None: + return team_allowed + if team_allowed is None: + return key_allowed + # Both are set - return the intersection + return list(set(key_allowed) & set(team_allowed)) + + @router.post( "/v1/search/{search_tool_name}", dependencies=[Depends(user_api_key_auth)], @@ -163,6 +224,24 @@ async def search( data["metadata"] = {} data["metadata"]["model_group"] = search_tool_name_value + # Access control check for search tools + resolved_search_tool_name = data.get("search_tool_name") + if resolved_search_tool_name: + from litellm.proxy.auth.auth_checks import search_tool_access_check + + await search_tool_access_check( + 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) try: @@ -258,6 +337,15 @@ async def list_search_tools( search_tools_list.append(tool_info) + # Filter search tools based on user's permissions + allowed_names = await _get_allowed_search_tool_names(user_api_key_dict) + if allowed_names is not None: + search_tools_list = [ + tool + for tool in search_tools_list + if tool.get("search_tool_name") in allowed_names + ] + return {"object": "list", "data": search_tools_list} except Exception as e: from litellm._logging import verbose_proxy_logger diff --git a/schema.prisma b/schema.prisma index fde9a466a28..dcbae8a9cff 100644 --- a/schema.prisma +++ b/schema.prisma @@ -269,6 +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[]? // 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 new file mode 100644 index 00000000000..0ce5f6fd559 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_search_tool_access.py @@ -0,0 +1,386 @@ +""" +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 +- 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, + 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 +# =========================================================================== + + +_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(f"{_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.""" + 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, + ) + 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.""" + token = UserAPIKeyAuth( + object_permission_id=None, + team_object_permission_id=None, + ) + 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, + ) + 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_perm = MagicMock() + mock_perm.search_tools = ["my-tool"] + + token = UserAPIKeyAuth( + object_permission_id="key-perm-id", + team_object_permission_id=None, + ) + 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, + ) + assert result is True + + +@pytest.mark.asyncio +async def test_should_deny_key_with_empty_search_tools(): + """Key with empty search_tools → deny.""" + mock_perm = MagicMock() + mock_perm.search_tools = [] + + token = UserAPIKeyAuth( + object_permission_id="key-perm-id", + team_object_permission_id=None, + ) + 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", + 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_team_perm = MagicMock() + mock_team_perm.search_tools = [] + + token = UserAPIKeyAuth( + object_permission_id=None, + team_object_permission_id="team-perm-id", + ) + 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", + 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).""" + key_perm = MagicMock() + key_perm.search_tools = ["the-tool"] + + team_perm = MagicMock() + team_perm.search_tools = [] + + async def mock_get_perm(object_permission_id, **kwargs): + if object_permission_id == "key-perm-id": + return key_perm + elif object_permission_id == "team-perm-id": + return team_perm + return None + + token = UserAPIKeyAuth( + object_permission_id="key-perm-id", + team_object_permission_id="team-perm-id", + ) + 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", + 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.""" + key_perm = MagicMock() + key_perm.search_tools = ["tool-x"] + + team_perm = MagicMock() + team_perm.search_tools = ["tool-x", "tool-y"] + + async def mock_get_perm(object_permission_id, **kwargs): + if object_permission_id == "key-perm-id": + return key_perm + elif object_permission_id == "team-perm-id": + return team_perm + return None + + token = UserAPIKeyAuth( + object_permission_id="key-perm-id", + team_object_permission_id="team-perm-id", + ) + 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, + ) + assert result is True + + +# =========================================================================== +# 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: + """Preserves 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