mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
commit
185db1941f
7 changed files with 612 additions and 0 deletions
|
|
@ -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([])
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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([])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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([])
|
||||
|
|
|
|||
386
tests/test_litellm/proxy/auth/test_search_tool_access.py
Normal file
386
tests/test_litellm/proxy/auth/test_search_tool_access.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue