diff --git a/backend/open_webui/models/prompts.py b/backend/open_webui/models/prompts.py index 47f74295db..70c05094b8 100644 --- a/backend/open_webui/models/prompts.py +++ b/backend/open_webui/models/prompts.py @@ -663,12 +663,38 @@ class PromptsTable: async def get_tags(self, db: Optional[AsyncSession] = None) -> list[str]: try: async with get_async_db_context(db) as db: - result = await db.execute(select(Prompt).filter_by(is_active=True)) - prompts = result.scalars().all() + result = await db.execute(select(Prompt.tags).filter(Prompt.is_active == True)) tags = set() - for prompt in prompts: - if prompt.tags: - for tag in prompt.tags: + for (tag_list,) in result.all(): + if tag_list: + for tag in tag_list: + if tag: + tags.add(tag) + return sorted(list(tags)) + except Exception: + return [] + + async def get_tags_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[str]: + try: + async with get_async_db_context(db) as db: + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_group_ids = [group.id for group in user_groups] + + query = select(Prompt.tags).filter(Prompt.is_active == True) + query = AccessGrants.has_permission_filter( + db=db, + query=query, + DocumentModel=Prompt, + filter={'user_id': user_id, 'group_ids': user_group_ids}, + resource_type='prompt', + permission='read', + ) + + result = await db.execute(query) + tags = set() + for (tag_list,) in result.all(): + if tag_list: + for tag in tag_list: if tag: tags.add(tag) return sorted(list(tags)) diff --git a/backend/open_webui/routers/prompts.py b/backend/open_webui/routers/prompts.py index 11901fc5a7..755034f880 100644 --- a/backend/open_webui/routers/prompts.py +++ b/backend/open_webui/routers/prompts.py @@ -61,13 +61,7 @@ async def get_prompts(user=Depends(get_verified_user), db: AsyncSession = Depend async def get_prompt_tags(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL: return await Prompts.get_tags(db=db) - else: - prompts = await Prompts.get_prompts_by_user_id(user.id, 'read', db=db) - tags = set() - for prompt in prompts: - if prompt.tags: - tags.update(prompt.tags) - return sorted(list(tags)) + return await Prompts.get_tags_by_user_id(user.id, db=db) @router.get('/list', response_model=PromptAccessListResponse)