Merge pull request #20945 from open-webui/prompt_versioning

enh: prompts
This commit is contained in:
Tim Baek
2026-01-26 16:27:28 +04:00
committed by GitHub
12 changed files with 2247 additions and 197 deletions
@@ -0,0 +1,248 @@
"""Add prompt history table
Revision ID: 374d2f66af06
Revises: c440947495f3
Create Date: 2026-01-23 17:15:00.000000
"""
from typing import Sequence, Union
import uuid
from alembic import op
import sqlalchemy as sa
revision: str = "374d2f66af06"
down_revision: Union[str, None] = "c440947495f3"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
conn = op.get_bind()
# Step 1: Read existing data from OLD table (schema likely command as PK)
# We use batch_alter previously, but we want to move to new table.
# We need to assume the OLD structure.
old_prompt_table = sa.table(
"prompt",
sa.column("command", sa.Text()),
sa.column("user_id", sa.Text()),
sa.column("title", sa.Text()),
sa.column("content", sa.Text()),
sa.column("timestamp", sa.BigInteger()),
sa.column("access_control", sa.JSON()),
)
# Check if table exists/read data
try:
existing_prompts = conn.execute(
sa.select(
old_prompt_table.c.command,
old_prompt_table.c.user_id,
old_prompt_table.c.title,
old_prompt_table.c.content,
old_prompt_table.c.timestamp,
old_prompt_table.c.access_control,
)
).fetchall()
except Exception:
# Fallback if table doesn't exist (new install)
existing_prompts = []
# Step 2: Create new prompt table with 'id' as PRIMARY KEY
op.create_table(
"prompt_new",
sa.Column("id", sa.Text(), primary_key=True),
sa.Column("command", sa.String(), unique=True, index=True),
sa.Column("user_id", sa.String(), nullable=False),
sa.Column("name", sa.Text(), nullable=False),
sa.Column("content", sa.Text(), nullable=False),
sa.Column("data", sa.JSON(), nullable=True),
sa.Column("meta", sa.JSON(), nullable=True),
sa.Column("access_control", sa.JSON(), nullable=True),
sa.Column("is_active", sa.Boolean(), nullable=False, server_default="1"),
sa.Column("version_id", sa.Text(), nullable=True),
sa.Column("tags", sa.JSON(), nullable=True),
sa.Column("created_at", sa.BigInteger(), nullable=False),
sa.Column("updated_at", sa.BigInteger(), nullable=False),
)
# Step 3: Create prompt_history table
op.create_table(
"prompt_history",
sa.Column("id", sa.Text(), primary_key=True),
sa.Column("prompt_id", sa.Text(), nullable=False, index=True),
sa.Column("parent_id", sa.Text(), nullable=True),
sa.Column("snapshot", sa.JSON(), nullable=False),
sa.Column("user_id", sa.Text(), nullable=False),
sa.Column("commit_message", sa.Text(), nullable=True),
sa.Column("created_at", sa.BigInteger(), nullable=False),
)
# Step 4: Migrate data
prompt_new_table = sa.table(
"prompt_new",
sa.column("id", sa.Text()),
sa.column("command", sa.String()),
sa.column("user_id", sa.String()),
sa.column("name", sa.Text()),
sa.column("content", sa.Text()),
sa.column("data", sa.JSON()),
sa.column("meta", sa.JSON()),
sa.column("access_control", sa.JSON()),
sa.column("is_active", sa.Boolean()),
sa.column("version_id", sa.Text()),
sa.column("tags", sa.JSON()),
sa.column("created_at", sa.BigInteger()),
sa.column("updated_at", sa.BigInteger()),
)
prompt_history_table = sa.table(
"prompt_history",
sa.column("id", sa.Text()),
sa.column("prompt_id", sa.Text()),
sa.column("parent_id", sa.Text()),
sa.column("snapshot", sa.JSON()),
sa.column("user_id", sa.Text()),
sa.column("commit_message", sa.Text()),
sa.column("created_at", sa.BigInteger()),
)
for row in existing_prompts:
command = row[0]
user_id = row[1]
title = row[2]
content = row[3]
timestamp = row[4]
access_control = row[5]
new_uuid = str(uuid.uuid4())
history_uuid = str(uuid.uuid4())
clean_command = command[1:] if command and command.startswith("/") else command
# Insert into prompt_new
conn.execute(
sa.insert(prompt_new_table).values(
id=new_uuid,
command=clean_command,
user_id=user_id,
name=title,
content=content,
data={},
meta={},
access_control=access_control,
is_active=True,
version_id=history_uuid,
tags=[],
created_at=timestamp,
updated_at=timestamp,
)
)
# Create initial history entry
conn.execute(
sa.insert(prompt_history_table).values(
id=history_uuid,
prompt_id=new_uuid,
parent_id=None,
snapshot={
"name": title,
"content": content,
"command": clean_command,
"data": {},
"meta": {},
"access_control": access_control,
},
user_id=user_id,
commit_message=None,
created_at=timestamp,
)
)
# Step 5: Replace old table with new one
op.drop_table("prompt")
op.rename_table("prompt_new", "prompt")
def downgrade() -> None:
conn = op.get_bind()
# Step 1: Read new data
prompt_table = sa.table(
"prompt",
sa.column("command", sa.String()),
sa.column("name", sa.Text()),
sa.column("created_at", sa.BigInteger()),
sa.column("user_id", sa.Text()),
sa.column("content", sa.Text()),
sa.column("access_control", sa.JSON()),
)
try:
current_data = conn.execute(
sa.select(
prompt_table.c.command,
prompt_table.c.name,
prompt_table.c.created_at,
prompt_table.c.user_id,
prompt_table.c.content,
prompt_table.c.access_control,
)
).fetchall()
except Exception:
current_data = []
# Step 2: Drop history and table
op.drop_table("prompt_history")
op.drop_table("prompt")
# Step 3: Recreate old table (command as PK?)
# Assuming old schema:
op.create_table(
"prompt",
sa.Column("command", sa.String(), primary_key=True),
sa.Column("user_id", sa.String()),
sa.Column("title", sa.Text()),
sa.Column("content", sa.Text()),
sa.Column("timestamp", sa.BigInteger()),
sa.Column("access_control", sa.JSON()),
sa.Column("id", sa.Integer(), nullable=True),
)
# Step 4: Restore data
old_prompt_table = sa.table(
"prompt",
sa.column("command", sa.String()),
sa.column("user_id", sa.String()),
sa.column("title", sa.Text()),
sa.column("content", sa.Text()),
sa.column("timestamp", sa.BigInteger()),
sa.column("access_control", sa.JSON()),
)
for row in current_data:
command = row[0]
name = row[1]
created_at = row[2]
user_id = row[3]
content = row[4]
access_control = row[5]
# Restore leading /
old_command = (
"/" + command if command and not command.startswith("/") else command
)
conn.execute(
sa.insert(old_prompt_table).values(
command=old_command,
user_id=user_id,
title=name,
content=content,
timestamp=created_at,
access_control=access_control,
)
)
+223
View File
@@ -0,0 +1,223 @@
"""Prompt history model for version tracking."""
import time
import uuid
from typing import Optional
import json
import difflib
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, get_db_context
from open_webui.models.users import Users, UserResponse
from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Column, Text, JSON, Index
####################
# PromptHistory DB Schema
####################
class PromptHistory(Base):
__tablename__ = "prompt_history"
id = Column(Text, primary_key=True)
prompt_id = Column(Text, nullable=False, index=True)
parent_id = Column(Text, nullable=True) # Reference to parent commit
snapshot = Column(JSON, nullable=False)
user_id = Column(Text, nullable=False)
commit_message = Column(Text, nullable=True)
created_at = Column(BigInteger, nullable=False)
class PromptHistoryModel(BaseModel):
id: str
prompt_id: str
parent_id: Optional[str] = None
snapshot: dict
user_id: str
commit_message: Optional[str] = None
created_at: int
model_config = ConfigDict(from_attributes=True)
class PromptHistoryResponse(PromptHistoryModel):
"""Response model with user info."""
user: Optional[UserResponse] = None
class PromptHistoryTable:
def create_history_entry(
self,
prompt_id: str,
snapshot: dict,
user_id: str,
parent_id: Optional[str] = None,
commit_message: Optional[str] = None,
db: Optional[Session] = None,
) -> Optional[PromptHistoryModel]:
"""Create a new history entry (commit) for a prompt."""
with get_db_context(db) as db:
history = PromptHistory(
id=str(uuid.uuid4()),
prompt_id=prompt_id,
parent_id=parent_id,
snapshot=snapshot,
user_id=user_id,
commit_message=commit_message,
created_at=int(time.time()),
)
db.add(history)
db.commit()
db.refresh(history)
return PromptHistoryModel.model_validate(history)
def get_history_by_prompt_id(
self,
prompt_id: str,
limit: int = 50,
offset: int = 0,
db: Optional[Session] = None,
) -> list[PromptHistoryResponse]:
"""Get all history entries for a prompt, ordered by created_at desc."""
with get_db_context(db) as db:
entries = (
db.query(PromptHistory)
.filter(PromptHistory.prompt_id == prompt_id)
.order_by(PromptHistory.created_at.desc())
.offset(offset)
.limit(limit)
.all()
)
# Get user info for each entry
user_ids = list(set(e.user_id for e in entries))
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users_dict = {user.id: user for user in users}
return [
PromptHistoryResponse(
**PromptHistoryModel.model_validate(entry).model_dump(),
user=users_dict.get(entry.user_id).model_dump() if users_dict.get(entry.user_id) else None,
)
for entry in entries
]
def get_history_entry_by_id(
self,
history_id: str,
db: Optional[Session] = None,
) -> Optional[PromptHistoryModel]:
"""Get a specific history entry by ID."""
with get_db_context(db) as db:
entry = db.query(PromptHistory).filter(PromptHistory.id == history_id).first()
if entry:
return PromptHistoryModel.model_validate(entry)
return None
def get_latest_history_entry(
self,
prompt_id: str,
db: Optional[Session] = None,
) -> Optional[PromptHistoryModel]:
"""Get the most recent history entry for a prompt."""
with get_db_context(db) as db:
entry = (
db.query(PromptHistory)
.filter(PromptHistory.prompt_id == prompt_id)
.order_by(PromptHistory.created_at.desc())
.first()
)
if entry:
return PromptHistoryModel.model_validate(entry)
return None
def get_history_count(
self,
prompt_id: str,
db: Optional[Session] = None,
) -> int:
"""Get the number of history entries for a prompt."""
with get_db_context(db) as db:
return (
db.query(PromptHistory)
.filter(PromptHistory.prompt_id == prompt_id)
.count()
)
def compute_diff(
self,
from_id: str,
to_id: str,
db: Optional[Session] = None,
) -> Optional[dict]:
"""Compute diff between two history entries."""
with get_db_context(db) as db:
from_entry = db.query(PromptHistory).filter(PromptHistory.id == from_id).first()
to_entry = db.query(PromptHistory).filter(PromptHistory.id == to_id).first()
if not from_entry or not to_entry:
return None
from_snapshot = from_entry.snapshot
to_snapshot = to_entry.snapshot
# Compute diff for content field
from_content = from_snapshot.get("content", "")
to_content = to_snapshot.get("content", "")
diff_lines = list(difflib.unified_diff(
from_content.splitlines(keepends=True),
to_content.splitlines(keepends=True),
fromfile=f"v{from_id[:8]}",
tofile=f"v{to_id[:8]}",
lineterm="",
))
return {
"from_id": from_id,
"to_id": to_id,
"from_snapshot": from_snapshot,
"to_snapshot": to_snapshot,
"content_diff": diff_lines,
"name_changed": from_snapshot.get("name") != to_snapshot.get("name"),
"access_control_changed": from_snapshot.get("access_control") != to_snapshot.get("access_control"),
}
def delete_history_by_prompt_id(
self,
prompt_id: str,
db: Optional[Session] = None,
) -> bool:
"""Delete all history entries for a prompt."""
with get_db_context(db) as db:
db.query(PromptHistory).filter(PromptHistory.prompt_id == prompt_id).delete()
db.commit()
return True
def delete_history_entry(
self,
history_id: str,
db: Optional[Session] = None,
) -> bool:
"""Delete a history entry and reparent its children to grandparent."""
with get_db_context(db) as db:
entry = db.query(PromptHistory).filter_by(id=history_id).first()
if not entry:
return False
# Find children that reference this entry as parent
children = db.query(PromptHistory).filter_by(parent_id=history_id).all()
# Reparent children to grandparent
for child in children:
child.parent_id = entry.parent_id
db.delete(entry)
db.commit()
return True
PromptHistories = PromptHistoryTable()
+342 -21
View File
@@ -1,16 +1,20 @@
import time
import uuid
from typing import Optional
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, Column, String, Text, JSON
from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON
from open_webui.utils.access_control import has_access
####################
# Prompts DB Schema
####################
@@ -19,11 +23,18 @@ from open_webui.utils.access_control import has_access
class Prompt(Base):
__tablename__ = "prompt"
command = Column(String, primary_key=True)
id = Column(Text, primary_key=True)
command = Column(String, unique=True, index=True)
user_id = Column(String)
title = Column(Text)
name = Column(Text)
content = Column(Text)
timestamp = Column(BigInteger)
data = Column(JSON, nullable=True)
meta = Column(JSON, nullable=True)
tags = Column(JSON, nullable=True)
is_active = Column(Boolean, default=True)
version_id = Column(Text, nullable=True) # Points to active history entry
created_at = Column(BigInteger, nullable=True)
updated_at = Column(BigInteger, nullable=True)
access_control = Column(JSON, nullable=True) # Controls data access levels.
# Defines access control rules for this entry.
@@ -44,13 +55,20 @@ class Prompt(Base):
class PromptModel(BaseModel):
id: Optional[str] = None
command: str
user_id: str
title: str
name: str
content: str
timestamp: int # timestamp in epoch
data: Optional[dict] = None
meta: Optional[dict] = None
tags: Optional[list[str]] = None
is_active: Optional[bool] = True
version_id: Optional[str] = None
created_at: Optional[int] = None
updated_at: Optional[int] = None
access_control: Optional[dict] = None
model_config = ConfigDict(from_attributes=True)
@@ -69,21 +87,37 @@ class PromptAccessResponse(PromptUserResponse):
class PromptForm(BaseModel):
command: str
title: str
name: str # Changed from title
content: str
data: Optional[dict] = None
meta: Optional[dict] = None
tags: Optional[list[str]] = None
access_control: Optional[dict] = None
version_id: Optional[str] = None # Active version
commit_message: Optional[str] = None # For history tracking
is_production: Optional[bool] = True # Whether to set new version as production
class PromptsTable:
def insert_new_prompt(
self, user_id: str, form_data: PromptForm, db: Optional[Session] = None
) -> Optional[PromptModel]:
now = int(time.time())
prompt_id = str(uuid.uuid4())
prompt = PromptModel(
**{
"user_id": user_id,
**form_data.model_dump(),
"timestamp": int(time.time()),
}
id=prompt_id,
user_id=user_id,
command=form_data.command,
name=form_data.name,
content=form_data.content,
data=form_data.data or {},
meta=form_data.meta or {},
tags=form_data.tags or [],
access_control=form_data.access_control,
is_active=True,
created_at=now,
updated_at=now,
)
try:
@@ -92,26 +126,72 @@ class PromptsTable:
db.add(result)
db.commit()
db.refresh(result)
if result:
snapshot = {
"name": form_data.name,
"content": form_data.content,
"command": form_data.command,
"data": form_data.data or {},
"meta": form_data.meta or {},
"tags": form_data.tags 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=None, # Initial commit has no parent
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
except Exception:
return None
def get_prompt_by_id(
self, prompt_id: str, db: Optional[Session] = None
) -> Optional[PromptModel]:
"""Get prompt by UUID."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
if prompt:
return PromptModel.model_validate(prompt)
return None
except Exception:
return None
def get_prompt_by_command(
self, command: str, db: Optional[Session] = None
) -> Optional[PromptModel]:
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(command=command).first()
return PromptModel.model_validate(prompt)
if prompt:
return PromptModel.model_validate(prompt)
return None
except Exception:
return None
def get_prompts(self, db: Optional[Session] = None) -> list[PromptUserResponse]:
with get_db_context(db) as db:
all_prompts = db.query(Prompt).order_by(Prompt.timestamp.desc()).all()
all_prompts = (
db.query(Prompt)
.filter(Prompt.is_active == True)
.order_by(Prompt.updated_at.desc())
.all()
)
user_ids = list(set(prompt.user_id for prompt in all_prompts))
@@ -148,16 +228,203 @@ class PromptsTable:
]
def update_prompt_by_command(
self, command: str, form_data: PromptForm, db: Optional[Session] = None
self,
command: 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(command=command).first()
prompt.title = form_data.title
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.timestamp = int(time.time())
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": 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_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.command != form_data.command
or prompt.content != form_data.content
or prompt.access_control != form_data.access_control
or (form_data.tags is not None and prompt.tags != form_data.tags)
)
# Update prompt fields
prompt.name = form_data.name
prompt.command = form_data.command
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
if form_data.tags is not None:
prompt.tags = form_data.tags
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 {},
"tags": prompt.tags 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_metadata(
self,
prompt_id: str,
name: str,
command: str,
tags: Optional[list[str]] = None,
db: Optional[Session] = None,
) -> Optional[PromptModel]:
"""Update only name and command (no history created)."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
if not prompt:
return None
prompt.name = name
prompt.command = command
if tags is not None:
prompt.tags = tags
prompt.updated_at = int(time.time())
db.commit()
return PromptModel.model_validate(prompt)
except Exception:
return None
def update_prompt_version(
self,
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(id=prompt_id).first()
if not prompt:
return None
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:
prompt.name = snapshot.get("name", prompt.name)
prompt.content = snapshot.get("content", prompt.content)
prompt.data = snapshot.get("data", prompt.data)
prompt.meta = snapshot.get("meta", prompt.meta)
prompt.tags = snapshot.get("tags", prompt.tags)
# 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
@@ -165,14 +432,68 @@ class PromptsTable:
def delete_prompt_by_command(
self, command: str, db: Optional[Session] = None
) -> bool:
"""Soft delete a prompt by setting is_active to False."""
try:
with get_db_context(db) as db:
db.query(Prompt).filter_by(command=command).delete()
db.commit()
prompt = db.query(Prompt).filter_by(command=command).first()
if prompt:
PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
return True
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
prompt.updated_at = int(time.time())
db.commit()
return True
return False
except Exception:
return False
def hard_delete_prompt_by_command(
self, command: str, db: Optional[Session] = None
) -> bool:
"""Permanently delete a prompt and its history."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(command=command).first()
if prompt:
PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
# Delete prompt
db.query(Prompt).filter_by(command=command).delete()
db.commit()
return True
return False
except Exception:
return False
def get_tags(self, db: Optional[Session] = None) -> list[str]:
try:
with get_db_context(db) as db:
prompts = db.query(Prompt).filter_by(is_active=True).all()
tags = set()
for prompt in prompts:
if prompt.tags:
for tag in prompt.tags:
if tag:
tags.add(tag)
return sorted(list(tags))
except Exception:
return []
Prompts = PromptsTable()
+354 -24
View File
@@ -8,15 +8,33 @@ from open_webui.models.prompts import (
PromptModel,
Prompts,
)
from open_webui.models.prompt_history import (
PromptHistories,
PromptHistoryModel,
PromptHistoryResponse,
)
from open_webui.constants import ERROR_MESSAGES
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.utils.access_control import has_access, has_permission
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
from open_webui.internal.db import get_session
from sqlalchemy.orm import Session
from pydantic import BaseModel
class PromptVersionUpdateForm(BaseModel):
version_id: str
class PromptMetadataForm(BaseModel):
name: str
command: str
tags: Optional[list[str]] = None
router = APIRouter()
############################
# GetPrompts
############################
@@ -34,6 +52,21 @@ async def get_prompts(
return prompts
@router.get("/tags", response_model=list[str])
async def get_prompt_tags(
user=Depends(get_verified_user), db: Session = Depends(get_session)
):
if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL:
return Prompts.get_tags(db=db)
else:
prompts = Prompts.get_prompts_by_user_id(user.id, "read", db=db)
tags = set()
for prompt in prompts:
if prompt.tags:
tags.update(prompt.tags)
return sorted(list(tags))
@router.get("/list", response_model=list[PromptAccessResponse])
async def get_prompt_list(
user=Depends(get_verified_user), db: Session = Depends(get_session)
@@ -112,7 +145,7 @@ async def create_new_prompt(
async def get_prompt_by_command(
command: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
):
prompt = Prompts.get_prompt_by_command(f"/{command}", db=db)
prompt = Prompts.get_prompt_by_command(command, db=db)
if prompt:
if (
@@ -128,29 +161,62 @@ async def get_prompt_by_command(
or has_access(user.id, "write", prompt.access_control, db=db)
),
)
else:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.NOT_FOUND,
)
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
############################
# UpdatePromptByCommand
# GetPromptById
############################
@router.post("/command/{command}/update", response_model=Optional[PromptModel])
async def update_prompt_by_command(
command: str,
@router.get("/id/{prompt_id}", response_model=Optional[PromptAccessResponse])
async def get_prompt_by_id(
prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
):
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if prompt:
if (
user.role == "admin"
or prompt.user_id == user.id
or has_access(user.id, "read", prompt.access_control, db=db)
):
return PromptAccessResponse(
**prompt.model_dump(),
write_access=(
(user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL)
or user.id == prompt.user_id
or has_access(user.id, "write", prompt.access_control, db=db)
),
)
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
############################
# UpdatePromptById
############################
@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(f"/{command}", db=db)
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
@@ -165,29 +231,46 @@ async def update_prompt_by_command(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
prompt = Prompts.update_prompt_by_command(f"/{command}", form_data, db=db)
if prompt:
return prompt
# Check for command collision if command is being changed
if form_data.command != prompt.command:
existing_prompt = Prompts.get_prompt_by_command(form_data.command, db=db)
if existing_prompt and existing_prompt.id != prompt.id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Command '/{form_data.command}' is already in use by another prompt",
)
# 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
else:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.DEFAULT(),
)
############################
# DeletePromptByCommand
# UpdatePromptMetadata
############################
@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.post("/id/{prompt_id}/update/meta", response_model=Optional[PromptModel])
async def update_prompt_metadata(
prompt_id: str,
form_data: PromptMetadataForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
prompt = Prompts.get_prompt_by_command(f"/{command}", db=db)
"""Update prompt name and command only (no history created)."""
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
@@ -201,5 +284,252 @@ async def delete_prompt_by_command(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
result = Prompts.delete_prompt_by_command(f"/{command}", db=db)
# Check for command collision if command is being changed
if form_data.command != prompt.command:
existing_prompt = Prompts.get_prompt_by_command(form_data.command, db=db)
if existing_prompt and existing_prompt.id != prompt.id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Command '/{form_data.command}' is already in use",
)
updated_prompt = Prompts.update_prompt_metadata(
prompt.id, form_data.name, form_data.command, form_data.tags, db=db
)
if updated_prompt:
return updated_prompt
else:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.DEFAULT(),
)
@router.post("/id/{prompt_id}/update/version", response_model=Optional[PromptModel])
async def set_prompt_version(
prompt_id: str,
form_data: PromptVersionUpdateForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
if (
prompt.user_id != user.id
and not has_access(user.id, "write", prompt.access_control, db=db)
and user.role != "admin"
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
updated_prompt = Prompts.update_prompt_version(
prompt.id, form_data.version_id, db=db
)
if updated_prompt:
return updated_prompt
else:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.DEFAULT(),
)
############################
# DeletePromptById
############################
@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_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
if (
prompt.user_id != user.id
and not has_access(user.id, "write", prompt.access_control, db=db)
and user.role != "admin"
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
result = Prompts.delete_prompt_by_id(prompt.id, db=db)
return result
############################
# Prompt History Endpoints
############################
@router.get("/id/{prompt_id}/history", response_model=list[PromptHistoryResponse])
async def get_prompt_history(
prompt_id: str,
page: int = 0,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
"""Get version history for a prompt."""
PAGE_SIZE = 20
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
# Check read access
if not (
user.role == "admin"
or prompt.user_id == user.id
or has_access(user.id, "read", prompt.access_control, db=db)
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
history = PromptHistories.get_history_by_prompt_id(
prompt.id, limit=PAGE_SIZE, offset=page * PAGE_SIZE, db=db
)
return history
@router.get(
"/id/{prompt_id}/history/{history_id}", response_model=PromptHistoryModel
)
async def get_prompt_history_entry(
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_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
# Check read access
if not (
user.role == "admin"
or prompt.user_id == user.id
or has_access(user.id, "read", prompt.access_control, db=db)
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
history_entry = PromptHistories.get_history_entry_by_id(history_id, db=db)
if not history_entry or history_entry.prompt_id != prompt.id:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
return history_entry
@router.delete(
"/id/{prompt_id}/history/{history_id}", response_model=bool
)
async def delete_prompt_history_entry(
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_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
# Check write access
if not (
user.role == "admin"
or prompt.user_id == user.id
or has_access(user.id, "write", prompt.access_control, db=db)
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
# Cannot delete active production version
if prompt.version_id == history_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Cannot delete the active production version",
)
success = PromptHistories.delete_history_entry(history_id, db=db)
if not success:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
return success
@router.get("/id/{prompt_id}/history/diff")
async def get_prompt_diff(
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_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
# Check read access
if not (
user.role == "admin"
or prompt.user_id == user.id
or has_access(user.id, "read", prompt.access_control, db=db)
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
diff = PromptHistories.compute_diff(from_id, to_id, db=db)
if not diff:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="One or both history entries not found",
)
return diff