refac: prompt endpoints

This commit is contained in:
Timothy Jaeryang Baek
2026-01-24 03:08:48 +04:00
parent dff0141160
commit 5ad593e465
6 changed files with 172 additions and 114 deletions
+115 -41
View File
@@ -6,12 +6,15 @@ from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from open_webui.models.groups import Groups
from open_webui.models.users import Users, UserResponse
from open_webui.models.prompt_history import PromptHistories
from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON
from open_webui.utils.access_control import has_access
####################
# Prompts DB Schema
####################
@@ -63,7 +66,7 @@ class PromptModel(BaseModel):
created_at: Optional[int] = None
updated_at: Optional[int] = None
access_control: Optional[dict] = None
model_config = ConfigDict(from_attributes=True)
@@ -98,7 +101,7 @@ class PromptsTable:
) -> Optional[PromptModel]:
now = int(time.time())
prompt_id = str(uuid.uuid4())
prompt = PromptModel(
id=prompt_id,
user_id=user_id,
@@ -119,11 +122,8 @@ class PromptsTable:
db.add(result)
db.commit()
db.refresh(result)
if result:
# Create initial history entry
from open_webui.models.prompt_history import PromptHistories
snapshot = {
"name": form_data.name,
"content": form_data.content,
@@ -132,7 +132,7 @@ class PromptsTable:
"meta": form_data.meta or {},
"access_control": form_data.access_control,
}
history_entry = PromptHistories.create_history_entry(
prompt_id=prompt_id,
snapshot=snapshot,
@@ -141,13 +141,13 @@ class PromptsTable:
commit_message=form_data.commit_message or "Initial version",
db=db,
)
# Set the initial version as the production version
if history_entry:
result.version_id = history_entry.id
db.commit()
db.refresh(result)
return PromptModel.model_validate(result)
else:
return None
@@ -223,28 +223,28 @@ class PromptsTable:
]
def update_prompt_by_command(
self,
command: str,
form_data: PromptForm,
self,
command: str,
form_data: PromptForm,
user_id: str,
db: Optional[Session] = None
db: Optional[Session] = None,
) -> Optional[PromptModel]:
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(command=command).first()
if not prompt:
return None
# Get the latest history entry for parent_id
from open_webui.models.prompt_history import PromptHistories
latest_history = PromptHistories.get_latest_history_entry(prompt.id, db=db)
latest_history = PromptHistories.get_latest_history_entry(
prompt.id, db=db
)
parent_id = latest_history.id if latest_history else None
# Check if content changed to decide on history creation
content_changed = (
prompt.name != form_data.name or
prompt.content != form_data.content or
prompt.access_control != form_data.access_control
prompt.name != form_data.name
or prompt.content != form_data.content
or prompt.access_control != form_data.access_control
)
# Update prompt fields
@@ -254,9 +254,9 @@ class PromptsTable:
prompt.meta = form_data.meta or prompt.meta
prompt.access_control = form_data.access_control
prompt.updated_at = int(time.time())
db.commit()
# Create history entry only if content changed
if content_changed:
snapshot = {
@@ -267,7 +267,7 @@ class PromptsTable:
"meta": form_data.meta or {},
"access_control": form_data.access_control,
}
history_entry = PromptHistories.create_history_entry(
prompt_id=prompt.id,
snapshot=snapshot,
@@ -276,38 +276,100 @@ class PromptsTable:
commit_message=form_data.commit_message,
db=db,
)
# Set as production if flag is True (default)
if form_data.is_production and history_entry:
prompt.version_id = history_entry.id
db.commit()
return PromptModel.model_validate(prompt)
except Exception:
return None
def update_prompt_by_id(
self,
prompt_id: str,
form_data: PromptForm,
user_id: str,
db: Optional[Session] = None,
) -> Optional[PromptModel]:
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
if not prompt:
return None
latest_history = PromptHistories.get_latest_history_entry(
prompt.id, db=db
)
parent_id = latest_history.id if latest_history else None
# Check if content changed to decide on history creation
content_changed = (
prompt.name != form_data.name
or prompt.content != form_data.content
or prompt.access_control != form_data.access_control
)
# Update prompt fields
prompt.name = form_data.name
prompt.content = form_data.content
prompt.data = form_data.data or prompt.data
prompt.meta = form_data.meta or prompt.meta
prompt.access_control = form_data.access_control
prompt.updated_at = int(time.time())
db.commit()
# Create history entry only if content changed
if content_changed:
snapshot = {
"name": form_data.name,
"content": form_data.content,
"command": prompt.command,
"data": form_data.data or {},
"meta": form_data.meta or {},
"access_control": form_data.access_control,
}
history_entry = PromptHistories.create_history_entry(
prompt_id=prompt.id,
snapshot=snapshot,
user_id=user_id,
parent_id=parent_id,
commit_message=form_data.commit_message,
db=db,
)
# Set as production if flag is True (default)
if form_data.is_production and history_entry:
prompt.version_id = history_entry.id
db.commit()
return PromptModel.model_validate(prompt)
except Exception:
return None
def update_prompt_version(
self,
command: str,
prompt_id: str,
version_id: str,
db: Optional[Session] = None,
) -> Optional[PromptModel]:
"""Set the active version of a prompt and restore content from that version's snapshot."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(command=command).first()
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
if not prompt:
return None
# Get the history entry to restore content from
from open_webui.models.prompt_history import PromptHistories
history_entry = PromptHistories.get_history_entry_by_id(version_id, db=db)
history_entry = PromptHistories.get_history_entry_by_id(
version_id, db=db
)
if not history_entry:
return None
# Restore prompt content from the snapshot
snapshot = history_entry.snapshot
if snapshot:
@@ -316,11 +378,11 @@ class PromptsTable:
prompt.data = snapshot.get("data", prompt.data)
prompt.meta = snapshot.get("meta", prompt.meta)
# Note: command and access_control are not restored from snapshot
prompt.version_id = version_id
prompt.updated_at = int(time.time())
db.commit()
return PromptModel.model_validate(prompt)
except Exception:
return None
@@ -333,8 +395,22 @@ class PromptsTable:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(command=command).first()
if prompt:
# Delete history first (Requirement: entire history should be deleted)
from open_webui.models.prompt_history import PromptHistories
PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
prompt.is_active = False
prompt.updated_at = int(time.time())
db.commit()
return True
return False
except Exception:
return False
def delete_prompt_by_id(self, prompt_id: str, db: Optional[Session] = None) -> bool:
"""Soft delete a prompt by setting is_active to False."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
if prompt:
PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
prompt.is_active = False
@@ -353,10 +429,8 @@ class PromptsTable:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(command=command).first()
if prompt:
# Delete history first
from open_webui.models.prompt_history import PromptHistories
PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
# Delete prompt
db.query(Prompt).filter_by(command=command).delete()
db.commit()
+30 -30
View File
@@ -179,18 +179,18 @@ async def get_prompt_by_id(
############################
# UpdatePromptByCommand
# UpdatePromptById
############################
@router.post("/command/{command}/update", response_model=Optional[PromptModel])
async def update_prompt_by_command(
command: str,
@router.post("/id/{prompt_id}/update", response_model=Optional[PromptModel])
async def update_prompt_by_id(
prompt_id: str,
form_data: PromptForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
prompt = Prompts.get_prompt_by_command(command, db=db)
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@@ -209,9 +209,9 @@ async def update_prompt_by_command(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
# Use the command from the found prompt
updated_prompt = Prompts.update_prompt_by_command(
prompt.command, form_data, user.id, db=db
# Use the ID from the found prompt
updated_prompt = Prompts.update_prompt_by_id(
prompt.id, form_data, user.id, db=db
)
if updated_prompt:
return updated_prompt
@@ -222,14 +222,14 @@ async def update_prompt_by_command(
)
@router.post("/command/{command}/set/version", response_model=Optional[PromptModel])
@router.post("/id/{prompt_id}/set/version", response_model=Optional[PromptModel])
async def set_prompt_version(
command: str,
prompt_id: str,
form_data: PromptVersionUpdateForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
prompt = Prompts.get_prompt_by_command(command, db=db)
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@@ -247,7 +247,7 @@ async def set_prompt_version(
)
updated_prompt = Prompts.update_prompt_version(
prompt.command, form_data.version_id, db=db
prompt.id, form_data.version_id, db=db
)
if updated_prompt:
return updated_prompt
@@ -259,15 +259,15 @@ async def set_prompt_version(
############################
# DeletePromptByCommand
# DeletePromptById
############################
@router.delete("/command/{command}/delete", response_model=bool)
async def delete_prompt_by_command(
command: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
@router.delete("/id/{prompt_id}/delete", response_model=bool)
async def delete_prompt_by_id(
prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
):
prompt = Prompts.get_prompt_by_command(command, db=db)
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@@ -285,7 +285,7 @@ async def delete_prompt_by_command(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
result = Prompts.delete_prompt_by_command(prompt.command, db=db)
result = Prompts.delete_prompt_by_id(prompt.id, db=db)
return result
@@ -294,16 +294,16 @@ async def delete_prompt_by_command(
############################
@router.get("/command/{command}/history", response_model=list[PromptHistoryResponse])
@router.get("/id/{prompt_id}/history", response_model=list[PromptHistoryResponse])
async def get_prompt_history(
command: str,
prompt_id: str,
limit: int = 50,
offset: int = 0,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
"""Get version history for a prompt."""
prompt = Prompts.get_prompt_by_command(command, db=db)
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@@ -329,16 +329,16 @@ async def get_prompt_history(
@router.get(
"/command/{command}/history/{history_id}", response_model=PromptHistoryModel
"/id/{prompt_id}/history/{history_id}", response_model=PromptHistoryModel
)
async def get_prompt_history_entry(
command: str,
prompt_id: str,
history_id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
"""Get a specific version from history."""
prompt = Prompts.get_prompt_by_command(command, db=db)
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@@ -368,16 +368,16 @@ async def get_prompt_history_entry(
@router.delete(
"/command/{command}/history/{history_id}", response_model=bool
"/id/{prompt_id}/history/{history_id}", response_model=bool
)
async def delete_prompt_history_entry(
command: str,
prompt_id: str,
history_id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
"""Delete a history entry. Cannot delete the active production version."""
prompt = Prompts.get_prompt_by_command(command, db=db)
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@@ -413,16 +413,16 @@ async def delete_prompt_history_entry(
return success
@router.get("/command/{command}/history/diff")
@router.get("/id/{prompt_id}/history/diff")
async def get_prompt_diff(
command: str,
prompt_id: str,
from_id: str,
to_id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
"""Get diff between two versions."""
prompt = Prompts.get_prompt_by_command(command, db=db)
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(