mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-28 05:27:35 +00:00
refac
Co-Authored-By: Classic298 <27028174+Classic298@users.noreply.github.com>
This commit is contained in:
parent
7abe11346a
commit
4807866a1c
2 changed files with 76 additions and 58 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue