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
+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]