mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: address greptile feedback on search_tools RBAC
- Migration: remove DEFAULT ARRAY[]::TEXT[] so existing rows get NULL (NULL = all access) instead of [] (no access) - Schema: remove @default([]) for search_tools in all 3 schema.prisma files - Use cached get_object_permission instead of raw DB queries in search_tool_access_check and _get_allowed_search_tool_names - Reject requests with missing search_tool_name (400) instead of silently bypassing access control - Fix vector store test: [] now correctly denies access (not allows) - Update search_tool_access_check tests to mock get_object_permission Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
bca07dca61
commit
d1a0b94919
7 changed files with 98 additions and 76 deletions
|
|
@ -1,2 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "search_tools" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "search_tools" TEXT[];
|
||||
|
|
|
|||
|
|
@ -269,7 +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[] @default([]) // Search tool names allowed for this key/team/org. Empty = no access (least privilege).
|
||||
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([])
|
||||
|
|
|
|||
|
|
@ -3612,11 +3612,16 @@ async def search_tool_access_check(
|
|||
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.
|
||||
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
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -3626,10 +3631,12 @@ async def search_tool_access_check(
|
|||
|
||||
# Check key-level permissions
|
||||
if valid_token is not None and valid_token.object_permission_id is not None:
|
||||
key_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": valid_token.object_permission_id},
|
||||
)
|
||||
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(
|
||||
|
|
@ -3645,10 +3652,12 @@ async def search_tool_access_check(
|
|||
else None
|
||||
)
|
||||
if team_object_permission_id is not None:
|
||||
team_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": team_object_permission_id},
|
||||
)
|
||||
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(
|
||||
|
|
|
|||
|
|
@ -269,7 +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[] @default([]) // Search tool names allowed for this key/team/org. Empty = no access (least privilege).
|
||||
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([])
|
||||
|
|
|
|||
|
|
@ -22,7 +22,12 @@ async def _get_allowed_search_tool_names(
|
|||
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
|
||||
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
|
||||
|
|
@ -30,25 +35,27 @@ async def _get_allowed_search_tool_names(
|
|||
key_allowed: Optional[List[str]] = None
|
||||
team_allowed: Optional[List[str]] = None
|
||||
|
||||
# Key-level permissions
|
||||
# Key-level permissions (via cached helper)
|
||||
if user_api_key_dict.object_permission_id is not None:
|
||||
key_perm = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={
|
||||
"object_permission_id": user_api_key_dict.object_permission_id
|
||||
},
|
||||
)
|
||||
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
|
||||
# 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 prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": team_perm_id},
|
||||
)
|
||||
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
|
||||
|
|
@ -226,6 +233,14 @@ async def search(
|
|||
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)
|
||||
|
|
|
|||
|
|
@ -269,7 +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[] @default([]) // Search tool names allowed for this key/team/org. Empty = no access (least privilege).
|
||||
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([])
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from typing import List, Optional
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from starlette import status
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
|
|
@ -149,10 +150,14 @@ class TestCanObjectCallSearchTools:
|
|||
# ===========================================================================
|
||||
|
||||
|
||||
_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("litellm.proxy.proxy_server.prisma_client", None):
|
||||
with patch(f"{_PROXY_SERVER}.prisma_client", None):
|
||||
result = await search_tool_access_check(
|
||||
search_tool_name="any-tool",
|
||||
valid_token=None,
|
||||
|
|
@ -163,8 +168,9 @@ async def test_should_allow_when_no_prisma_client():
|
|||
@pytest.mark.asyncio
|
||||
async def test_should_allow_when_no_token():
|
||||
"""No valid_token → allow."""
|
||||
mock_prisma = MagicMock()
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
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,
|
||||
|
|
@ -175,12 +181,13 @@ async def test_should_allow_when_no_token():
|
|||
@pytest.mark.asyncio
|
||||
async def test_should_allow_when_no_permission_ids():
|
||||
"""Token with no object_permission_id or team_object_permission_id → allow."""
|
||||
mock_prisma = MagicMock()
|
||||
token = UserAPIKeyAuth(
|
||||
object_permission_id=None,
|
||||
team_object_permission_id=None,
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
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,
|
||||
|
|
@ -191,18 +198,17 @@ async def test_should_allow_when_no_permission_ids():
|
|||
@pytest.mark.asyncio
|
||||
async def test_should_allow_key_with_matching_permission():
|
||||
"""Key with object_permission that includes the tool → allow."""
|
||||
mock_prisma = MagicMock()
|
||||
mock_perm = MagicMock()
|
||||
mock_perm.search_tools = ["my-tool"]
|
||||
mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
return_value=mock_perm
|
||||
)
|
||||
|
||||
token = UserAPIKeyAuth(
|
||||
object_permission_id="key-perm-id",
|
||||
team_object_permission_id=None,
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
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,
|
||||
|
|
@ -213,18 +219,17 @@ async def test_should_allow_key_with_matching_permission():
|
|||
@pytest.mark.asyncio
|
||||
async def test_should_deny_key_with_empty_search_tools():
|
||||
"""Key with empty search_tools → deny."""
|
||||
mock_prisma = MagicMock()
|
||||
mock_perm = MagicMock()
|
||||
mock_perm.search_tools = []
|
||||
mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
return_value=mock_perm
|
||||
)
|
||||
|
||||
token = UserAPIKeyAuth(
|
||||
object_permission_id="key-perm-id",
|
||||
team_object_permission_id=None,
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
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",
|
||||
|
|
@ -236,18 +241,17 @@ async def test_should_deny_key_with_empty_search_tools():
|
|||
@pytest.mark.asyncio
|
||||
async def test_should_deny_team_with_empty_search_tools():
|
||||
"""Team with empty search_tools → deny."""
|
||||
mock_prisma = MagicMock()
|
||||
mock_team_perm = MagicMock()
|
||||
mock_team_perm.search_tools = []
|
||||
mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
return_value=mock_team_perm
|
||||
)
|
||||
|
||||
token = UserAPIKeyAuth(
|
||||
object_permission_id=None,
|
||||
team_object_permission_id="team-perm-id",
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
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",
|
||||
|
|
@ -259,30 +263,27 @@ async def test_should_deny_team_with_empty_search_tools():
|
|||
@pytest.mark.asyncio
|
||||
async def test_should_deny_when_key_allows_but_team_denies():
|
||||
"""Key allows, team denies → deny (team check second)."""
|
||||
mock_prisma = MagicMock()
|
||||
|
||||
key_perm = MagicMock()
|
||||
key_perm.search_tools = ["the-tool"]
|
||||
|
||||
team_perm = MagicMock()
|
||||
team_perm.search_tools = []
|
||||
|
||||
async def mock_find(where):
|
||||
if where["object_permission_id"] == "key-perm-id":
|
||||
async def mock_get_perm(object_permission_id, **kwargs):
|
||||
if object_permission_id == "key-perm-id":
|
||||
return key_perm
|
||||
elif where["object_permission_id"] == "team-perm-id":
|
||||
elif object_permission_id == "team-perm-id":
|
||||
return team_perm
|
||||
return None
|
||||
|
||||
mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
side_effect=mock_find
|
||||
)
|
||||
|
||||
token = UserAPIKeyAuth(
|
||||
object_permission_id="key-perm-id",
|
||||
team_object_permission_id="team-perm-id",
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
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",
|
||||
|
|
@ -294,30 +295,27 @@ async def test_should_deny_when_key_allows_but_team_denies():
|
|||
@pytest.mark.asyncio
|
||||
async def test_should_allow_when_both_key_and_team_allow():
|
||||
"""Both key and team allow → allow."""
|
||||
mock_prisma = MagicMock()
|
||||
|
||||
key_perm = MagicMock()
|
||||
key_perm.search_tools = ["tool-x"]
|
||||
|
||||
team_perm = MagicMock()
|
||||
team_perm.search_tools = ["tool-x", "tool-y"]
|
||||
|
||||
async def mock_find(where):
|
||||
if where["object_permission_id"] == "key-perm-id":
|
||||
async def mock_get_perm(object_permission_id, **kwargs):
|
||||
if object_permission_id == "key-perm-id":
|
||||
return key_perm
|
||||
elif where["object_permission_id"] == "team-perm-id":
|
||||
elif object_permission_id == "team-perm-id":
|
||||
return team_perm
|
||||
return None
|
||||
|
||||
mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
side_effect=mock_find
|
||||
)
|
||||
|
||||
token = UserAPIKeyAuth(
|
||||
object_permission_id="key-perm-id",
|
||||
team_object_permission_id="team-perm-id",
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
|
||||
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,
|
||||
|
|
@ -378,18 +376,18 @@ class TestSearchToolErrorTypes:
|
|||
|
||||
|
||||
class TestVectorStoreAccessNotBroken:
|
||||
"""should preserve existing vector store semantics: empty list = allow ALL."""
|
||||
"""Vector store access follows same semantics: None = all access, [] = no access."""
|
||||
|
||||
def test_should_allow_all_when_vector_stores_is_empty(self):
|
||||
"""Vector stores: empty list = access to ALL (existing behavior)."""
|
||||
def test_should_deny_all_when_vector_stores_is_empty(self):
|
||||
"""Vector stores: empty list = no access (matches search_tools semantics)."""
|
||||
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
|
||||
with pytest.raises(ProxyException):
|
||||
_can_object_call_vector_stores(
|
||||
object_type="key",
|
||||
vector_store_ids_to_run=["store-1"],
|
||||
object_permissions=perm,
|
||||
)
|
||||
|
||||
def test_should_allow_when_vector_stores_is_none(self):
|
||||
perm = MagicMock()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue