33import uuid
44from 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
77from sqlalchemy .ext .asyncio import AsyncSession
88from open_webui .internal .db import Base , JSONField , get_async_db_context
99from 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