mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
[Fix] Search Tools: Handle wildcard "*" default in permissions and move auth helper
The DB migration sets search_tools DEFAULT ARRAY['*'], but the auth check and listing filter treated "*" as a literal name, causing 401s and empty list responses for all callers with the default permission. - Add wildcard handling in _can_object_call_search_tools - Add _normalize_search_tools_wildcard to convert ["*"] → None (no restriction) - Move get_allowed_search_tool_names from endpoints.py to auth_checks.py to eliminate cross-router import between search_tool_management and endpoints Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
d8704a59db
commit
9cfc5b97de
4 changed files with 200 additions and 67 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue