Coverage for open_webui/models/prompts.py: 52%
381 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1"""Prompt template models, forms, and database operations."""
3from __future__ import annotations
5import logging
6import time
7import uuid
8from typing import Optional
10log = logging.getLogger(__name__)
12from open_webui.internal.db import Base, JSONField, get_async_db_context
13from open_webui.models.access_grants import AccessGrantModel, AccessGrants
14from open_webui.models.groups import Groups
15from open_webui.models.prompt_history import PromptHistories
16from open_webui.models.users import User, UserModel, UserResponse, Users
17from open_webui.utils.misc import json_text_variants
18from pydantic import BaseModel, ConfigDict, Field
19from sqlalchemy import JSON, BigInteger, Boolean, Column, String, Text, cast, delete, func, or_, select, text, update
20from sqlalchemy.ext.asyncio import AsyncSession
23class Prompt(Base): # versioned template
24 """Slash-command prompt with history tracking and access control."""
26 __tablename__ = 'prompt'
28 id = Column(Text, primary_key=True)
29 command = Column(String, unique=True, index=True)
30 user_id = Column(String, index=True) # owner user id
31 name = Column(Text)
32 content = Column(Text) # the prompt template body
33 data = Column(JSON, nullable=True) # structured prompt parameters
34 meta = Column(JSON, nullable=True) # freeform metadata (description, etc.)
35 tags = Column(JSON, nullable=True)
36 is_active = Column(Boolean, default=True)
37 version_id = Column(Text, nullable=True) # Points to active history entry
38 created_at = Column(BigInteger, nullable=True)
39 updated_at = Column(BigInteger, nullable=True)
42class PromptModel(BaseModel):
43 id: str | None = None
44 command: str
45 user_id: str
46 name: str
47 content: str
48 data: dict | None = None
49 meta: dict | None = None
50 tags: list[str] | None = None
51 is_active: bool | None = True
52 version_id: str | None = None
53 created_at: int | None = None
54 updated_at: int | None = None
55 access_grants: list[AccessGrantModel] = Field(default_factory=list)
57 model_config = ConfigDict(from_attributes=True) # allows ORM model binding
60# --- form / schema definitions ---
61# Forms
62####################
65class PromptUserResponse(PromptModel):
66 user: UserResponse | None = None
69class PromptAccessResponse(PromptUserResponse):
70 write_access: bool | None = False
73class PromptListResponse(BaseModel):
74 items: list[PromptUserResponse]
75 total: int
78class PromptAccessListResponse(BaseModel):
79 items: list[PromptAccessResponse]
80 total: int
83class PromptForm(BaseModel):
84 command: str
85 name: str # Changed from title
86 content: str
87 data: dict | None = None
88 meta: dict | None = None
89 tags: list[str] | None = None
90 access_grants: list[dict] | None = None
91 version_id: str | None = None # Active version
92 commit_message: str | None = None # For history tracking
93 is_production: bool | None = True # Whether to set new version as production
96class PromptsTable:
97 async def _get_access_grants(self, prompt_id: str, db: AsyncSession | None = None) -> list[AccessGrantModel]:
98 return await AccessGrants.get_grants_by_resource('prompt', prompt_id, db=db)
100 async def _to_prompt_model(
101 self,
102 prompt: Prompt,
103 access_grants: list[AccessGrantModel] | None = None,
104 db: AsyncSession | None = None,
105 ) -> PromptModel:
106 prompt_model = PromptModel.model_validate(prompt)
107 prompt_model.access_grants = (
108 access_grants if access_grants is not None else await self._get_access_grants(prompt_model.id, db=db)
109 )
110 return prompt_model
112 async def insert_new_prompt(
113 self, user_id: str, form_data: PromptForm, db: AsyncSession | None = None
114 ) -> PromptModel | None:
115 now = int(time.time())
116 prompt_id = str(uuid.uuid4())
118 async with get_async_db_context(db) as session:
119 try:
120 record = Prompt(
121 id=prompt_id,
122 user_id=user_id,
123 command=form_data.command,
124 name=form_data.name,
125 content=form_data.content,
126 data=form_data.data or {},
127 meta=form_data.meta or {},
128 tags=form_data.tags or [],
129 is_active=True,
130 created_at=now,
131 updated_at=now,
132 )
133 session.add(record)
134 await session.commit()
136 await AccessGrants.set_access_grants(
137 'prompt',
138 prompt_id,
139 form_data.access_grants,
140 db=session,
141 ) # persist sharing rules
143 if not record: # shouldn't happen, but guard anyway 143 ↛ 144line 143 didn't jump to line 144 because the condition on line 143 was never true
144 return None
146 # Build the initial version snapshot.
147 grants = await self._get_access_grants(prompt_id, db=session)
148 snapshot = {
149 'name': form_data.name,
150 'content': form_data.content,
151 'command': form_data.command,
152 'data': form_data.data or {},
153 'meta': form_data.meta or {},
154 'tags': form_data.tags or [],
155 'access_grants': [g.model_dump() for g in grants],
156 }
158 history_entry = await PromptHistories.create_history_entry(
159 prompt_id=prompt_id,
160 snapshot=snapshot,
161 user_id=user_id,
162 parent_id=None,
163 commit_message=form_data.commit_message or 'Initial version',
164 db=session,
165 ) # creates the first version entry
167 # Pin the initial history entry as the production version.
168 if history_entry: 168 ↛ 172line 168 didn't jump to line 172 because the condition on line 168 was always true
169 record.version_id = history_entry.id
170 await session.commit()
172 return await self._to_prompt_model(record, db=session)
173 except Exception as e:
174 log.exception('Error creating prompt: %s', e)
175 return None
177 async def get_prompt_by_id(self, prompt_id: str, db: AsyncSession | None = None) -> PromptModel | None:
178 try:
179 async with get_async_db_context(db) as session:
180 result = await session.execute(
181 select(Prompt).filter_by(id=prompt_id),
182 )
183 prompt = result.scalars().first() # None when not found
184 if not prompt:
185 return None
186 return await self._to_prompt_model(prompt, db=session)
187 except Exception: # connection / integrity error
188 return
190 async def get_prompt_by_command(self, command: str, db: AsyncSession | None = None) -> PromptModel | None:
191 """Look up a prompt by its unique slash-command string."""
192 async with get_async_db_context(db) as session:
193 match = (await session.execute(select(Prompt).where(Prompt.command == command))).scalars().first()
194 if match is None:
195 return
196 return await self._to_prompt_model(match, db=session)
197 # --- context manager always returns above ---
198 return
200 async def get_prompts(self, db: AsyncSession | None = None) -> list[PromptUserResponse]:
201 """Return all active prompts ordered by most recently updated."""
202 async with get_async_db_context(db) as session:
203 active = (
204 (
205 await session.execute(
206 select(Prompt).where(Prompt.is_active.is_(True)).order_by(Prompt.updated_at.desc())
207 )
208 )
209 .scalars()
210 .all()
211 )
213 user_ids = list(set(p.user_id for p in active))
214 prompt_ids = [p.id for p in active]
216 users = await Users.get_users_by_user_ids(user_ids, db=session) if user_ids else []
217 users_dict = {u.id: u for u in users}
218 grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=session)
220 prompts = []
221 for prompt in active:
222 user = users_dict.get(prompt.user_id)
223 prompts.append(
224 PromptUserResponse.model_validate(
225 {
226 **(
227 await self._to_prompt_model(
228 prompt,
229 access_grants=grants_map.get(prompt.id, []),
230 db=session,
231 )
232 ).model_dump(),
233 'user': user.model_dump() if user else None,
234 }
235 )
236 )
238 return prompts
240 async def get_prompts_by_user_id(
241 self, user_id: str, permission: str = 'write', db: AsyncSession | None = None
242 ) -> list[PromptUserResponse]:
243 async with get_async_db_context(db) as session:
244 user_groups = await Groups.get_groups_by_member_id(user_id, db=session)
245 user_group_ids = [group.id for group in user_groups]
247 query = select(Prompt).filter(Prompt.is_active == True).order_by(Prompt.updated_at.desc())
248 query = AccessGrants.has_permission_filter(
249 db=db,
250 query=query,
251 DocumentModel=Prompt,
252 filter={'user_id': user_id, 'group_ids': user_group_ids},
253 resource_type='prompt',
254 permission=permission,
255 )
257 result = await session.execute(query)
258 accessible_prompts = result.scalars().all()
260 if not accessible_prompts:
261 return []
263 prompt_ids = [p.id for p in accessible_prompts]
264 owner_ids = list({p.user_id for p in accessible_prompts})
266 users = await Users.get_users_by_user_ids(owner_ids, db=session)
267 users_dict = {u.id: u for u in users}
268 grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=session)
270 results = []
271 for prompt in accessible_prompts:
272 user = users_dict.get(prompt.user_id)
273 results.append(
274 PromptUserResponse.model_validate(
275 {
276 **(
277 await self._to_prompt_model(
278 prompt,
279 access_grants=grants_map.get(prompt.id, []),
280 db=db,
281 )
282 ).model_dump(),
283 'user': user.model_dump() if user else None,
284 }
285 )
286 )
287 return results
289 async def search_prompts(
290 self,
291 user_id: str,
292 filter: dict = {},
293 skip: int = 0,
294 limit: int = 30,
295 db: AsyncSession | None = None,
296 ) -> PromptListResponse:
297 async with get_async_db_context(db) as session:
298 # Join with User table for user filtering and sorting
299 query = select(Prompt, User).outerjoin(User, User.id == Prompt.user_id)
301 if filter:
302 query_key = filter.get('query')
303 if query_key:
304 query = query.filter(
305 or_(
306 Prompt.name.ilike(f'%{query_key}%'),
307 Prompt.command.ilike(f'%{query_key}%'),
308 Prompt.content.ilike(f'%{query_key}%'),
309 User.name.ilike(f'%{query_key}%'),
310 User.email.ilike(f'%{query_key}%'),
311 )
312 )
314 view_option = filter.get('view_option')
315 if view_option == 'created': 315 ↛ 316line 315 didn't jump to line 316 because the condition on line 315 was never true
316 query = query.filter(Prompt.user_id == user_id)
317 elif view_option == 'shared': 317 ↛ 318line 317 didn't jump to line 318 because the condition on line 317 was never true
318 query = query.filter(Prompt.user_id != user_id)
320 # Apply access grant filtering
321 query = AccessGrants.has_permission_filter(
322 db=db,
323 query=query,
324 DocumentModel=Prompt,
325 filter=filter,
326 resource_type='prompt',
327 permission='read',
328 )
330 tag = filter.get('tag')
331 if tag:
332 bind = await session.connection()
333 dialect_name = bind.dialect.name
334 tag_lower = tag.lower()
336 if dialect_name == 'sqlite': 336 ↛ 341line 336 didn't jump to line 341 because the condition on line 336 was always true
337 tag_lower = tag.replace('\\', '\\\\').replace('%', '\\%').replace('_', '\\_')
338 tag_clause = text(
339 "EXISTS (SELECT 1 FROM json_each(prompt.tags) t WHERE t.value LIKE :tag_val ESCAPE '\\')"
340 )
341 elif dialect_name == 'postgresql':
342 tag_clause = text(
343 'EXISTS (SELECT 1 FROM json_array_elements_text(prompt.tags) t WHERE LOWER(t) = :tag_val)'
344 )
345 else:
346 # Fallback for dialects with no JSON array function: LIKE on the text.
347 tags_text = func.lower(cast(Prompt.tags, String))
348 tag_clause = or_(
349 *(tags_text.like(f'%"{variant}"%') for variant in json_text_variants(tag_lower))
350 )
351 tag_lower = None
353 if tag_lower is not None: 353 ↛ 356line 353 didn't jump to line 356 because the condition on line 353 was always true
354 query = query.filter(tag_clause.params(tag_val=tag_lower))
355 else:
356 query = query.filter(tag_clause)
358 order_by = filter.get('order_by')
359 direction = filter.get('direction')
361 if order_by == 'name': 361 ↛ 362line 361 didn't jump to line 362 because the condition on line 361 was never true
362 if direction == 'asc':
363 query = query.order_by(Prompt.name.asc())
364 else:
365 query = query.order_by(Prompt.name.desc())
366 elif order_by == 'created_at': 366 ↛ 367line 366 didn't jump to line 367 because the condition on line 366 was never true
367 if direction == 'asc':
368 query = query.order_by(Prompt.created_at.asc())
369 else:
370 query = query.order_by(Prompt.created_at.desc())
371 elif order_by == 'updated_at': 371 ↛ 372line 371 didn't jump to line 372 because the condition on line 371 was never true
372 if direction == 'asc':
373 query = query.order_by(Prompt.updated_at.asc())
374 else:
375 query = query.order_by(Prompt.updated_at.desc())
376 else:
377 query = query.order_by(Prompt.updated_at.desc())
378 else:
379 query = query.order_by(Prompt.updated_at.desc())
381 # Count BEFORE pagination
382 count_result = await session.execute(select(func.count()).select_from(query.subquery()))
383 total = count_result.scalar()
385 if skip:
386 query = query.offset(skip)
387 if limit:
388 query = query.limit(limit)
390 result = await session.execute(query)
391 items = result.all()
393 prompt_ids = [prompt.id for prompt, _ in items]
394 grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=session)
396 prompts = []
397 for prompt, user in items:
398 prompts.append(
399 PromptUserResponse(
400 **(
401 await self._to_prompt_model(
402 prompt,
403 access_grants=grants_map.get(prompt.id, []),
404 db=db,
405 )
406 ).model_dump(),
407 user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
408 )
409 )
411 return PromptListResponse(items=prompts, total=total)
413 async def update_prompt_by_command(
414 self,
415 command: str,
416 form_data: PromptForm,
417 user_id: str,
418 db: AsyncSession | None = None,
419 ) -> PromptModel | None:
420 if not command:
421 return None
422 try: # database transaction
423 async with get_async_db_context(db) as session:
424 result = await session.execute(select(Prompt).filter_by(command=command))
425 prompt = result.scalars().first()
426 if not prompt:
427 return None
429 latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=session)
430 parent_id = latest_history.id if latest_history else None
431 current_access_grants = await self._get_access_grants(prompt.id, db=session)
433 # Check if content changed to decide on history creation
434 content_changed = (
435 prompt.name != form_data.name
436 or prompt.content != form_data.content
437 or form_data.access_grants is not None
438 )
440 # Update prompt fields
441 prompt.name = form_data.name
442 prompt.content = form_data.content
443 prompt.data = form_data.data or prompt.data
444 prompt.meta = form_data.meta or prompt.meta
445 prompt.updated_at = int(time.time())
446 if form_data.access_grants is not None:
447 await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=session)
448 current_access_grants = await self._get_access_grants(prompt.id, db=session)
450 await session.commit()
452 # Create history entry only if content changed
453 if content_changed:
454 snapshot = {
455 'name': form_data.name,
456 'content': form_data.content,
457 'command': command,
458 'data': form_data.data or {},
459 'meta': form_data.meta or {},
460 'access_grants': [grant.model_dump() for grant in current_access_grants],
461 }
463 history_entry = await PromptHistories.create_history_entry(
464 prompt_id=prompt.id,
465 snapshot=snapshot,
466 user_id=user_id,
467 parent_id=parent_id,
468 commit_message=form_data.commit_message,
469 db=db,
470 )
472 # Set as production if flag is True (default)
473 if form_data.is_production and history_entry:
474 prompt.version_id = history_entry.id
475 await session.commit()
477 return await self._to_prompt_model(prompt, db=session)
478 except Exception:
479 return None
481 async def update_prompt_by_id(
482 self,
483 prompt_id: str,
484 form_data: PromptForm,
485 user_id: str,
486 db: AsyncSession | None = None,
487 ) -> PromptModel | None:
488 try:
489 async with get_async_db_context(db) as session:
490 result = await session.execute(select(Prompt).filter_by(id=prompt_id))
491 prompt = result.scalars().first()
492 if not prompt:
493 return None
495 latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=session)
496 parent_id = latest_history.id if latest_history else None
497 current_access_grants = await self._get_access_grants(prompt.id, db=session)
499 # Check if content changed to decide on history creation
500 content_changed = (
501 prompt.name != form_data.name
502 or prompt.command != form_data.command
503 or prompt.content != form_data.content
504 or form_data.access_grants is not None
505 or (form_data.tags is not None and prompt.tags != form_data.tags)
506 )
508 # Update prompt fields
509 prompt.command = form_data.command
511 if form_data.is_production:
512 prompt.name = form_data.name
513 prompt.content = form_data.content
514 prompt.data = form_data.data or prompt.data
515 prompt.meta = form_data.meta or prompt.meta
517 if form_data.tags is not None:
518 prompt.tags = form_data.tags
520 if form_data.access_grants is not None:
521 await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=session)
522 current_access_grants = await self._get_access_grants(prompt.id, db=session)
524 prompt.updated_at = int(time.time())
526 await session.commit()
528 # Create history entry only if content changed
529 if content_changed:
530 snapshot = {
531 'name': form_data.name,
532 'content': form_data.content,
533 'command': prompt.command,
534 'data': form_data.data or {},
535 'meta': form_data.meta or {},
536 'tags': form_data.tags if form_data.tags is not None else (prompt.tags or []),
537 'access_grants': [grant.model_dump() for grant in current_access_grants],
538 }
540 history_entry = await PromptHistories.create_history_entry(
541 prompt_id=prompt.id,
542 snapshot=snapshot,
543 user_id=user_id,
544 parent_id=parent_id,
545 commit_message=form_data.commit_message,
546 db=db,
547 )
549 # Set as production if flag is True (default)
550 if form_data.is_production and history_entry:
551 prompt.version_id = history_entry.id
552 await session.commit()
554 return await self._to_prompt_model(prompt, db=session)
555 except Exception:
556 return None
558 async def update_prompt_metadata(
559 self,
560 prompt_id: str,
561 name: str,
562 command: str,
563 tags: list[str] | None = None,
564 db: AsyncSession | None = None,
565 ) -> PromptModel | None:
566 """Update only name, command, and tags (no history created)."""
567 try:
568 async with get_async_db_context(db) as session:
569 result = await session.execute(select(Prompt).filter_by(id=prompt_id))
570 prompt = result.scalars().first()
571 if not prompt:
572 return None
574 prompt.name = name
575 prompt.command = command
577 if tags is not None:
578 prompt.tags = tags
580 prompt.updated_at = int(time.time())
581 await session.commit()
583 return await self._to_prompt_model(prompt, db=session)
584 except Exception:
585 return None
587 async def update_prompt_version(
588 self,
589 prompt_id: str,
590 version_id: str,
591 db: AsyncSession | None = None,
592 ) -> PromptModel | None:
593 """Set the active version of a prompt and restore content from that version's snapshot."""
594 try:
595 async with get_async_db_context(db) as session:
596 result = await session.execute(select(Prompt).filter_by(id=prompt_id))
597 prompt = result.scalars().first()
598 if not prompt:
599 return None
601 history_entry = await PromptHistories.get_history_entry_by_id(version_id, db=session)
603 # Reject a version_id from another prompt; restoring it would copy a foreign snapshot in.
604 if not history_entry or history_entry.prompt_id != prompt_id:
605 return None
607 # Restore prompt content from the snapshot
608 snapshot = history_entry.snapshot
609 if snapshot: 609 ↛ 617line 609 didn't jump to line 617 because the condition on line 609 was always true
610 prompt.name = snapshot.get('name', prompt.name)
611 prompt.content = snapshot.get('content', prompt.content)
612 prompt.data = snapshot.get('data', prompt.data)
613 prompt.meta = snapshot.get('meta', prompt.meta)
614 prompt.tags = snapshot.get('tags', prompt.tags)
615 # Note: command and access_grants are not restored from snapshot
617 prompt.version_id = version_id
618 prompt.updated_at = int(time.time())
619 await session.commit()
621 return await self._to_prompt_model(prompt, db=session)
622 except Exception as e: # connection error
623 log.error(f'Failed to restore prompt version: {e}')
624 return None # restoration failed
626 async def toggle_prompt_active(
627 self,
628 prompt_id: str,
629 db: AsyncSession | None = None,
630 ) -> PromptModel | None:
631 """Flip the is_active flag on a prompt."""
632 if not prompt_id: 632 ↛ 633line 632 didn't jump to line 633 because the condition on line 632 was never true
633 return None
634 try: # activation state toggle
635 async with get_async_db_context(db) as session:
636 result = await session.execute(select(Prompt).filter_by(id=prompt_id))
637 prompt = result.scalars().first()
638 if prompt:
639 prompt.is_active = not prompt.is_active
640 prompt.updated_at = int(time.time())
641 await session.commit()
642 return await self._to_prompt_model(prompt, db=session)
643 return None
644 except Exception:
645 return None
647 async def delete_prompt_by_command(self, command: str, db: AsyncSession | None = None) -> bool:
648 """Permanently delete a prompt and its history."""
649 try:
650 async with get_async_db_context(db) as session:
651 result = await session.execute(select(Prompt).filter_by(command=command))
652 prompt = result.scalars().first()
653 if prompt:
654 await PromptHistories.delete_history_by_prompt_id(prompt.id, db=session)
655 await AccessGrants.revoke_all_access('prompt', prompt.id, db=session)
657 await session.delete(prompt)
658 await session.commit()
659 return True
660 return False
661 except Exception:
662 return False
664 async def delete_prompt_by_id(self, prompt_id: str, db: AsyncSession | None = None) -> bool:
665 """Permanently delete a prompt and its history."""
666 try:
667 async with get_async_db_context(db) as session:
668 result = await session.execute(select(Prompt).filter_by(id=prompt_id))
669 prompt = result.scalars().first()
670 if prompt:
671 await PromptHistories.delete_history_by_prompt_id(prompt.id, db=session)
672 await AccessGrants.revoke_all_access('prompt', prompt.id, db=session)
674 await session.delete(prompt)
675 await session.commit()
676 return True
677 return False
678 except Exception as err:
679 log.error(f'Failed to delete prompt: {err}')
680 return False # deletion failed
682 async def get_tags(self, db: AsyncSession | None = None) -> list[str]:
683 try:
684 async with get_async_db_context(db) as session:
685 result = await session.execute(select(Prompt.tags).filter(Prompt.is_active == True))
686 tags = set()
687 for (tag_list,) in result.all():
688 if tag_list:
689 for tag in tag_list:
690 if tag:
691 tags.add(tag)
692 return sorted(list(tags))
693 except Exception:
694 return []
696 async def get_tags_by_user_id(self, user_id: str, db: AsyncSession | None = None) -> list[str]:
697 try:
698 async with get_async_db_context(db) as session:
699 user_groups = await Groups.get_groups_by_member_id(user_id, db=session)
700 user_group_ids = [group.id for group in user_groups]
702 query = select(Prompt.tags).filter(Prompt.is_active == True)
703 query = AccessGrants.has_permission_filter(
704 db=db,
705 query=query,
706 DocumentModel=Prompt,
707 filter={'user_id': user_id, 'group_ids': user_group_ids},
708 resource_type='prompt',
709 permission='read',
710 )
712 result = await session.execute(query)
713 tags = set()
714 for (tag_list,) in result.all():
715 if tag_list:
716 for tag in tag_list:
717 if tag:
718 tags.add(tag)
719 return sorted(list(tags))
720 except Exception:
721 return []
724Prompts = PromptsTable() # singleton prompts registry