diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 8b67a59e34d..7f8def1f785 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3597,6 +3597,10 @@ def _can_object_call_search_tools( if object_permissions.search_tools is None: return True + # Wildcard "*" = all tools accessible (migration default) + if "*" in object_permissions.search_tools: + 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: @@ -3675,3 +3679,72 @@ async def search_tool_access_check( ) return True + + +def _normalize_search_tools_wildcard( + search_tools: Optional[List[str]], +) -> Optional[List[str]]: + """Treat ["*"] (the migration default) as None (no restriction).""" + if search_tools is not None and "*" in search_tools: + return None + return search_tools + + +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.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 = _normalize_search_tools_wildcard(key_perm.search_tools) + + # 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 = _normalize_search_tools_wildcard(team_perm.search_tools) + + # 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)) diff --git a/litellm/proxy/search_endpoints/endpoints.py b/litellm/proxy/search_endpoints/endpoints.py index b67c9fe2443..1f08785237c 100644 --- a/litellm/proxy/search_endpoints/endpoints.py +++ b/litellm/proxy/search_endpoints/endpoints.py @@ -6,73 +6,13 @@ from fastapi.responses import ORJSONResponse from litellm._logging import verbose_proxy_logger from litellm.proxy._types import * +from litellm.proxy.auth.auth_checks import get_allowed_search_tool_names from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing 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)], @@ -338,7 +278,7 @@ 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) + allowed_names = await get_allowed_search_tool_names(user_api_key_dict) if allowed_names is not None: search_tools_list = [ tool diff --git a/litellm/proxy/search_endpoints/search_tool_management.py b/litellm/proxy/search_endpoints/search_tool_management.py index c9d415c0b9a..2ba7bcf32fb 100644 --- a/litellm/proxy/search_endpoints/search_tool_management.py +++ b/litellm/proxy/search_endpoints/search_tool_management.py @@ -8,6 +8,7 @@ from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel from litellm._logging import verbose_proxy_logger +from litellm.proxy.auth.auth_checks import get_allowed_search_tool_names from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.search_endpoints.search_tool_registry import SearchToolRegistry from litellm.types.search import ( @@ -164,11 +165,7 @@ async def list_search_tools( ) # Filter based on caller's key/team permissions - from litellm.proxy.search_endpoints.endpoints import ( - _get_allowed_search_tool_names, - ) - - allowed_names = await _get_allowed_search_tool_names(user_api_key_dict) + allowed_names = await get_allowed_search_tool_names(user_api_key_dict) if allowed_names is not None: search_tool_configs = [ tool 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 0ce5f6fd559..877f6ac988b 100644 --- a/tests/test_litellm/proxy/auth/test_search_tool_access.py +++ b/tests/test_litellm/proxy/auth/test_search_tool_access.py @@ -22,6 +22,8 @@ from litellm.proxy._types import ( from litellm.proxy.auth.auth_checks import ( _can_object_call_search_tools, _can_object_call_vector_stores, + _normalize_search_tools_wildcard, + get_allowed_search_tool_names, search_tool_access_check, ) @@ -124,6 +126,26 @@ class TestCanObjectCallSearchTools: ) assert exc_info.value.type == ProxyErrorTypes.org_search_tool_access_denied + def test_should_allow_any_tool_when_wildcard_present(self): + """search_tools=["*"] (migration default) → allow any tool.""" + perm = _make_object_permission(search_tools=["*"]) + result = _can_object_call_search_tools( + object_type="key", + search_tool_name="any-tool", + object_permissions=perm, + ) + assert result is True + + def test_should_allow_any_tool_when_wildcard_mixed_with_names(self): + """search_tools=["*", "tool-a"] → wildcard dominates, allow any.""" + perm = _make_object_permission(search_tools=["*", "tool-a"]) + result = _can_object_call_search_tools( + object_type="key", + search_tool_name="tool-b", + object_permissions=perm, + ) + assert result is True + 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"]) @@ -384,3 +406,104 @@ class TestVectorStoreAccessNotBroken: object_permissions=perm, ) assert exc_info.value.type == ProxyErrorTypes.key_vector_store_access_denied + + +# =========================================================================== +# _normalize_search_tools_wildcard +# =========================================================================== + + +class TestNormalizeSearchToolsWildcard: + def test_none_stays_none(self): + assert _normalize_search_tools_wildcard(None) is None + + def test_empty_list_stays_empty(self): + assert _normalize_search_tools_wildcard([]) == [] + + def test_explicit_names_unchanged(self): + assert _normalize_search_tools_wildcard(["a", "b"]) == ["a", "b"] + + def test_wildcard_only_returns_none(self): + assert _normalize_search_tools_wildcard(["*"]) is None + + def test_wildcard_mixed_returns_none(self): + assert _normalize_search_tools_wildcard(["*", "tool-a"]) is None + + +# =========================================================================== +# get_allowed_search_tool_names — wildcard handling +# =========================================================================== + + +@pytest.mark.asyncio +async def test_get_allowed_should_return_none_when_key_has_wildcard(): + """Key with ["*"] → None (no restriction).""" + key_perm = MagicMock() + key_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=key_perm)): + result = await get_allowed_search_tool_names(token) + assert result is None + + +@pytest.mark.asyncio +async def test_get_allowed_should_intersect_wildcard_key_with_team_list(): + """Key ["*"] + team ["tool-a"] → ["tool-a"].""" + key_perm = MagicMock() + key_perm.search_tools = ["*"] + + team_perm = MagicMock() + team_perm.search_tools = ["tool-a"] + + 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 get_allowed_search_tool_names(token) + assert result == ["tool-a"] + + +@pytest.mark.asyncio +async def test_get_allowed_should_return_none_when_both_wildcard(): + """Key ["*"] + team ["*"] → None (no restriction).""" + key_perm = MagicMock() + key_perm.search_tools = ["*"] + + 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)): + result = await get_allowed_search_tool_names(token) + assert result is None