Skip to content

Commit e7e752f

Browse files
committed
refac
1 parent f44b7a0 commit e7e752f

1 file changed

Lines changed: 39 additions & 34 deletions

File tree

backend/open_webui/models/prompts.py

Lines changed: 39 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33
import uuid
44
from typing import Optional
55

6-
from sqlalchemy import select, delete, update, or_, func, cast, String
6+
from sqlalchemy import select, delete, update, or_, func, text, cast, String
77
from sqlalchemy.ext.asyncio import AsyncSession
88
from open_webui.internal.db import Base, JSONField, get_async_db_context
99
from open_webui.models.groups import Groups
@@ -260,12 +260,12 @@ async def search_prompts(
260260
) -> PromptListResponse:
261261
async with get_async_db_context(db) as db:
262262
# Join with User table for user filtering and sorting
263-
stmt = select(Prompt, User).outerjoin(User, User.id == Prompt.user_id)
263+
query = select(Prompt, User).outerjoin(User, User.id == Prompt.user_id)
264264

265265
if filter:
266266
query_key = filter.get('query')
267267
if query_key:
268-
stmt = stmt.filter(
268+
query = query.filter(
269269
or_(
270270
Prompt.name.ilike(f'%{query_key}%'),
271271
Prompt.command.ilike(f'%{query_key}%'),
@@ -277,14 +277,14 @@ async def search_prompts(
277277

278278
view_option = filter.get('view_option')
279279
if view_option == 'created':
280-
stmt = stmt.filter(Prompt.user_id == user_id)
280+
query = query.filter(Prompt.user_id == user_id)
281281
elif view_option == 'shared':
282-
stmt = stmt.filter(Prompt.user_id != user_id)
282+
query = query.filter(Prompt.user_id != user_id)
283283

284284
# Apply access grant filtering
285-
stmt = AccessGrants.has_permission_filter(
285+
query = AccessGrants.has_permission_filter(
286286
db=db,
287-
query=stmt,
287+
query=query,
288288
DocumentModel=Prompt,
289289
filter=filter,
290290
resource_type='prompt',
@@ -293,56 +293,61 @@ async def search_prompts(
293293

294294
tag = filter.get('tag')
295295
if tag:
296-
# SQLite stores JSON text via json.dumps(ensure_ascii=True),
297-
# so non-ASCII chars are \uXXXX-escaped. PostgreSQL native JSONB
298-
# stores literal Unicode. Use the right pattern for each.
299-
if db.bind.dialect.name == 'sqlite':
300-
if tag.isascii():
301-
tags_text = func.lower(cast(Prompt.tags, String))
302-
pattern = f'%{json.dumps(tag.lower())}%'
303-
else:
304-
# LOWER() is ASCII-only; non-ASCII codepoints would
305-
# produce different \uXXXX escapes when lowered.
306-
tags_text = cast(Prompt.tags, String)
307-
pattern = f'%{json.dumps(tag)}%'
296+
bind = await db.connection()
297+
dialect_name = bind.dialect.name
298+
tag_lower = tag.lower()
299+
300+
if dialect_name == 'sqlite':
301+
tag_clause = text(
302+
"EXISTS (SELECT 1 FROM json_each(prompt.tags) t WHERE LOWER(t.value) = :tag_val)"
303+
)
304+
elif dialect_name == 'postgresql':
305+
tag_clause = text(
306+
"EXISTS (SELECT 1 FROM json_array_elements_text(prompt.tags) t WHERE LOWER(t) = :tag_val)"
307+
)
308+
else:
309+
# Fallback: LIKE on serialised JSON text (ASCII-safe only)
310+
tag_clause = func.lower(cast(Prompt.tags, String)).like(f'%{json.dumps(tag_lower, ensure_ascii=False)}%')
311+
tag_lower = None
312+
313+
if tag_lower is not None:
314+
query = query.filter(tag_clause.params(tag_val=tag_lower))
308315
else:
309-
tags_text = func.lower(cast(Prompt.tags, String))
310-
pattern = f'%{json.dumps(tag.lower(), ensure_ascii=False)}%'
311-
stmt = stmt.filter(tags_text.like(pattern))
316+
query = query.filter(tag_clause)
312317

313318
order_by = filter.get('order_by')
314319
direction = filter.get('direction')
315320

316321
if order_by == 'name':
317322
if direction == 'asc':
318-
stmt = stmt.order_by(Prompt.name.asc())
323+
query = query.order_by(Prompt.name.asc())
319324
else:
320-
stmt = stmt.order_by(Prompt.name.desc())
325+
query = query.order_by(Prompt.name.desc())
321326
elif order_by == 'created_at':
322327
if direction == 'asc':
323-
stmt = stmt.order_by(Prompt.created_at.asc())
328+
query = query.order_by(Prompt.created_at.asc())
324329
else:
325-
stmt = stmt.order_by(Prompt.created_at.desc())
330+
query = query.order_by(Prompt.created_at.desc())
326331
elif order_by == 'updated_at':
327332
if direction == 'asc':
328-
stmt = stmt.order_by(Prompt.updated_at.asc())
333+
query = query.order_by(Prompt.updated_at.asc())
329334
else:
330-
stmt = stmt.order_by(Prompt.updated_at.desc())
335+
query = query.order_by(Prompt.updated_at.desc())
331336
else:
332-
stmt = stmt.order_by(Prompt.updated_at.desc())
337+
query = query.order_by(Prompt.updated_at.desc())
333338
else:
334-
stmt = stmt.order_by(Prompt.updated_at.desc())
339+
query = query.order_by(Prompt.updated_at.desc())
335340

336341
# Count BEFORE pagination
337-
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
342+
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
338343
total = count_result.scalar()
339344

340345
if skip:
341-
stmt = stmt.offset(skip)
346+
query = query.offset(skip)
342347
if limit:
343-
stmt = stmt.limit(limit)
348+
query = query.limit(limit)
344349

345-
result = await db.execute(stmt)
350+
result = await db.execute(query)
346351
items = result.all()
347352

348353
prompt_ids = [prompt.id for prompt, _ in items]

0 commit comments

Comments
 (0)