Co-Authored-By: Classic298 <27028174+Classic298@users.noreply.github.com>
This commit is contained in:
Timothy Jaeryang Baek 2026-08-23 01:31:30 -04:00
parent 7abe11346a
commit 4807866a1c
2 changed files with 76 additions and 58 deletions

View file

@ -169,20 +169,37 @@ class ToolsTable:
for tool in tools
}
async def get_tools(self, defer_content: bool = False, db: AsyncSession | None = None) -> list[ToolUserModel]:
async def get_tools(
self,
defer_content: bool = False,
db: AsyncSession | None = None,
user_id: str | None = None,
user_group_ids: set[str] | None = None,
permission: str = 'read',
) -> list[ToolUserModel]:
async with get_async_db_context(db) as db:
if defer_content:
# Skip Tool.content (plugin source, potentially large) via a
# column select; Row attributes satisfy from_attributes.
result = await db.execute(
select(
Tool.id, Tool.user_id, Tool.name, Tool.specs, Tool.meta, Tool.updated_at, Tool.created_at
).order_by(Tool.updated_at.desc())
# Skip Tool.content (plugin source, potentially large) via a
# column select; Row attributes satisfy from_attributes.
stmt = (
select(Tool.id, Tool.user_id, Tool.name, Tool.specs, Tool.meta, Tool.updated_at, Tool.created_at)
if defer_content
else select(Tool)
).order_by(Tool.updated_at.desc())
if user_id is not None:
if user_group_ids is None:
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)}
stmt = AccessGrants.has_permission_filter(
db=db,
query=stmt,
DocumentModel=Tool,
filter={'user_id': user_id, 'group_ids': user_group_ids},
resource_type='tool',
permission=permission,
)
all_tools = result.all()
else:
result = await db.execute(select(Tool).order_by(Tool.updated_at.desc()))
all_tools = result.scalars().all()
result = await db.execute(stmt)
all_tools = result.all() if defer_content else result.scalars().all()
user_ids = list(set(tool.user_id for tool in all_tools))
tool_ids = [tool.id for tool in all_tools]
@ -217,20 +234,15 @@ class ToolsTable:
defer_content: bool = False,
db: AsyncSession | None = None,
) -> list[ToolUserModel]:
tools = await self.get_tools(defer_content=defer_content, db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
# One grants query for all non-owned tools instead of one per tool
accessible_ids = await AccessGrants.get_accessible_resource_ids(
user_id=user_id,
resource_type='tool',
resource_ids=[tool.id for tool in tools if tool.user_id != user_id],
permission=permission,
user_group_ids=user_group_ids,
return await self.get_tools(
defer_content=defer_content,
db=db,
user_id=user_id,
user_group_ids=user_group_ids,
permission=permission,
)
return [tool for tool in tools if tool.user_id == user_id or tool.id in accessible_ids]
async def get_tool_valves_by_id(self, id: str, db: AsyncSession | None = None) -> dict | None:
try:

View file

@ -71,11 +71,22 @@ async def get_tools(
db: AsyncSession = Depends(get_async_session),
):
tools = []
bypass_access_control = user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL
user_group_ids = (
set()
if bypass_access_control
else {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
)
# Local Tools
if ENABLE_PLUGINS:
tools_cache = get_tools_cache(request)
for tool in await Tools.get_tools(defer_content=True, db=db):
for tool in await Tools.get_tools(
defer_content=True,
db=db,
user_id=None if bypass_access_control else user.id,
user_group_ids=user_group_ids,
):
tool_module = tools_cache.get(tool.id)
has_user_valves = (
hasattr(tool_module, 'UserValves')
@ -166,31 +177,19 @@ async def get_tools(
)
)
if not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL):
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
filtered_tools = []
for tool in tools:
if tool.user_id == user.id:
filtered_tools.append(tool)
elif str(tool.id).startswith('server:'):
if await has_access(
user.id,
'read',
server_access_grants.get(str(tool.id), []),
user_group_ids,
db=db,
):
filtered_tools.append(tool)
elif await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tool.id,
permission='read',
user_group_ids=user_group_ids,
if not bypass_access_control:
tools = [
tool
for tool in tools
if not str(tool.id).startswith('server:')
or await has_access(
user.id,
'read',
server_access_grants.get(str(tool.id), []),
user_group_ids,
db=db,
):
filtered_tools.append(tool)
tools = filtered_tools
)
]
if query:
q = query.casefold()
@ -209,17 +208,23 @@ async def get_tool_list(user=Depends(get_verified_user), db: AsyncSession = Depe
if not ENABLE_PLUGINS:
return []
if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL:
tools = await Tools.get_tools(defer_content=True, db=db)
else:
tools = await Tools.get_tools_by_user_id(user.id, 'read', defer_content=True, db=db)
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
bypass_access_control = user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL
user_group_ids = (
set()
if bypass_access_control
else {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
)
tools = await Tools.get_tools(
defer_content=True,
db=db,
user_id=None if bypass_access_control else user.id,
user_group_ids=user_group_ids,
)
result = []
for tool in tools:
has_write = (
(user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
bypass_access_control
or user.id == tool.user_id
or any(
g.permission == 'write'
@ -334,10 +339,11 @@ async def export_tools(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL:
return await Tools.get_tools(db=db)
else:
return await Tools.get_tools_by_user_id(user.id, 'read', db=db)
bypass_access_control = user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL
return await Tools.get_tools(
db=db,
user_id=None if bypass_access_control else user.id,
)
############################