This commit is contained in:
Timothy Jaeryang Baek
2026-04-17 14:09:35 +09:00
parent f44b7a01f5
commit e7e752f8e7
+39 -34
View File
@@ -3,7 +3,7 @@ import time
import uuid
from typing import Optional
from sqlalchemy import select, delete, update, or_, func, cast, String
from sqlalchemy import select, delete, update, or_, func, text, cast, String
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.groups import Groups
@@ -260,12 +260,12 @@ class PromptsTable:
) -> PromptListResponse:
async with get_async_db_context(db) as db:
# Join with User table for user filtering and sorting
stmt = select(Prompt, User).outerjoin(User, User.id == Prompt.user_id)
query = select(Prompt, User).outerjoin(User, User.id == Prompt.user_id)
if filter:
query_key = filter.get('query')
if query_key:
stmt = stmt.filter(
query = query.filter(
or_(
Prompt.name.ilike(f'%{query_key}%'),
Prompt.command.ilike(f'%{query_key}%'),
@@ -277,14 +277,14 @@ class PromptsTable:
view_option = filter.get('view_option')
if view_option == 'created':
stmt = stmt.filter(Prompt.user_id == user_id)
query = query.filter(Prompt.user_id == user_id)
elif view_option == 'shared':
stmt = stmt.filter(Prompt.user_id != user_id)
query = query.filter(Prompt.user_id != user_id)
# Apply access grant filtering
stmt = AccessGrants.has_permission_filter(
query = AccessGrants.has_permission_filter(
db=db,
query=stmt,
query=query,
DocumentModel=Prompt,
filter=filter,
resource_type='prompt',
@@ -293,56 +293,61 @@ class PromptsTable:
tag = filter.get('tag')
if tag:
# SQLite stores JSON text via json.dumps(ensure_ascii=True),
# so non-ASCII chars are \uXXXX-escaped. PostgreSQL native JSONB
# stores literal Unicode. Use the right pattern for each.
if db.bind.dialect.name == 'sqlite':
if tag.isascii():
tags_text = func.lower(cast(Prompt.tags, String))
pattern = f'%{json.dumps(tag.lower())}%'
else:
# LOWER() is ASCII-only; non-ASCII codepoints would
# produce different \uXXXX escapes when lowered.
tags_text = cast(Prompt.tags, String)
pattern = f'%{json.dumps(tag)}%'
bind = await db.connection()
dialect_name = bind.dialect.name
tag_lower = tag.lower()
if dialect_name == 'sqlite':
tag_clause = text(
"EXISTS (SELECT 1 FROM json_each(prompt.tags) t WHERE LOWER(t.value) = :tag_val)"
)
elif dialect_name == 'postgresql':
tag_clause = text(
"EXISTS (SELECT 1 FROM json_array_elements_text(prompt.tags) t WHERE LOWER(t) = :tag_val)"
)
else:
tags_text = func.lower(cast(Prompt.tags, String))
pattern = f'%{json.dumps(tag.lower(), ensure_ascii=False)}%'
stmt = stmt.filter(tags_text.like(pattern))
# Fallback: LIKE on serialised JSON text (ASCII-safe only)
tag_clause = func.lower(cast(Prompt.tags, String)).like(f'%{json.dumps(tag_lower, ensure_ascii=False)}%')
tag_lower = None
if tag_lower is not None:
query = query.filter(tag_clause.params(tag_val=tag_lower))
else:
query = query.filter(tag_clause)
order_by = filter.get('order_by')
direction = filter.get('direction')
if order_by == 'name':
if direction == 'asc':
stmt = stmt.order_by(Prompt.name.asc())
query = query.order_by(Prompt.name.asc())
else:
stmt = stmt.order_by(Prompt.name.desc())
query = query.order_by(Prompt.name.desc())
elif order_by == 'created_at':
if direction == 'asc':
stmt = stmt.order_by(Prompt.created_at.asc())
query = query.order_by(Prompt.created_at.asc())
else:
stmt = stmt.order_by(Prompt.created_at.desc())
query = query.order_by(Prompt.created_at.desc())
elif order_by == 'updated_at':
if direction == 'asc':
stmt = stmt.order_by(Prompt.updated_at.asc())
query = query.order_by(Prompt.updated_at.asc())
else:
stmt = stmt.order_by(Prompt.updated_at.desc())
query = query.order_by(Prompt.updated_at.desc())
else:
stmt = stmt.order_by(Prompt.updated_at.desc())
query = query.order_by(Prompt.updated_at.desc())
else:
stmt = stmt.order_by(Prompt.updated_at.desc())
query = query.order_by(Prompt.updated_at.desc())
# Count BEFORE pagination
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
total = count_result.scalar()
if skip:
stmt = stmt.offset(skip)
query = query.offset(skip)
if limit:
stmt = stmt.limit(limit)
query = query.limit(limit)
result = await db.execute(stmt)
result = await db.execute(query)
items = result.all()
prompt_ids = [prompt.id for prompt, _ in items]