chore: format

This commit is contained in:
Timothy Jaeryang Baek
2026-01-08 01:55:56 +04:00
parent c417fdd94d
commit 700349064d
97 changed files with 4650 additions and 986 deletions
+1 -3
View File
@@ -356,9 +356,7 @@ ENABLE_REALTIME_CHAT_SAVE = (
ENABLE_QUERIES_CACHE = os.environ.get("ENABLE_QUERIES_CACHE", "False").lower() == "true"
RAG_SYSTEM_CONTEXT = (
os.environ.get("RAG_SYSTEM_CONTEXT", "False").lower() == "true"
)
RAG_SYSTEM_CONTEXT = os.environ.get("RAG_SYSTEM_CONTEXT", "False").lower() == "true"
####################################
# REDIS
+12 -4
View File
@@ -135,7 +135,9 @@ class AuthsTable:
except Exception:
return None
def authenticate_user_by_api_key(self, api_key: str, db: Optional[Session] = None) -> Optional[UserModel]:
def authenticate_user_by_api_key(
self, api_key: str, db: Optional[Session] = None
) -> Optional[UserModel]:
log.info(f"authenticate_user_by_api_key: {api_key}")
# if no api_key, return None
if not api_key:
@@ -147,7 +149,9 @@ class AuthsTable:
except Exception:
return False
def authenticate_user_by_email(self, email: str, db: Optional[Session] = None) -> Optional[UserModel]:
def authenticate_user_by_email(
self, email: str, db: Optional[Session] = None
) -> Optional[UserModel]:
log.info(f"authenticate_user_by_email: {email}")
try:
with get_db_context(db) as db:
@@ -158,7 +162,9 @@ class AuthsTable:
except Exception:
return None
def update_user_password_by_id(self, id: str, new_password: str, db: Optional[Session] = None) -> bool:
def update_user_password_by_id(
self, id: str, new_password: str, db: Optional[Session] = None
) -> bool:
try:
with get_db_context(db) as db:
result = (
@@ -169,7 +175,9 @@ class AuthsTable:
except Exception:
return False
def update_email_by_id(self, id: str, email: str, db: Optional[Session] = None) -> bool:
def update_email_by_id(
self, id: str, email: str, db: Optional[Session] = None
) -> bool:
try:
with get_db_context(db) as db:
result = db.query(Auth).filter_by(id=id).update({"email": email})
+25 -7
View File
@@ -149,7 +149,9 @@ class FeedbackTable:
log.exception(f"Error creating a new feedback: {e}")
return None
def get_feedback_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FeedbackModel]:
def get_feedback_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[FeedbackModel]:
try:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id).first()
@@ -172,7 +174,11 @@ class FeedbackTable:
return None
def get_feedback_items(
self, filter: dict = {}, skip: int = 0, limit: int = 30, db: Optional[Session] = None
self,
filter: dict = {},
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
) -> FeedbackListResponse:
with get_db_context(db) as db:
query = db.query(Feedback, User).join(User, Feedback.user_id == User.id)
@@ -244,7 +250,9 @@ class FeedbackTable:
.all()
]
def get_feedbacks_by_type(self, type: str, db: Optional[Session] = None) -> list[FeedbackModel]:
def get_feedbacks_by_type(
self, type: str, db: Optional[Session] = None
) -> list[FeedbackModel]:
with get_db_context(db) as db:
return [
FeedbackModel.model_validate(feedback)
@@ -254,7 +262,9 @@ class FeedbackTable:
.all()
]
def get_feedbacks_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[FeedbackModel]:
def get_feedbacks_by_user_id(
self, user_id: str, db: Optional[Session] = None
) -> list[FeedbackModel]:
with get_db_context(db) as db:
return [
FeedbackModel.model_validate(feedback)
@@ -285,7 +295,11 @@ class FeedbackTable:
return FeedbackModel.model_validate(feedback)
def update_feedback_by_id_and_user_id(
self, id: str, user_id: str, form_data: FeedbackForm, db: Optional[Session] = None
self,
id: str,
user_id: str,
form_data: FeedbackForm,
db: Optional[Session] = None,
) -> Optional[FeedbackModel]:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
@@ -313,7 +327,9 @@ class FeedbackTable:
db.commit()
return True
def delete_feedback_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool:
def delete_feedback_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[Session] = None
) -> bool:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
if not feedback:
@@ -322,7 +338,9 @@ class FeedbackTable:
db.commit()
return True
def delete_feedbacks_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool:
def delete_feedbacks_by_user_id(
self, user_id: str, db: Optional[Session] = None
) -> bool:
with get_db_context(db) as db:
feedbacks = db.query(Feedback).filter_by(user_id=user_id).all()
if not feedbacks:
+21 -5
View File
@@ -84,7 +84,11 @@ class FolderUpdateForm(BaseModel):
class FolderTable:
def insert_new_folder(
self, user_id: str, form_data: FolderForm, parent_id: Optional[str] = None, db: Optional[Session] = None
self,
user_id: str,
form_data: FolderForm,
parent_id: Optional[str] = None,
db: Optional[Session] = None,
) -> Optional[FolderModel]:
with get_db_context(db) as db:
id = str(uuid.uuid4())
@@ -149,7 +153,9 @@ class FolderTable:
except Exception:
return None
def get_folders_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[FolderModel]:
def get_folders_by_user_id(
self, user_id: str, db: Optional[Session] = None
) -> list[FolderModel]:
with get_db_context(db) as db:
return [
FolderModel.model_validate(folder)
@@ -157,7 +163,11 @@ class FolderTable:
]
def get_folder_by_parent_id_and_user_id_and_name(
self, parent_id: Optional[str], user_id: str, name: str, db: Optional[Session] = None
self,
parent_id: Optional[str],
user_id: str,
name: str,
db: Optional[Session] = None,
) -> Optional[FolderModel]:
try:
with get_db_context(db) as db:
@@ -213,7 +223,11 @@ class FolderTable:
return
def update_folder_by_id_and_user_id(
self, id: str, user_id: str, form_data: FolderUpdateForm, db: Optional[Session] = None
self,
id: str,
user_id: str,
form_data: FolderUpdateForm,
db: Optional[Session] = None,
) -> Optional[FolderModel]:
try:
with get_db_context(db) as db:
@@ -278,7 +292,9 @@ class FolderTable:
log.error(f"update_folder: {e}")
return
def delete_folder_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> list[str]:
def delete_folder_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[Session] = None
) -> list[str]:
try:
folder_ids = []
with get_db_context(db) as db:
+27 -8
View File
@@ -104,7 +104,11 @@ class FunctionValves(BaseModel):
class FunctionsTable:
def insert_new_function(
self, user_id: str, type: str, form_data: FunctionForm, db: Optional[Session] = None
self,
user_id: str,
type: str,
form_data: FunctionForm,
db: Optional[Session] = None,
) -> Optional[FunctionModel]:
function = FunctionModel(
**{
@@ -131,7 +135,10 @@ class FunctionsTable:
return None
def sync_functions(
self, user_id: str, functions: list[FunctionWithValvesModel], db: Optional[Session] = None
self,
user_id: str,
functions: list[FunctionWithValvesModel],
db: Optional[Session] = None,
) -> list[FunctionWithValvesModel]:
# Synchronize functions for a user by updating existing ones, inserting new ones, and removing those that are no longer present.
try:
@@ -178,7 +185,9 @@ class FunctionsTable:
log.exception(f"Error syncing functions for user {user_id}: {e}")
return []
def get_function_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FunctionModel]:
def get_function_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[FunctionModel]:
try:
with get_db_context(db) as db:
function = db.get(Function, id)
@@ -206,7 +215,9 @@ class FunctionsTable:
FunctionModel.model_validate(function) for function in functions
]
def get_function_list(self, db: Optional[Session] = None) -> list[FunctionUserResponse]:
def get_function_list(
self, db: Optional[Session] = None
) -> list[FunctionUserResponse]:
with get_db_context(db) as db:
functions = db.query(Function).order_by(Function.updated_at.desc()).all()
user_ids = list(set(func.user_id for func in functions))
@@ -245,7 +256,9 @@ class FunctionsTable:
for function in db.query(Function).filter_by(type=type).all()
]
def get_global_filter_functions(self, db: Optional[Session] = None) -> list[FunctionModel]:
def get_global_filter_functions(
self, db: Optional[Session] = None
) -> list[FunctionModel]:
with get_db_context(db) as db:
return [
FunctionModel.model_validate(function)
@@ -254,7 +267,9 @@ class FunctionsTable:
.all()
]
def get_global_action_functions(self, db: Optional[Session] = None) -> list[FunctionModel]:
def get_global_action_functions(
self, db: Optional[Session] = None
) -> list[FunctionModel]:
with get_db_context(db) as db:
return [
FunctionModel.model_validate(function)
@@ -263,7 +278,9 @@ class FunctionsTable:
.all()
]
def get_function_valves_by_id(self, id: str, db: Optional[Session] = None) -> Optional[dict]:
def get_function_valves_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[dict]:
with get_db_context(db) as db:
try:
function = db.get(Function, id)
@@ -352,7 +369,9 @@ class FunctionsTable:
)
return None
def update_function_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[FunctionModel]:
def update_function_by_id(
self, id: str, updated: dict, db: Optional[Session] = None
) -> Optional[FunctionModel]:
with get_db_context(db) as db:
try:
db.query(Function).filter_by(id=id).update(
+49 -15
View File
@@ -1,4 +1,3 @@
import json
import logging
import time
@@ -210,7 +209,12 @@ class KnowledgeTable:
return knowledge_bases
def search_knowledge_bases(
self, user_id: str, filter: dict, skip: int = 0, limit: int = 30, db: Optional[Session] = None
self,
user_id: str,
filter: dict,
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
) -> KnowledgeListResponse:
try:
with get_db_context(db) as db:
@@ -320,7 +324,9 @@ class KnowledgeTable:
if user
else None
),
collection=KnowledgeModel.model_validate(knowledge).model_dump(),
collection=KnowledgeModel.model_validate(
knowledge
).model_dump(),
)
)
@@ -330,20 +336,26 @@ class KnowledgeTable:
print("search_knowledge_files error:", e)
return KnowledgeFileListResponse(items=[], total=0)
def check_access_by_user_id(self, id, user_id, permission="write", db: Optional[Session] = None) -> bool:
def check_access_by_user_id(
self, id, user_id, permission="write", db: Optional[Session] = None
) -> bool:
knowledge = self.get_knowledge_by_id(id, db=db)
if not knowledge:
return False
if knowledge.user_id == user_id:
return True
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
user_group_ids = {
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
}
return has_access(user_id, permission, knowledge.access_control, user_group_ids)
def get_knowledge_bases_by_user_id(
self, user_id: str, permission: str = "write", db: Optional[Session] = None
) -> list[KnowledgeUserModel]:
knowledge_bases = self.get_knowledge_bases(db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
user_group_ids = {
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
}
return [
knowledge_base
for knowledge_base in knowledge_bases
@@ -353,7 +365,9 @@ class KnowledgeTable:
)
]
def get_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> Optional[KnowledgeModel]:
def get_knowledge_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[KnowledgeModel]:
try:
with get_db_context(db) as db:
knowledge = db.query(Knowledge).filter_by(id=id).first()
@@ -371,12 +385,16 @@ class KnowledgeTable:
if knowledge.user_id == user_id:
return knowledge
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
user_group_ids = {
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
}
if has_access(user_id, "write", knowledge.access_control, user_group_ids):
return knowledge
return None
def get_knowledges_by_file_id(self, file_id: str, db: Optional[Session] = None) -> list[KnowledgeModel]:
def get_knowledges_by_file_id(
self, file_id: str, db: Optional[Session] = None
) -> list[KnowledgeModel]:
try:
with get_db_context(db) as db:
knowledges = (
@@ -474,7 +492,9 @@ class KnowledgeTable:
print(e)
return KnowledgeFileListResponse(items=[], total=0)
def get_files_by_id(self, knowledge_id: str, db: Optional[Session] = None) -> list[FileModel]:
def get_files_by_id(
self, knowledge_id: str, db: Optional[Session] = None
) -> list[FileModel]:
try:
with get_db_context(db) as db:
files = (
@@ -487,7 +507,9 @@ class KnowledgeTable:
except Exception:
return []
def get_file_metadatas_by_id(self, knowledge_id: str, db: Optional[Session] = None) -> list[FileMetadataResponse]:
def get_file_metadatas_by_id(
self, knowledge_id: str, db: Optional[Session] = None
) -> list[FileMetadataResponse]:
try:
with get_db_context(db) as db:
files = self.get_files_by_id(knowledge_id, db=db)
@@ -496,7 +518,11 @@ class KnowledgeTable:
return []
def add_file_to_knowledge_by_id(
self, knowledge_id: str, file_id: str, user_id: str, db: Optional[Session] = None
self,
knowledge_id: str,
file_id: str,
user_id: str,
db: Optional[Session] = None,
) -> Optional[KnowledgeFileModel]:
with get_db_context(db) as db:
knowledge_file = KnowledgeFileModel(
@@ -522,7 +548,9 @@ class KnowledgeTable:
except Exception:
return None
def remove_file_from_knowledge_by_id(self, knowledge_id: str, file_id: str, db: Optional[Session] = None) -> bool:
def remove_file_from_knowledge_by_id(
self, knowledge_id: str, file_id: str, db: Optional[Session] = None
) -> bool:
try:
with get_db_context(db) as db:
db.query(KnowledgeFile).filter_by(
@@ -533,7 +561,9 @@ class KnowledgeTable:
except Exception:
return False
def reset_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> Optional[KnowledgeModel]:
def reset_knowledge_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[KnowledgeModel]:
try:
with get_db_context(db) as db:
# Delete all knowledge_file entries for this knowledge_id
@@ -554,7 +584,11 @@ class KnowledgeTable:
return None
def update_knowledge_by_id(
self, id: str, form_data: KnowledgeForm, overwrite: bool = False, db: Optional[Session] = None
self,
id: str,
form_data: KnowledgeForm,
overwrite: bool = False,
db: Optional[Session] = None,
) -> Optional[KnowledgeModel]:
try:
with get_db_context(db) as db:
+12 -4
View File
@@ -94,7 +94,9 @@ class MemoriesTable:
except Exception:
return None
def get_memories_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[MemoryModel]:
def get_memories_by_user_id(
self, user_id: str, db: Optional[Session] = None
) -> list[MemoryModel]:
with get_db_context(db) as db:
try:
memories = db.query(Memory).filter_by(user_id=user_id).all()
@@ -102,7 +104,9 @@ class MemoriesTable:
except Exception:
return None
def get_memory_by_id(self, id: str, db: Optional[Session] = None) -> Optional[MemoryModel]:
def get_memory_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[MemoryModel]:
with get_db_context(db) as db:
try:
memory = db.get(Memory, id)
@@ -121,7 +125,9 @@ class MemoriesTable:
except Exception:
return False
def delete_memories_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool:
def delete_memories_by_user_id(
self, user_id: str, db: Optional[Session] = None
) -> bool:
with get_db_context(db) as db:
try:
db.query(Memory).filter_by(user_id=user_id).delete()
@@ -131,7 +137,9 @@ class MemoriesTable:
except Exception:
return False
def delete_memory_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool:
def delete_memory_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[Session] = None
) -> bool:
with get_db_context(db) as db:
try:
memory = db.get(Memory, id)
+7 -6
View File
@@ -538,15 +538,16 @@ class MessageTable:
)
if start_timestamp:
query_builder = query_builder.filter(Message.created_at >= start_timestamp)
query_builder = query_builder.filter(
Message.created_at >= start_timestamp
)
if end_timestamp:
query_builder = query_builder.filter(Message.created_at <= end_timestamp)
query_builder = query_builder.filter(
Message.created_at <= end_timestamp
)
messages = (
query_builder
.order_by(Message.created_at.desc())
.limit(limit)
.all()
query_builder.order_by(Message.created_at.desc()).limit(limit).all()
)
return [MessageModel.model_validate(msg) for msg in messages]
+24 -7
View File
@@ -222,7 +222,9 @@ class ModelsTable:
self, user_id: str, permission: str = "write", db: Optional[Session] = None
) -> list[ModelUserResponse]:
models = self.get_models(db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
user_group_ids = {
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
}
return [
model
for model in models
@@ -273,7 +275,12 @@ class ModelsTable:
return query
def search_models(
self, user_id: str, filter: dict = {}, skip: int = 0, limit: int = 30, db: Optional[Session] = None
self,
user_id: str,
filter: dict = {},
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
) -> ModelListResponse:
with get_db_context(db) as db:
# Join GroupMember so we can order by group_id when requested
@@ -359,7 +366,9 @@ class ModelsTable:
return ModelListResponse(items=models, total=total)
def get_model_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ModelModel]:
def get_model_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[ModelModel]:
try:
with get_db_context(db) as db:
model = db.get(Model, id)
@@ -367,7 +376,9 @@ class ModelsTable:
except Exception:
return None
def get_models_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[ModelModel]:
def get_models_by_ids(
self, ids: list[str], db: Optional[Session] = None
) -> list[ModelModel]:
try:
with get_db_context(db) as db:
models = db.query(Model).filter(Model.id.in_(ids)).all()
@@ -375,7 +386,9 @@ class ModelsTable:
except Exception:
return []
def toggle_model_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ModelModel]:
def toggle_model_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[ModelModel]:
with get_db_context(db) as db:
try:
is_active = db.query(Model).filter_by(id=id).first().is_active
@@ -392,7 +405,9 @@ class ModelsTable:
except Exception:
return None
def update_model_by_id(self, id: str, model: ModelForm, db: Optional[Session] = None) -> Optional[ModelModel]:
def update_model_by_id(
self, id: str, model: ModelForm, db: Optional[Session] = None
) -> Optional[ModelModel]:
try:
with get_db_context(db) as db:
# update only the fields that are present in the model
@@ -428,7 +443,9 @@ class ModelsTable:
except Exception:
return False
def sync_models(self, user_id: str, models: list[ModelModel], db: Optional[Session] = None) -> list[ModelModel]:
def sync_models(
self, user_id: str, models: list[ModelModel], db: Optional[Session] = None
) -> list[ModelModel]:
try:
with get_db_context(db) as db:
# Get existing models
+10 -2
View File
@@ -260,8 +260,16 @@ class NoteTable:
normalized_query = query_key.replace("-", "").replace(" ", "")
query = query.filter(
or_(
func.replace(func.replace(Note.title, "-", ""), " ", "").ilike(f"%{normalized_query}%"),
func.replace(func.replace(cast(Note.data["content"]["md"], Text), "-", ""), " ", "").ilike(f"%{normalized_query}%"),
func.replace(
func.replace(Note.title, "-", ""), " ", ""
).ilike(f"%{normalized_query}%"),
func.replace(
func.replace(
cast(Note.data["content"]["md"], Text), "-", ""
),
" ",
"",
).ilike(f"%{normalized_query}%"),
)
)
+15 -5
View File
@@ -143,7 +143,9 @@ class OAuthSessionTable:
log.error(f"Error creating OAuth session: {e}")
return None
def get_session_by_id(self, session_id: str, db: Optional[Session] = None) -> Optional[OAuthSessionModel]:
def get_session_by_id(
self, session_id: str, db: Optional[Session] = None
) -> Optional[OAuthSessionModel]:
"""Get OAuth session by ID"""
try:
with get_db_context(db) as db:
@@ -197,7 +199,9 @@ class OAuthSessionTable:
log.error(f"Error getting OAuth session by provider and user ID: {e}")
return None
def get_sessions_by_user_id(self, user_id: str, db: Optional[Session] = None) -> List[OAuthSessionModel]:
def get_sessions_by_user_id(
self, user_id: str, db: Optional[Session] = None
) -> List[OAuthSessionModel]:
"""Get all OAuth sessions for a user"""
try:
with get_db_context(db) as db:
@@ -241,7 +245,9 @@ class OAuthSessionTable:
log.error(f"Error updating OAuth session tokens: {e}")
return None
def delete_session_by_id(self, session_id: str, db: Optional[Session] = None) -> bool:
def delete_session_by_id(
self, session_id: str, db: Optional[Session] = None
) -> bool:
"""Delete an OAuth session"""
try:
with get_db_context(db) as db:
@@ -252,7 +258,9 @@ class OAuthSessionTable:
log.error(f"Error deleting OAuth session: {e}")
return False
def delete_sessions_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool:
def delete_sessions_by_user_id(
self, user_id: str, db: Optional[Session] = None
) -> bool:
"""Delete all OAuth sessions for a user"""
try:
with get_db_context(db) as db:
@@ -263,7 +271,9 @@ class OAuthSessionTable:
log.error(f"Error deleting OAuth sessions by user ID: {e}")
return False
def delete_sessions_by_provider(self, provider: str, db: Optional[Session] = None) -> bool:
def delete_sessions_by_provider(
self, provider: str, db: Optional[Session] = None
) -> bool:
"""Delete all OAuth sessions for a provider"""
try:
with get_db_context(db) as db:
+9 -4
View File
@@ -1,4 +1,3 @@
import time
from typing import Optional
@@ -100,7 +99,9 @@ class PromptsTable:
except Exception:
return None
def get_prompt_by_command(self, command: str, db: Optional[Session] = None) -> Optional[PromptModel]:
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()
@@ -135,7 +136,9 @@ class PromptsTable:
self, user_id: str, permission: str = "write", db: Optional[Session] = None
) -> list[PromptUserResponse]:
prompts = self.get_prompts(db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
user_group_ids = {
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
}
return [
prompt
@@ -159,7 +162,9 @@ class PromptsTable:
except Exception:
return None
def delete_prompt_by_command(self, command: str, db: Optional[Session] = None) -> bool:
def delete_prompt_by_command(
self, command: str, db: Optional[Session] = None
) -> bool:
try:
with get_db_context(db) as db:
db.query(Prompt).filter_by(command=command).delete()
+9 -3
View File
@@ -51,7 +51,9 @@ class TagChatIdForm(BaseModel):
class TagTable:
def insert_new_tag(self, name: str, user_id: str, db: Optional[Session] = None) -> Optional[TagModel]:
def insert_new_tag(
self, name: str, user_id: str, db: Optional[Session] = None
) -> Optional[TagModel]:
with get_db_context(db) as db:
id = name.replace(" ", "_").lower()
tag = TagModel(**{"id": id, "user_id": user_id, "name": name})
@@ -79,7 +81,9 @@ class TagTable:
except Exception:
return None
def get_tags_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[TagModel]:
def get_tags_by_user_id(
self, user_id: str, db: Optional[Session] = None
) -> list[TagModel]:
with get_db_context(db) as db:
return [
TagModel.model_validate(tag)
@@ -97,7 +101,9 @@ class TagTable:
)
]
def delete_tag_by_name_and_user_id(self, name: str, user_id: str, db: Optional[Session] = None) -> bool:
def delete_tag_by_name_and_user_id(
self, name: str, user_id: str, db: Optional[Session] = None
) -> bool:
try:
with get_db_context(db) as db:
id = name.replace(" ", "_").lower()
+20 -6
View File
@@ -115,7 +115,11 @@ class ToolValves(BaseModel):
class ToolsTable:
def insert_new_tool(
self, user_id: str, form_data: ToolForm, specs: list[dict], db: Optional[Session] = None
self,
user_id: str,
form_data: ToolForm,
specs: list[dict],
db: Optional[Session] = None,
) -> Optional[ToolModel]:
with get_db_context(db) as db:
tool = ToolModel(
@@ -141,7 +145,9 @@ class ToolsTable:
log.exception(f"Error creating a new tool: {e}")
return None
def get_tool_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ToolModel]:
def get_tool_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[ToolModel]:
try:
with get_db_context(db) as db:
tool = db.get(Tool, id)
@@ -175,7 +181,9 @@ class ToolsTable:
self, user_id: str, permission: str = "write", db: Optional[Session] = None
) -> list[ToolUserModel]:
tools = self.get_tools(db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
user_group_ids = {
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
}
return [
tool
@@ -184,7 +192,9 @@ class ToolsTable:
or has_access(user_id, permission, tool.access_control, user_group_ids)
]
def get_tool_valves_by_id(self, id: str, db: Optional[Session] = None) -> Optional[dict]:
def get_tool_valves_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[dict]:
try:
with get_db_context(db) as db:
tool = db.get(Tool, id)
@@ -193,7 +203,9 @@ class ToolsTable:
log.exception(f"Error getting tool valves by id {id}")
return None
def update_tool_valves_by_id(self, id: str, valves: dict, db: Optional[Session] = None) -> Optional[ToolValves]:
def update_tool_valves_by_id(
self, id: str, valves: dict, db: Optional[Session] = None
) -> Optional[ToolValves]:
try:
with get_db_context(db) as db:
db.query(Tool).filter_by(id=id).update(
@@ -249,7 +261,9 @@ class ToolsTable:
)
return None
def update_tool_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[ToolModel]:
def update_tool_by_id(
self, id: str, updated: dict, db: Optional[Session] = None
) -> Optional[ToolModel]:
try:
with get_db_context(db) as db:
db.query(Tool).filter_by(id=id).update(
+42 -14
View File
@@ -269,7 +269,9 @@ class UsersTable:
else:
return None
def get_user_by_id(self, id: str, db: Optional[Session] = None) -> Optional[UserModel]:
def get_user_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
@@ -277,7 +279,9 @@ class UsersTable:
except Exception:
return None
def get_user_by_api_key(self, api_key: str, db: Optional[Session] = None) -> Optional[UserModel]:
def get_user_by_api_key(
self, api_key: str, db: Optional[Session] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = (
@@ -290,7 +294,9 @@ class UsersTable:
except Exception:
return None
def get_user_by_email(self, email: str, db: Optional[Session] = None) -> Optional[UserModel]:
def get_user_by_email(
self, email: str, db: Optional[Session] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(email=email).first()
@@ -298,7 +304,9 @@ class UsersTable:
except Exception:
return None
def get_user_by_oauth_sub(self, provider: str, sub: str, db: Optional[Session] = None) -> Optional[UserModel]:
def get_user_by_oauth_sub(
self, provider: str, sub: str, db: Optional[Session] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db: # type: Session
dialect_name = db.bind.dialect.name
@@ -455,7 +463,9 @@ class UsersTable:
"total": total,
}
def get_users_by_group_id(self, group_id: str, db: Optional[Session] = None) -> list[UserModel]:
def get_users_by_group_id(
self, group_id: str, db: Optional[Session] = None
) -> list[UserModel]:
with get_db_context(db) as db:
users = (
db.query(User)
@@ -465,7 +475,9 @@ class UsersTable:
)
return [UserModel.model_validate(user) for user in users]
def get_users_by_user_ids(self, user_ids: list[str], db: Optional[Session] = None) -> list[UserStatusModel]:
def get_users_by_user_ids(
self, user_ids: list[str], db: Optional[Session] = None
) -> list[UserStatusModel]:
with get_db_context(db) as db:
users = db.query(User).filter(User.id.in_(user_ids)).all()
return [UserModel.model_validate(user) for user in users]
@@ -486,7 +498,9 @@ class UsersTable:
except Exception:
return None
def get_user_webhook_url_by_id(self, id: str, db: Optional[Session] = None) -> Optional[str]:
def get_user_webhook_url_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[str]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
@@ -511,7 +525,9 @@ class UsersTable:
)
return query.count()
def update_user_role_by_id(self, id: str, role: str, db: Optional[Session] = None) -> Optional[UserModel]:
def update_user_role_by_id(
self, id: str, role: str, db: Optional[Session] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
db.query(User).filter_by(id=id).update({"role": role})
@@ -552,7 +568,9 @@ class UsersTable:
return None
@throttle(DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL)
def update_last_active_by_id(self, id: str, db: Optional[Session] = None) -> Optional[UserModel]:
def update_last_active_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
db.query(User).filter_by(id=id).update(
@@ -597,7 +615,9 @@ class UsersTable:
except Exception:
return None
def update_user_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[UserModel]:
def update_user_by_id(
self, id: str, updated: dict, db: Optional[Session] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
db.query(User).filter_by(id=id).update(updated)
@@ -610,7 +630,9 @@ class UsersTable:
print(e)
return None
def update_user_settings_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[UserModel]:
def update_user_settings_by_id(
self, id: str, updated: dict, db: Optional[Session] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
@@ -651,7 +673,9 @@ class UsersTable:
except Exception:
return False
def get_user_api_key_by_id(self, id: str, db: Optional[Session] = None) -> Optional[str]:
def get_user_api_key_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[str]:
try:
with get_db_context(db) as db:
api_key = db.query(ApiKey).filter_by(user_id=id).first()
@@ -659,7 +683,9 @@ class UsersTable:
except Exception:
return None
def update_user_api_key_by_id(self, id: str, api_key: str, db: Optional[Session] = None) -> bool:
def update_user_api_key_by_id(
self, id: str, api_key: str, db: Optional[Session] = None
) -> bool:
try:
with get_db_context(db) as db:
db.query(ApiKey).filter_by(user_id=id).delete()
@@ -690,7 +716,9 @@ class UsersTable:
except Exception:
return False
def get_valid_user_ids(self, user_ids: list[str], db: Optional[Session] = None) -> list[str]:
def get_valid_user_ids(
self, user_ids: list[str], db: Optional[Session] = None
) -> list[str]:
with get_db_context(db) as db:
users = db.query(User).filter(User.id.in_(user_ids)).all()
return [user.id for user in users]
+83 -25
View File
@@ -117,7 +117,9 @@ async def get_channels(
last_message = Messages.get_last_message_by_channel_id(channel.id, db=db)
last_message_at = last_message.created_at if last_message else None
channel_member = Channels.get_member_by_channel_and_user_id(channel.id, user.id, db=db)
channel_member = Channels.get_member_by_channel_and_user_id(
channel.id, user.id, db=db
)
unread_count = (
Messages.get_unread_message_count(
channel.id, user.id, channel_member.last_read_at, db=db
@@ -135,7 +137,10 @@ async def get_channels(
]
users = [
UserIdNameStatusResponse(
**{**user.model_dump(), "is_active": Users.is_user_active(user.id, db=db)}
**{
**user.model_dump(),
"is_active": Users.is_user_active(user.id, db=db),
}
)
for user in Users.get_users_by_user_ids(user_ids, db=db)
]
@@ -187,11 +192,15 @@ async def get_dm_channel_by_user_id(
)
try:
existing_channel = Channels.get_dm_channel_by_user_ids([user.id, user_id], db=db)
existing_channel = Channels.get_dm_channel_by_user_ids(
[user.id, user_id], db=db
)
if existing_channel:
participant_ids = [
member.user_id
for member in Channels.get_members_by_channel_id(existing_channel.id, db=db)
for member in Channels.get_members_by_channel_id(
existing_channel.id, db=db
)
]
await emit_to_users(
@@ -203,7 +212,9 @@ async def get_dm_channel_by_user_id(
f"channel:{existing_channel.id}", participant_ids
)
Channels.update_member_active_status(existing_channel.id, user.id, True, db=db)
Channels.update_member_active_status(
existing_channel.id, user.id, True, db=db
)
return ChannelModel(**existing_channel.model_dump())
channel = Channels.insert_new_channel(
@@ -288,7 +299,9 @@ async def create_new_channel(
f"channel:{existing_channel.id}", participant_ids
)
Channels.update_member_active_status(existing_channel.id, user.id, True, db=db)
Channels.update_member_active_status(
existing_channel.id, user.id, True, db=db
)
return ChannelModel(**existing_channel.model_dump())
channel = Channels.insert_new_channel(form_data, user.id, db=db)
@@ -353,17 +366,23 @@ async def get_channel_by_id(
)
user_ids = [
member.user_id for member in Channels.get_members_by_channel_id(channel.id, db=db)
member.user_id
for member in Channels.get_members_by_channel_id(channel.id, db=db)
]
users = [
UserIdNameStatusResponse(
**{**user.model_dump(), "is_active": Users.is_user_active(user.id, db=db)}
**{
**user.model_dump(),
"is_active": Users.is_user_active(user.id, db=db),
}
)
for user in Users.get_users_by_user_ids(user_ids, db=db)
]
channel_member = Channels.get_member_by_channel_and_user_id(channel.id, user.id, db=db)
channel_member = Channels.get_member_by_channel_and_user_id(
channel.id, user.id, db=db
)
unread_count = Messages.get_unread_message_count(
channel.id, user.id, channel_member.last_read_at if channel_member else None
)
@@ -373,7 +392,9 @@ async def get_channel_by_id(
**channel.model_dump(),
"user_ids": user_ids,
"users": users,
"is_manager": Channels.is_user_channel_manager(channel.id, user.id, db=db),
"is_manager": Channels.is_user_channel_manager(
channel.id, user.id, db=db
),
"write_access": True,
"user_count": len(user_ids),
"last_read_at": channel_member.last_read_at if channel_member else None,
@@ -389,12 +410,18 @@ async def get_channel_by_id(
)
write_access = has_access(
user.id, type="write", access_control=channel.access_control, strict=False, db=db
user.id,
type="write",
access_control=channel.access_control,
strict=False,
db=db,
)
user_count = len(get_users_with_access("read", channel.access_control))
channel_member = Channels.get_member_by_channel_and_user_id(channel.id, user.id, db=db)
channel_member = Channels.get_member_by_channel_and_user_id(
channel.id, user.id, db=db
)
unread_count = Messages.get_unread_message_count(
channel.id, user.id, channel_member.last_read_at if channel_member else None
)
@@ -404,7 +431,9 @@ async def get_channel_by_id(
**channel.model_dump(),
"user_ids": user_ids,
"users": users,
"is_manager": Channels.is_user_channel_manager(channel.id, user.id, db=db),
"is_manager": Channels.is_user_channel_manager(
channel.id, user.id, db=db
),
"write_access": write_access or user.role == "admin",
"user_count": user_count,
"last_read_at": channel_member.last_read_at if channel_member else None,
@@ -453,7 +482,8 @@ async def get_channel_members_by_id(
if channel.type == "dm":
user_ids = [
member.user_id for member in Channels.get_members_by_channel_id(channel.id, db=db)
member.user_id
for member in Channels.get_members_by_channel_id(channel.id, db=db)
]
users = Users.get_users_by_user_ids(user_ids, db=db)
total = len(users)
@@ -533,7 +563,9 @@ async def update_is_active_member_by_id_and_user_id(
status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND
)
Channels.update_member_active_status(channel.id, user.id, form_data.is_active, db=db)
Channels.update_member_active_status(
channel.id, user.id, form_data.is_active, db=db
)
return True
@@ -626,7 +658,9 @@ async def remove_members_by_id(
)
try:
deleted = Channels.remove_members_from_channel(channel.id, form_data.user_ids, db=db)
deleted = Channels.remove_members_from_channel(
channel.id, form_data.user_ids, db=db
)
return deleted
except Exception as e:
@@ -794,7 +828,9 @@ async def get_channel_messages(
**message.model_dump(),
"reply_count": len(thread_replies),
"latest_reply_at": latest_thread_reply_at,
"reactions": Messages.get_reactions_by_message_id(message.id, db=db),
"reactions": Messages.get_reactions_by_message_id(
message.id, db=db
),
"user": UserNameResponse(**users[message.user_id].model_dump()),
}
)
@@ -857,7 +893,9 @@ async def get_pinned_channel_messages(
MessageWithReactionsResponse(
**{
**message.model_dump(),
"reactions": Messages.get_reactions_by_message_id(message.id, db=db),
"reactions": Messages.get_reactions_by_message_id(
message.id, db=db
),
"user": UserNameResponse(**users[message.user_id].model_dump()),
}
)
@@ -871,7 +909,9 @@ async def get_pinned_channel_messages(
############################
async def send_notification(name, webui_url, channel, message, active_user_ids, db=None):
async def send_notification(
name, webui_url, channel, message, active_user_ids, db=None
):
users = get_users_with_access("read", channel.access_control)
for user in users:
@@ -966,7 +1006,9 @@ async def model_response_handler(request, channel, message, user, db=None):
for thread_message in thread_messages:
message_user = None
if thread_message.user_id not in message_users:
message_user = Users.get_user_by_id(thread_message.user_id, db=db)
message_user = Users.get_user_by_id(
thread_message.user_id, db=db
)
message_users[thread_message.user_id] = message_user
else:
message_user = message_users[thread_message.user_id]
@@ -1098,7 +1140,11 @@ async def new_message_handler(
)
else:
if user.role != "admin" and not has_access(
user.id, type="write", access_control=channel.access_control, strict=False, db=db
user.id,
type="write",
access_control=channel.access_control,
strict=False,
db=db,
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()
@@ -1417,7 +1463,9 @@ async def get_channel_thread_messages(
status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()
)
message_list = Messages.get_messages_by_parent_id(id, message_id, skip, limit, db=db)
message_list = Messages.get_messages_by_parent_id(
id, message_id, skip, limit, db=db
)
if not message_list:
return []
@@ -1434,7 +1482,9 @@ async def get_channel_thread_messages(
**message.model_dump(),
"reply_count": 0,
"latest_reply_at": None,
"reactions": Messages.get_reactions_by_message_id(message.id, db=db),
"reactions": Messages.get_reactions_by_message_id(
message.id, db=db
),
"user": UserNameResponse(**users[message.user_id].model_dump()),
}
)
@@ -1554,7 +1604,11 @@ async def add_reaction_to_message(
)
else:
if user.role != "admin" and not has_access(
user.id, type="write", access_control=channel.access_control, strict=False, db=db
user.id,
type="write",
access_control=channel.access_control,
strict=False,
db=db,
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()
@@ -1629,7 +1683,11 @@ async def remove_reaction_by_id_and_user_id_and_name(
)
else:
if user.role != "admin" and not has_access(
user.id, type="write", access_control=channel.access_control, strict=False, db=db
user.id,
type="write",
access_control=channel.access_control,
strict=False,
db=db,
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()
+1 -1
View File
@@ -383,7 +383,7 @@ def calculate_chat_stats(
def generate_chat_stats_jsonl_generator(user_id, filter):
"""
Synchronous generator for streaming chat stats export.
NOTE: We intentionally do NOT pass a shared db session here. Instead, we let
each batch create its own short-lived session via get_db_context(None).
This is critical for SQLite in low-resource environments because:
+39 -15
View File
@@ -18,8 +18,6 @@ log = logging.getLogger(__name__)
router = APIRouter()
@router.get("/ef")
async def get_embeddings(request: Request):
return {"result": await request.app.state.EMBEDDING_FUNCTION("hello world")}
@@ -32,7 +30,9 @@ async def get_embeddings(request: Request):
@router.get("/", response_model=list[MemoryModel])
async def get_memories(
request: Request, user=Depends(get_verified_user), db: Session = Depends(get_session)
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
if not request.app.state.config.ENABLE_MEMORIES:
raise HTTPException(
@@ -40,7 +40,9 @@ async def get_memories(
detail=ERROR_MESSAGES.NOT_FOUND,
)
if not has_permission(user.id, "features.memories", request.app.state.config.USER_PERMISSIONS):
if not has_permission(
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
@@ -75,7 +77,9 @@ async def add_memory(
detail=ERROR_MESSAGES.NOT_FOUND,
)
if not has_permission(user.id, "features.memories", request.app.state.config.USER_PERMISSIONS):
if not has_permission(
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
@@ -112,7 +116,10 @@ class QueryMemoryForm(BaseModel):
@router.post("/query")
async def query_memory(
request: Request, form_data: QueryMemoryForm, user=Depends(get_verified_user), db: Session = Depends(get_session)
request: Request,
form_data: QueryMemoryForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
if not request.app.state.config.ENABLE_MEMORIES:
raise HTTPException(
@@ -120,7 +127,9 @@ async def query_memory(
detail=ERROR_MESSAGES.NOT_FOUND,
)
if not has_permission(user.id, "features.memories", request.app.state.config.USER_PERMISSIONS):
if not has_permission(
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
@@ -146,7 +155,9 @@ async def query_memory(
############################
@router.post("/reset", response_model=bool)
async def reset_memory_from_vector_db(
request: Request, user=Depends(get_verified_user), db: Session = Depends(get_session)
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
if not request.app.state.config.ENABLE_MEMORIES:
raise HTTPException(
@@ -154,12 +165,14 @@ async def reset_memory_from_vector_db(
detail=ERROR_MESSAGES.NOT_FOUND,
)
if not has_permission(user.id, "features.memories", request.app.state.config.USER_PERMISSIONS):
if not has_permission(
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
VECTOR_DB_CLIENT.delete_collection(f"user-memory-{user.id}")
memories = Memories.get_memories_by_user_id(user.id, db=db)
@@ -198,7 +211,9 @@ async def reset_memory_from_vector_db(
@router.delete("/delete/user", response_model=bool)
async def delete_memory_by_user_id(
request: Request, user=Depends(get_verified_user), db: Session = Depends(get_session)
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
if not request.app.state.config.ENABLE_MEMORIES:
raise HTTPException(
@@ -206,7 +221,9 @@ async def delete_memory_by_user_id(
detail=ERROR_MESSAGES.NOT_FOUND,
)
if not has_permission(user.id, "features.memories", request.app.state.config.USER_PERMISSIONS):
if not has_permission(
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
@@ -243,7 +260,9 @@ async def update_memory_by_id(
detail=ERROR_MESSAGES.NOT_FOUND,
)
if not has_permission(user.id, "features.memories", request.app.state.config.USER_PERMISSIONS):
if not has_permission(
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
@@ -283,7 +302,10 @@ async def update_memory_by_id(
@router.delete("/{memory_id}", response_model=bool)
async def delete_memory_by_id(
memory_id: str, request: Request, user=Depends(get_verified_user), db: Session = Depends(get_session)
memory_id: str,
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
if not request.app.state.config.ENABLE_MEMORIES:
raise HTTPException(
@@ -291,7 +313,9 @@ async def delete_memory_by_id(
detail=ERROR_MESSAGES.NOT_FOUND,
)
if not has_permission(user.id, "features.memories", request.app.state.config.USER_PERMISSIONS):
if not has_permission(
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
+16 -4
View File
@@ -1302,7 +1302,10 @@ async def generate_chat_completion(
if not (
user.id == model_info.user_id
or has_access(
user.id, type="read", access_control=model_info.access_control, db=db
user.id,
type="read",
access_control=model_info.access_control,
db=db,
)
):
raise HTTPException(
@@ -1409,7 +1412,10 @@ async def generate_openai_completion(
if not (
user.id == model_info.user_id
or has_access(
user.id, type="read", access_control=model_info.access_control, db=db
user.id,
type="read",
access_control=model_info.access_control,
db=db,
)
):
raise HTTPException(
@@ -1493,7 +1499,10 @@ async def generate_openai_chat_completion(
if not (
user.id == model_info.user_id
or has_access(
user.id, type="read", access_control=model_info.access_control, db=db
user.id,
type="read",
access_control=model_info.access_control,
db=db,
)
):
raise HTTPException(
@@ -1592,7 +1601,10 @@ async def get_openai_models(
model_info = Models.get_model_by_id(model["id"], db=db)
if model_info:
if user.id == model_info.user_id or has_access(
user.id, type="read", access_control=model_info.access_control, db=db
user.id,
type="read",
access_control=model_info.access_control,
db=db,
):
filtered_models.append(model)
models = filtered_models
+4 -1
View File
@@ -837,7 +837,10 @@ async def generate_chat_completion(
if not (
user.id == model_info.user_id
or has_access(
user.id, type="read", access_control=model_info.access_control, db=db
user.id,
type="read",
access_control=model_info.access_control,
db=db,
)
):
raise HTTPException(
+20 -6
View File
@@ -23,7 +23,9 @@ router = APIRouter()
@router.get("/", response_model=list[PromptModel])
async def get_prompts(user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def get_prompts(
user=Depends(get_verified_user), db: Session = Depends(get_session)
):
if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL:
prompts = Prompts.get_prompts(db=db)
else:
@@ -33,7 +35,9 @@ async def get_prompts(user=Depends(get_verified_user), db: Session = Depends(get
@router.get("/list", response_model=list[PromptAccessResponse])
async def get_prompt_list(user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def get_prompt_list(
user=Depends(get_verified_user), db: Session = Depends(get_session)
):
if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL:
prompts = Prompts.get_prompts(db=db)
else:
@@ -59,11 +63,17 @@ async def get_prompt_list(user=Depends(get_verified_user), db: Session = Depends
@router.post("/create", response_model=Optional[PromptModel])
async def create_new_prompt(
request: Request, form_data: PromptForm, user=Depends(get_verified_user), db: Session = Depends(get_session)
request: Request,
form_data: PromptForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
if user.role != "admin" and not (
has_permission(
user.id, "workspace.prompts", request.app.state.config.USER_PERMISSIONS, db=db
user.id,
"workspace.prompts",
request.app.state.config.USER_PERMISSIONS,
db=db,
)
or has_permission(
user.id,
@@ -99,7 +109,9 @@ async def create_new_prompt(
@router.get("/command/{command}", response_model=Optional[PromptAccessResponse])
async def get_prompt_by_command(command: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
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)
if prompt:
@@ -169,7 +181,9 @@ async def update_prompt_by_command(
@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)):
async def delete_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)
if not prompt:
raise HTTPException(
+7 -5
View File
@@ -246,11 +246,13 @@ def get_user_ids_from_room(room):
active_session_ids = get_session_ids_from_room(room)
active_user_ids = list(
set([
SESSION_POOL.get(session_id)["id"]
for session_id in active_session_ids
if SESSION_POOL.get(session_id) is not None
])
set(
[
SESSION_POOL.get(session_id)["id"]
for session_id in active_session_ids
if SESSION_POOL.get(session_id) is not None
]
)
)
return active_user_ids
+20 -10
View File
@@ -1378,7 +1378,9 @@ async def view_knowledge_file(
if (
user_role == "admin"
or knowledge_base.user_id == user_id
or has_access(user_id, "read", knowledge_base.access_control, user_group_ids)
or has_access(
user_id, "read", knowledge_base.access_control, user_group_ids
)
):
has_knowledge_access = True
knowledge_info = {"id": knowledge_base.id, "name": knowledge_base.name}
@@ -1463,7 +1465,9 @@ async def query_knowledge_bases(
if knowledge and (
user_role == "admin"
or knowledge.user_id == user_id
or has_access(user_id, "read", knowledge.access_control, user_group_ids)
or has_access(
user_id, "read", knowledge.access_control, user_group_ids
)
):
collection_names.append(item_id)
@@ -1482,12 +1486,14 @@ async def query_knowledge_bases(
or has_access(user_id, "read", note.access_control)
):
content = note.data.get("content", {}).get("md", "")
note_results.append({
"content": content,
"source": note.title,
"note_id": note.id,
"type": "note",
})
note_results.append(
{
"content": content,
"source": note.title,
"note_id": note.id,
"type": "note",
}
)
elif knowledge_ids:
# User specified specific KBs
@@ -1496,7 +1502,9 @@ async def query_knowledge_bases(
if knowledge and (
user_role == "admin"
or knowledge.user_id == user_id
or has_access(user_id, "read", knowledge.access_control, user_group_ids)
or has_access(
user_id, "read", knowledge.access_control, user_group_ids
)
):
collection_names.append(knowledge_id)
else:
@@ -1535,7 +1543,9 @@ async def query_knowledge_bases(
for idx, doc in enumerate(documents):
chunk_info = {
"content": doc,
"source": metadatas[idx].get("source", metadatas[idx].get("name", "Unknown")),
"source": metadatas[idx].get(
"source", metadatas[idx].get("name", "Unknown")
),
"file_id": metadatas[idx].get("file_id", ""),
}
if idx < len(distances):
+2 -5
View File
@@ -948,7 +948,7 @@ def get_image_urls(delta_images, request, metadata, user) -> list[str]:
def add_file_context(messages: list, chat_id: str, user) -> list:
"""
Add file URLs to user messages for native function calling.
Add file URLs to messages for native function calling.
"""
if not chat_id or chat_id.startswith("local:"):
return messages
@@ -961,7 +961,6 @@ def add_file_context(messages: list, chat_id: str, user) -> list:
stored_messages = get_message_list(
history.get("messages", {}), history.get("currentId")
)
stored_user_messages = [msg for msg in stored_messages if msg.get("role") == "user"]
def format_file_tag(file):
attrs = f'type="{file.get("type", "file")}" url="{file["url"]}"'
@@ -971,9 +970,7 @@ def add_file_context(messages: list, chat_id: str, user) -> list:
attrs += f' name="{file["name"]}"'
return f"<file {attrs}/>"
user_messages = [msg for msg in messages if msg.get("role") == "user"]
for message, stored_message in zip(user_messages, stored_user_messages):
for message, stored_message in zip(messages, stored_messages):
files_with_urls = [
file for file in stored_message.get("files", []) if file.get("url")
]
+19 -8
View File
@@ -360,7 +360,12 @@ def get_builtin_tools(
# Helper to get model capabilities (defaults to True if not specified)
def get_model_capability(name: str, default: bool = True) -> bool:
return model.get("info", {}).get("meta", {}).get("capabilities", {}).get(name, default)
return (
model.get("info", {})
.get("meta", {})
.get("capabilities", {})
.get(name, default)
)
# Time utilities - always available for date calculations
builtin_functions.extend([get_current_timestamp, calculate_timestamp])
@@ -375,7 +380,13 @@ def get_builtin_tools(
else:
# No model knowledge - allow full KB browsing
builtin_functions.extend(
[list_knowledge_bases, search_knowledge_bases, search_knowledge_files, view_knowledge_file, query_knowledge_bases]
[
list_knowledge_bases,
search_knowledge_bases,
search_knowledge_files,
view_knowledge_file,
query_knowledge_bases,
]
)
# Chats tools - search and fetch user's chat history
@@ -386,9 +397,9 @@ def get_builtin_tools(
builtin_functions.extend([search_memories, add_memory, replace_memory_content])
# Add web search tools if enabled globally AND model has web_search capability
if getattr(request.app.state.config, "ENABLE_WEB_SEARCH", False) and get_model_capability(
"web_search"
):
if getattr(
request.app.state.config, "ENABLE_WEB_SEARCH", False
) and get_model_capability("web_search"):
builtin_functions.extend([search_web, fetch_url])
# Add image generation/edit tools if enabled globally AND model has image_generation capability
@@ -396,9 +407,9 @@ def get_builtin_tools(
request.app.state.config, "ENABLE_IMAGE_GENERATION", False
) and get_model_capability("image_generation"):
builtin_functions.append(generate_image)
if getattr(request.app.state.config, "ENABLE_IMAGE_EDIT", False) and get_model_capability(
"image_generation"
):
if getattr(
request.app.state.config, "ENABLE_IMAGE_EDIT", False
) and get_model_capability("image_generation"):
builtin_functions.append(edit_image)
# Notes tools - search, view, create, and update user's notes (if notes enabled globally)