Merge pull request #24075 from BerriAI/litellm_search-tool-permissions-f5e4

[Feature] Search Tools: Add access control via object permissions
This commit is contained in:
yuneng-jiang 2026-03-21 10:23:08 -07:00 • committed by GitHub
commit 185db1941f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 612 additions and 0 deletions

View file

@ -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([])

View file

@ -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,

View file

@ -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

View file

@ -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([])

View file

@ -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

View file

@ -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([])

View file

@ -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