From 4807866a1cf47340f1b5ea76fded95f8114305f9 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 23 Aug 2026 01:31:30 -0400 Subject: [PATCH] refac Co-Authored-By: Classic298 <27028174+Classic298@users.noreply.github.com> --- backend/open_webui/models/tools.py | 56 +++++++++++++-------- backend/open_webui/routers/tools.py | 78 ++++++++++++++++------------- 2 files changed, 76 insertions(+), 58 deletions(-) diff --git a/backend/open_webui/models/tools.py b/backend/open_webui/models/tools.py index cbc21854ea..24e019aea9 100644 --- a/backend/open_webui/models/tools.py +++ b/backend/open_webui/models/tools.py @@ -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: diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index fb44bfc0a7..53ef0bdfdf 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -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, + ) ############################