From cc94a90b4d4d690bc7cb9f7124f2d6e552973970 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Tue, 19 May 2026 21:03:23 +0400 Subject: [PATCH] refac Co-Authored-By: Algorithm5838 <108630393+Algorithm5838@users.noreply.github.com> --- backend/open_webui/models/tools.py | 17 +++++++++++++++++ backend/open_webui/utils/tools.py | 5 ++++- 2 files changed, 21 insertions(+), 1 deletion(-) diff --git a/backend/open_webui/models/tools.py b/backend/open_webui/models/tools.py index 9a1a202292..c73a2d7bb7 100644 --- a/backend/open_webui/models/tools.py +++ b/backend/open_webui/models/tools.py @@ -151,6 +151,23 @@ class ToolsTable: except Exception: return None + async def get_tools_by_ids( + self, tool_ids: list[str], db: AsyncSession | None = None + ) -> dict[str, ToolModel]: + """Batch-fetch multiple tools by ID, returning a dict keyed by tool ID.""" + if not tool_ids: + return {} + async with get_async_db_context(db) as db: + result = await db.execute(select(Tool).where(Tool.id.in_(tool_ids))) + tools = result.scalars().all() + grants_map = await AccessGrants.get_grants_by_resources( + 'tool', [tool.id for tool in tools], db=db + ) + return { + tool.id: await self._to_tool_model(tool, access_grants=grants_map.get(tool.id, []), db=db) + for tool in tools + } + async def get_tools(self, defer_content: bool = False, db: AsyncSession | None = None) -> list[ToolUserModel]: async with get_async_db_context(db) as db: stmt = select(Tool).order_by(Tool.updated_at.desc()) diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index da91834102..4017d0c48d 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -166,8 +166,11 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr # Get user's group memberships for access control checks user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + # Batch-fetch all DB tools in one query instead of one per tool_id + tool_models = await Tools.get_tools_by_ids(tool_ids) + for tool_id in tool_ids: - tool = await Tools.get_tool_by_id(tool_id) + tool = tool_models.get(tool_id) if tool: # Check access control for local tools if (