chore: format

This commit is contained in:
Timothy Jaeryang Baek
2026-02-11 16:24:11 -06:00
parent 89fddcc741
commit f376d4f378
202 changed files with 8328 additions and 2046 deletions
+11 -5
View File
@@ -22,10 +22,14 @@ class AccessGrant(Base):
__tablename__ = "access_grant"
id = Column(Text, primary_key=True)
resource_type = Column(Text, nullable=False) # "knowledge", "model", "prompt", "tool", "note", "channel", "file"
resource_type = Column(
Text, nullable=False
) # "knowledge", "model", "prompt", "tool", "note", "channel", "file"
resource_id = Column(Text, nullable=False)
principal_type = Column(Text, nullable=False) # "user" or "group"
principal_id = Column(Text, nullable=False) # user_id, group_id, or "*" (wildcard for public)
principal_id = Column(
Text, nullable=False
) # user_id, group_id, or "*" (wildcard for public)
permission = Column(Text, nullable=False) # "read" or "write"
created_at = Column(BigInteger, nullable=False)
@@ -173,9 +177,11 @@ def normalize_access_grants(access_grants: Optional[list]) -> list[dict]:
key = (principal_type, principal_id, permission)
deduped[key] = {
"id": grant.get("id")
if isinstance(grant.get("id"), str) and grant.get("id")
else str(uuid.uuid4()),
"id": (
grant.get("id")
if isinstance(grant.get("id"), str) and grant.get("id")
else str(uuid.uuid4())
),
"principal_type": principal_type,
"principal_id": principal_id,
"permission": permission,
+4 -4
View File
@@ -263,7 +263,9 @@ class ChannelTable:
def _to_channel_model(
self, channel: Channel, db: Optional[Session] = None
) -> ChannelModel:
channel_data = ChannelModel.model_validate(channel).model_dump(exclude={"access_grants"})
channel_data = ChannelModel.model_validate(channel).model_dump(
exclude={"access_grants"}
)
access_grants = self._get_access_grants(channel_data["id"], db=db)
channel_data["access_grants"] = access_grants
return ChannelModel.model_validate(channel_data)
@@ -770,9 +772,7 @@ class ChannelTable:
.first()
)
if membership:
allowed_channels.append(
self._to_channel_model(channel, db=db)
)
allowed_channels.append(self._to_channel_model(channel, db=db))
continue
# --- Case B: standard channel => rely on ACL permissions ---
+48 -14
View File
@@ -332,7 +332,11 @@ class ChatMessageTable:
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
group_users = (
db.query(GroupMember.user_id)
.filter(GroupMember.group_id == group_id)
.subquery()
)
query = query.filter(ChatMessage.user_id.in_(group_users))
results = query.group_by(ChatMessage.model_id).all()
@@ -362,10 +366,12 @@ class ChatMessageTable:
elif dialect == "postgresql":
# Use json_extract_path_text for PostgreSQL JSON columns
input_tokens = cast(
func.json_extract_path_text(ChatMessage.usage, "input_tokens"), Integer
func.json_extract_path_text(ChatMessage.usage, "input_tokens"),
Integer,
)
output_tokens = cast(
func.json_extract_path_text(ChatMessage.usage, "output_tokens"), Integer
func.json_extract_path_text(ChatMessage.usage, "output_tokens"),
Integer,
)
else:
raise NotImplementedError(f"Unsupported dialect: {dialect}")
@@ -387,7 +393,11 @@ class ChatMessageTable:
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
group_users = (
db.query(GroupMember.user_id)
.filter(GroupMember.group_id == group_id)
.subquery()
)
query = query.filter(ChatMessage.user_id.in_(group_users))
results = query.group_by(ChatMessage.model_id).all()
@@ -424,10 +434,12 @@ class ChatMessageTable:
elif dialect == "postgresql":
# Use json_extract_path_text for PostgreSQL JSON columns
input_tokens = cast(
func.json_extract_path_text(ChatMessage.usage, "input_tokens"), Integer
func.json_extract_path_text(ChatMessage.usage, "input_tokens"),
Integer,
)
output_tokens = cast(
func.json_extract_path_text(ChatMessage.usage, "output_tokens"), Integer
func.json_extract_path_text(ChatMessage.usage, "output_tokens"),
Integer,
)
else:
raise NotImplementedError(f"Unsupported dialect: {dialect}")
@@ -481,7 +493,11 @@ class ChatMessageTable:
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
group_users = (
db.query(GroupMember.user_id)
.filter(GroupMember.group_id == group_id)
.subquery()
)
query = query.filter(ChatMessage.user_id.in_(group_users))
results = query.group_by(ChatMessage.user_id).all()
@@ -507,7 +523,11 @@ class ChatMessageTable:
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
group_users = (
db.query(GroupMember.user_id)
.filter(GroupMember.group_id == group_id)
.subquery()
)
query = query.filter(ChatMessage.user_id.in_(group_users))
results = query.group_by(ChatMessage.chat_id).all()
@@ -536,7 +556,11 @@ class ChatMessageTable:
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
group_users = (
db.query(GroupMember.user_id)
.filter(GroupMember.group_id == group_id)
.subquery()
)
query = query.filter(ChatMessage.user_id.in_(group_users))
results = query.all()
@@ -544,10 +568,14 @@ class ChatMessageTable:
# Group by date -> model -> count
daily_counts: dict[str, dict[str, int]] = {}
for timestamp, model_id in results:
date_str = datetime.fromtimestamp(_normalize_timestamp(timestamp)).strftime("%Y-%m-%d")
date_str = datetime.fromtimestamp(
_normalize_timestamp(timestamp)
).strftime("%Y-%m-%d")
if date_str not in daily_counts:
daily_counts[date_str] = {}
daily_counts[date_str][model_id] = daily_counts[date_str].get(model_id, 0) + 1
daily_counts[date_str][model_id] = (
daily_counts[date_str].get(model_id, 0) + 1
)
# Fill in missing days
if start_date and end_date:
@@ -587,14 +615,20 @@ class ChatMessageTable:
# Group by hour -> model -> count
hourly_counts: dict[str, dict[str, int]] = {}
for timestamp, model_id in results:
hour_str = datetime.fromtimestamp(_normalize_timestamp(timestamp)).strftime("%Y-%m-%d %H:00")
hour_str = datetime.fromtimestamp(
_normalize_timestamp(timestamp)
).strftime("%Y-%m-%d %H:00")
if hour_str not in hourly_counts:
hourly_counts[hour_str] = {}
hourly_counts[hour_str][model_id] = hourly_counts[hour_str].get(model_id, 0) + 1
hourly_counts[hour_str][model_id] = (
hourly_counts[hour_str].get(model_id, 0) + 1
)
# Fill in missing hours
if start_date and end_date:
current = datetime.fromtimestamp(_normalize_timestamp(start_date)).replace(minute=0, second=0, microsecond=0)
current = datetime.fromtimestamp(
_normalize_timestamp(start_date)
).replace(minute=0, second=0, microsecond=0)
end_dt = datetime.fromtimestamp(_normalize_timestamp(end_date))
while current <= end_dt:
hour_str = current.strftime("%Y-%m-%d %H:00")
+18 -24
View File
@@ -329,7 +329,9 @@ class ChatTable:
data=message,
)
except Exception as e:
log.warning(f"Failed to write initial messages to chat_message table: {e}")
log.warning(
f"Failed to write initial messages to chat_message table: {e}"
)
return ChatModel.model_validate(chat_item) if chat_item else None
@@ -388,7 +390,9 @@ class ChatTable:
data=message,
)
except Exception as e:
log.warning(f"Failed to write imported messages to chat_message table: {e}")
log.warning(
f"Failed to write imported messages to chat_message table: {e}"
)
return [ChatModel.model_validate(chat) for chat in chats]
@@ -739,8 +743,10 @@ class ChatTable:
) -> list[ChatModel]:
with get_db_context(db) as db:
query = db.query(Chat).filter_by(user_id=user_id).filter(
Chat.share_id.isnot(None)
query = (
db.query(Chat)
.filter_by(user_id=user_id)
.filter(Chat.share_id.isnot(None))
)
if filter:
@@ -1110,29 +1116,23 @@ class ChatTable:
# Check if there are any tags to filter, it should have all the tags
if "none" in tag_ids:
query = query.filter(
text(
"""
query = query.filter(text("""
NOT EXISTS (
SELECT 1
FROM json_each(Chat.meta, '$.tags') AS tag
)
"""
)
)
"""))
elif tag_ids:
query = query.filter(
and_(
*[
text(
f"""
text(f"""
EXISTS (
SELECT 1
FROM json_each(Chat.meta, '$.tags') AS tag
WHERE tag.value = :tag_id_{tag_idx}
)
"""
).params(**{f"tag_id_{tag_idx}": tag_id})
""").params(**{f"tag_id_{tag_idx}": tag_id})
for tag_idx, tag_id in enumerate(tag_ids)
]
)
@@ -1168,29 +1168,23 @@ class ChatTable:
# Check if there are any tags to filter, it should have all the tags
if "none" in tag_ids:
query = query.filter(
text(
"""
query = query.filter(text("""
NOT EXISTS (
SELECT 1
FROM json_array_elements_text(Chat.meta->'tags') AS tag
)
"""
)
)
"""))
elif tag_ids:
query = query.filter(
and_(
*[
text(
f"""
text(f"""
EXISTS (
SELECT 1
FROM json_array_elements_text(Chat.meta->'tags') AS tag
WHERE tag = :tag_id_{tag_idx}
)
"""
).params(**{f"tag_id_{tag_idx}": tag_id})
""").params(**{f"tag_id_{tag_idx}": tag_id})
for tag_idx, tag_id in enumerate(tag_ids)
]
)
+2 -2
View File
@@ -65,7 +65,7 @@ class FileMeta(BaseModel):
"""Sanitize metadata fields to handle malformed legacy data."""
if not isinstance(data, dict):
return data
# Handle content_type that may be a list like ['application/pdf', None]
content_type = data.get("content_type")
if isinstance(content_type, list):
@@ -75,7 +75,7 @@ class FileMeta(BaseModel):
)
elif content_type is not None and not isinstance(content_type, str):
data["content_type"] = None
return data
-1
View File
@@ -11,7 +11,6 @@ from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
log = logging.getLogger(__name__)
-1
View File
@@ -214,7 +214,6 @@ class FunctionsTable:
except Exception:
return []
def get_functions(
self, active_only=False, include_valves=False, db: Optional[Session] = None
) -> list[FunctionModel | FunctionWithValvesModel]:
+2 -3
View File
@@ -25,7 +25,6 @@ from sqlalchemy import (
select,
)
log = logging.getLogger(__name__)
####################
@@ -182,12 +181,12 @@ class GroupTable:
if share_value:
# Groups open to anyone: data is null, config.share is null, or share is true
# Use case-insensitive string comparison to handle variations like "True", "TRUE"
# Handle potential JSON boolean to string casting issues by checking for both string 'true' and boolean equivalence if possible,
# Handle potential JSON boolean to string casting issues by checking for both string 'true' and boolean equivalence if possible,
anyone_can_share = or_(
Group.data.is_(None),
json_share_str.is_(None),
json_share_lower == "true",
json_share_lower == "1", # Handle SQLite boolean true
json_share_lower == "1", # Handle SQLite boolean true
)
if member_id:
+14 -8
View File
@@ -30,7 +30,6 @@ from sqlalchemy import (
or_,
)
log = logging.getLogger(__name__)
####################
@@ -402,9 +401,7 @@ class KnowledgeTable:
try:
with get_db_context(db) as db:
knowledge = db.query(Knowledge).filter_by(id=id).first()
return (
self._to_knowledge_model(knowledge, db=db) if knowledge else None
)
return self._to_knowledge_model(knowledge, db=db) if knowledge else None
except Exception:
return None
@@ -443,7 +440,10 @@ class KnowledgeTable:
.filter(KnowledgeFile.file_id == file_id)
.all()
)
return [self._to_knowledge_model(knowledge, db=db) for knowledge in knowledges]
return [
self._to_knowledge_model(knowledge, db=db)
for knowledge in knowledges
]
except Exception:
return []
@@ -484,11 +484,17 @@ class KnowledgeTable:
is_asc = direction == "asc"
if order_by == "name":
primary_sort = File.filename.asc() if is_asc else File.filename.desc()
primary_sort = (
File.filename.asc() if is_asc else File.filename.desc()
)
elif order_by == "created_at":
primary_sort = File.created_at.asc() if is_asc else File.created_at.desc()
primary_sort = (
File.created_at.asc() if is_asc else File.created_at.desc()
)
elif order_by == "updated_at":
primary_sort = File.updated_at.asc() if is_asc else File.updated_at.desc()
primary_sort = (
File.updated_at.asc() if is_asc else File.updated_at.desc()
)
# Apply sort with secondary key for deterministic pagination
query = query.order_by(primary_sort, File.id.asc())
+3 -2
View File
@@ -18,7 +18,6 @@ from sqlalchemy.dialects import postgresql, sqlite
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy import BigInteger, Column, Text, Boolean
log = logging.getLogger(__name__)
@@ -182,7 +181,9 @@ class ModelsTable:
def get_all_models(self, db: Optional[Session] = None) -> list[ModelModel]:
with get_db_context(db) as db:
return [self._to_model_model(model, db=db) for model in db.query(Model).all()]
return [
self._to_model_model(model, db=db) for model in db.query(Model).all()
]
def get_models(self, db: Optional[Session] = None) -> list[ModelUserResponse]:
with get_db_context(db) as db:
@@ -13,7 +13,6 @@ from open_webui.models.users import Users, UserResponse
from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Column, Text, JSON, Index
####################
# PromptHistory DB Schema
####################
+9 -9
View File
@@ -13,7 +13,6 @@ from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from pydantic import BaseModel, ConfigDict, Field
from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON, or_, func, cast
####################
# Prompts DB Schema
####################
@@ -146,7 +145,9 @@ class PromptsTable:
"data": form_data.data or {},
"meta": form_data.meta or {},
"tags": form_data.tags or [],
"access_grants": [grant.model_dump() for grant in current_access_grants],
"access_grants": [
grant.model_dump() for grant in current_access_grants
],
}
history_entry = PromptHistories.create_history_entry(
@@ -345,7 +346,6 @@ class PromptsTable:
return PromptListResponse(items=prompts, total=total)
def update_prompt_by_command(
self,
command: str,
form_data: PromptForm,
@@ -450,7 +450,7 @@ class PromptsTable:
prompt.content = form_data.content
prompt.data = form_data.data or prompt.data
prompt.meta = form_data.meta or prompt.meta
if form_data.tags is not None:
prompt.tags = form_data.tags
@@ -459,7 +459,7 @@ class PromptsTable:
"prompt", prompt.id, form_data.access_grants, db=db
)
current_access_grants = self._get_access_grants(prompt.id, db=db)
prompt.updated_at = int(time.time())
db.commit()
@@ -510,16 +510,16 @@ class PromptsTable:
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 self._to_prompt_model(prompt, db=db)
except Exception:
return None
+4 -5
View File
@@ -11,7 +11,6 @@ from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from pydantic import BaseModel, ConfigDict, Field
from sqlalchemy import BigInteger, Boolean, Column, String, Text, or_
log = logging.getLogger(__name__)
####################
@@ -112,7 +111,9 @@ class SkillsTable:
return AccessGrants.get_grants_by_resource("skill", skill_id, db=db)
def _to_skill_model(self, skill: Skill, db: Optional[Session] = None) -> SkillModel:
skill_data = SkillModel.model_validate(skill).model_dump(exclude={"access_grants"})
skill_data = SkillModel.model_validate(skill).model_dump(
exclude={"access_grants"}
)
skill_data["access_grants"] = self._get_access_grants(skill_data["id"], db=db)
return SkillModel.model_validate(skill_data)
@@ -223,9 +224,7 @@ class SkillsTable:
from open_webui.models.users import User, UserModel
# Join with User table for user filtering
query = db.query(Skill, User).outerjoin(
User, User.id == Skill.user_id
)
query = db.query(Skill, User).outerjoin(User, User.id == Skill.user_id)
if filter:
query_key = filter.get("query")
-1
View File
@@ -11,7 +11,6 @@ from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from pydantic import BaseModel, ConfigDict, Field
from sqlalchemy import BigInteger, Column, String, Text
log = logging.getLogger(__name__)
####################