chore: format
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}%"),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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")
|
||||
]
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user