From 5921a19519e117947db7cd79b638e96045017412 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Tue, 6 Jan 2026 02:19:57 +0400 Subject: [PATCH] refac --- backend/open_webui/models/chats.py | 161 ++++++++++++++++++------- backend/open_webui/tools/builtin.py | 54 ++++++--- backend/open_webui/utils/middleware.py | 2 + backend/open_webui/utils/tools.py | 2 + 4 files changed, 160 insertions(+), 59 deletions(-) diff --git a/backend/open_webui/models/chats.py b/backend/open_webui/models/chats.py index ea48044bc..12359eec9 100644 --- a/backend/open_webui/models/chats.py +++ b/backend/open_webui/models/chats.py @@ -281,7 +281,9 @@ class ChatTable: return changed - def insert_new_chat(self, user_id: str, form_data: ChatForm, db: Optional[Session] = None) -> Optional[ChatModel]: + def insert_new_chat( + self, user_id: str, form_data: ChatForm, db: Optional[Session] = None + ) -> Optional[ChatModel]: with get_db_context(db) as db: id = str(uuid.uuid4()) chat = ChatModel( @@ -332,7 +334,10 @@ class ChatTable: return chat def import_chats( - self, user_id: str, chat_import_forms: list[ChatImportForm], db: Optional[Session] = None + self, + user_id: str, + chat_import_forms: list[ChatImportForm], + db: Optional[Session] = None, ) -> list[ChatModel]: with get_db_context(db) as db: chats = [] @@ -345,7 +350,9 @@ class ChatTable: db.commit() return [ChatModel.model_validate(chat) for chat in chats] - def update_chat_by_id(self, id: str, chat: dict, db: Optional[Session] = None) -> Optional[ChatModel]: + def update_chat_by_id( + self, id: str, chat: dict, db: Optional[Session] = None + ) -> Optional[ChatModel]: try: with get_db_context(db) as db: chat_item = db.get(Chat, id) @@ -466,25 +473,28 @@ class ChatTable: def add_message_files_by_id_and_message_id( self, id: str, message_id: str, files: list[dict] ) -> list[dict]: - chat = self.get_chat_by_id(id) - if chat is None: - return None + with get_db_context() as db: + chat = self.get_chat_by_id(id, db=db) + if chat is None: + return None - chat = chat.chat - history = chat.get("history", {}) + chat = chat.chat + history = chat.get("history", {}) - message_files = [] + message_files = [] - if message_id in history.get("messages", {}): - message_files = history["messages"][message_id].get("files", []) - message_files = message_files + files - history["messages"][message_id]["files"] = message_files + if message_id in history.get("messages", {}): + message_files = history["messages"][message_id].get("files", []) + message_files = message_files + files + history["messages"][message_id]["files"] = message_files - chat["history"] = history - self.update_chat_by_id(id, chat) - return message_files + chat["history"] = history + self.update_chat_by_id(id, chat, db=db) + return message_files - def insert_shared_chat_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + def insert_shared_chat_by_chat_id( + self, chat_id: str, db: Optional[Session] = None + ) -> Optional[ChatModel]: with get_db_context(db) as db: # Get the existing chat to share chat = db.get(Chat, chat_id) @@ -522,7 +532,9 @@ class ChatTable: db.commit() return shared_chat if (shared_result and result) else None - def update_shared_chat_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + def update_shared_chat_by_chat_id( + self, chat_id: str, db: Optional[Session] = None + ) -> Optional[ChatModel]: try: with get_db_context(db) as db: chat = db.get(Chat, chat_id) @@ -546,7 +558,9 @@ class ChatTable: except Exception: return None - def delete_shared_chat_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> bool: + def delete_shared_chat_by_chat_id( + self, chat_id: str, db: Optional[Session] = None + ) -> bool: try: with get_db_context(db) as db: db.query(Chat).filter_by(user_id=f"shared-{chat_id}").delete() @@ -556,7 +570,9 @@ class ChatTable: except Exception: return False - def unarchive_all_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: + def unarchive_all_chats_by_user_id( + self, user_id: str, db: Optional[Session] = None + ) -> bool: try: with get_db_context(db) as db: db.query(Chat).filter_by(user_id=user_id).update({"archived": False}) @@ -578,7 +594,9 @@ class ChatTable: except Exception: return None - def toggle_chat_pinned_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + def toggle_chat_pinned_by_id( + self, id: str, db: Optional[Session] = None + ) -> Optional[ChatModel]: try: with get_db_context(db) as db: chat = db.get(Chat, id) @@ -590,7 +608,9 @@ class ChatTable: except Exception: return None - def toggle_chat_archive_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + def toggle_chat_archive_by_id( + self, id: str, db: Optional[Session] = None + ) -> Optional[ChatModel]: try: with get_db_context(db) as db: chat = db.get(Chat, id) @@ -603,7 +623,9 @@ class ChatTable: except Exception: return None - def archive_all_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: + def archive_all_chats_by_user_id( + self, user_id: str, db: Optional[Session] = None + ) -> bool: try: with get_db_context(db) as db: db.query(Chat).filter_by(user_id=user_id).update({"archived": True}) @@ -740,7 +762,11 @@ class ChatTable: ] def get_chat_list_by_chat_ids( - self, chat_ids: list[str], skip: int = 0, limit: int = 50, db: Optional[Session] = None + self, + chat_ids: list[str], + skip: int = 0, + limit: int = 50, + db: Optional[Session] = None, ) -> list[ChatModel]: with get_db_context(db) as db: all_chats = ( @@ -752,7 +778,9 @@ class ChatTable: ) return [ChatModel.model_validate(chat) for chat in all_chats] - def get_chat_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + def get_chat_by_id( + self, id: str, db: Optional[Session] = None + ) -> Optional[ChatModel]: try: with get_db_context(db) as db: chat_item = db.get(Chat, id) @@ -767,7 +795,9 @@ class ChatTable: except Exception: return None - def get_chat_by_share_id(self, id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + def get_chat_by_share_id( + self, id: str, db: Optional[Session] = None + ) -> Optional[ChatModel]: try: with get_db_context(db) as db: # it is possible that the shared link was deleted. hence, @@ -781,7 +811,9 @@ class ChatTable: except Exception: return None - def get_chat_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + def get_chat_by_id_and_user_id( + self, id: str, user_id: str, db: Optional[Session] = None + ) -> Optional[ChatModel]: try: with get_db_context(db) as db: chat = db.query(Chat).filter_by(id=id, user_id=user_id).first() @@ -789,7 +821,9 @@ class ChatTable: except Exception: return None - def get_chats(self, skip: int = 0, limit: int = 50, db: Optional[Session] = None) -> list[ChatModel]: + def get_chats( + self, skip: int = 0, limit: int = 50, db: Optional[Session] = None + ) -> list[ChatModel]: with get_db_context(db) as db: all_chats = ( db.query(Chat) @@ -844,7 +878,9 @@ class ChatTable: } ) - def get_pinned_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[ChatModel]: + def get_pinned_chats_by_user_id( + self, user_id: str, db: Optional[Session] = None + ) -> list[ChatModel]: with get_db_context(db) as db: all_chats = ( db.query(Chat) @@ -853,7 +889,9 @@ class ChatTable: ) return [ChatModel.model_validate(chat) for chat in all_chats] - def get_archived_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[ChatModel]: + def get_archived_chats_by_user_id( + self, user_id: str, db: Optional[Session] = None + ) -> list[ChatModel]: with get_db_context(db) as db: all_chats = ( db.query(Chat) @@ -1074,7 +1112,12 @@ class ChatTable: return [ChatModel.model_validate(chat) for chat in all_chats] def get_chats_by_folder_id_and_user_id( - self, folder_id: str, user_id: str, skip: int = 0, limit: int = 60, db: Optional[Session] = None + self, + folder_id: str, + user_id: str, + skip: int = 0, + limit: int = 60, + db: Optional[Session] = None, ) -> list[ChatModel]: with get_db_context(db) as db: query = db.query(Chat).filter_by(folder_id=folder_id, user_id=user_id) @@ -1121,14 +1164,21 @@ class ChatTable: except Exception: return None - def get_chat_tags_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> list[TagModel]: + def get_chat_tags_by_id_and_user_id( + self, id: str, user_id: str, db: Optional[Session] = None + ) -> list[TagModel]: with get_db_context(db) as db: chat = db.get(Chat, id) tags = chat.meta.get("tags", []) return [Tags.get_tag_by_name_and_user_id(tag, user_id) for tag in tags] def get_chat_list_by_user_id_and_tag_name( - self, user_id: str, tag_name: str, skip: int = 0, limit: int = 50, db: Optional[Session] = None + self, + user_id: str, + tag_name: str, + skip: int = 0, + limit: int = 50, + db: Optional[Session] = None, ) -> list[ChatModel]: with get_db_context(db) as db: query = db.query(Chat).filter_by(user_id=user_id) @@ -1181,7 +1231,9 @@ class ChatTable: except Exception: return None - def count_chats_by_tag_name_and_user_id(self, tag_name: str, user_id: str, db: Optional[Session] = None) -> int: + def count_chats_by_tag_name_and_user_id( + self, tag_name: str, user_id: str, db: Optional[Session] = None + ) -> int: with get_db_context(db) as db: # Assuming `get_db()` returns a session object query = db.query(Chat).filter_by(user_id=user_id, archived=False) @@ -1217,7 +1269,9 @@ class ChatTable: return count - def count_chats_by_folder_id_and_user_id(self, folder_id: str, user_id: str, db: Optional[Session] = None) -> int: + def count_chats_by_folder_id_and_user_id( + self, folder_id: str, user_id: str, db: Optional[Session] = None + ) -> int: with get_db_context(db) as db: query = db.query(Chat).filter_by(user_id=user_id) @@ -1246,7 +1300,9 @@ class ChatTable: except Exception: return False - def delete_all_tags_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool: + def delete_all_tags_by_id_and_user_id( + self, id: str, user_id: str, db: Optional[Session] = None + ) -> bool: try: with get_db_context(db) as db: chat = db.get(Chat, id) @@ -1270,7 +1326,9 @@ class ChatTable: except Exception: return False - def delete_chat_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool: + def delete_chat_by_id_and_user_id( + self, id: str, user_id: str, db: Optional[Session] = None + ) -> bool: try: with get_db_context(db) as db: db.query(Chat).filter_by(id=id, user_id=user_id).delete() @@ -1280,7 +1338,9 @@ class ChatTable: except Exception: return False - def delete_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: + def delete_chats_by_user_id( + self, user_id: str, db: Optional[Session] = None + ) -> bool: try: with get_db_context(db) as db: self.delete_shared_chats_by_user_id(user_id, db=db) @@ -1305,7 +1365,11 @@ class ChatTable: return False def move_chats_by_user_id_and_folder_id( - self, user_id: str, folder_id: str, new_folder_id: Optional[str], db: Optional[Session] = None + self, + user_id: str, + folder_id: str, + new_folder_id: Optional[str], + db: Optional[Session] = None, ) -> bool: try: with get_db_context(db) as db: @@ -1318,7 +1382,9 @@ class ChatTable: except Exception: return False - def delete_shared_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: + def delete_shared_chats_by_user_id( + self, user_id: str, db: Optional[Session] = None + ) -> bool: try: with get_db_context(db) as db: chats_by_user = db.query(Chat).filter_by(user_id=user_id).all() @@ -1332,7 +1398,12 @@ class ChatTable: return False def insert_chat_files( - self, chat_id: str, message_id: str, file_ids: list[str], user_id: str, db: Optional[Session] = None + self, + chat_id: str, + message_id: str, + file_ids: list[str], + user_id: str, + db: Optional[Session] = None, ) -> Optional[list[ChatFileModel]]: if not file_ids: return None @@ -1398,7 +1469,9 @@ class ChatTable: ChatFileModel.model_validate(chat_file) for chat_file in all_chat_files ] - def delete_chat_file(self, chat_id: str, file_id: str, db: Optional[Session] = None) -> bool: + def delete_chat_file( + self, chat_id: str, file_id: str, db: Optional[Session] = None + ) -> bool: try: with get_db_context(db) as db: db.query(ChatFile).filter_by(chat_id=chat_id, file_id=file_id).delete() @@ -1407,7 +1480,9 @@ class ChatTable: except Exception: return False - def get_shared_chats_by_file_id(self, file_id: str, db: Optional[Session] = None) -> list[ChatModel]: + def get_shared_chats_by_file_id( + self, file_id: str, db: Optional[Session] = None + ) -> list[ChatModel]: with get_db_context(db) as db: # Join Chat and ChatFile tables to get shared chats associated with the file_id all_chats = ( diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index e5dc9f423..11069ae01 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -214,6 +214,8 @@ async def generate_image( __request__: Request = None, __user__: dict = None, __event_emitter__: callable = None, + __chat_id__: str = None, + __message_id__: str = None, ) -> str: """ Generate an image based on a text prompt. @@ -233,15 +235,24 @@ async def generate_image( user=user, ) + # Prepare file entries for the images + image_files = [{"type": "image", "url": img["url"]} for img in images] + + # Persist files to DB if chat context is available + if __chat_id__ and __message_id__ and images: + image_files = Chats.add_message_files_by_id_and_message_id( + __chat_id__, + __message_id__, + image_files, + ) + # Emit the images to the UI if event emitter is available - if __event_emitter__ and images: + if __event_emitter__ and image_files: await __event_emitter__( { - "type": "files", + "type": "chat:message:files", "data": { - "files": [ - {"type": "image", "url": img["url"]} for img in images - ] + "files": image_files, }, } ) @@ -263,17 +274,19 @@ async def generate_image( async def edit_image( prompt: str, - image_url: str, + image_urls: list[str], __request__: Request = None, __user__: dict = None, __event_emitter__: callable = None, + __chat_id__: str = None, + __message_id__: str = None, ) -> str: """ - Edit an existing image based on a text prompt. + Edit existing images based on a text prompt. - :param prompt: A description of the changes to make to the image - :param image_url: The URL of the image to edit - :return: Confirmation that the image was edited, or an error message + :param prompt: A description of the changes to make to the images + :param image_urls: A list of URLs of the images to edit + :return: Confirmation that the images were edited, or an error message """ if __request__ is None: return json.dumps({"error": "Request context not available"}) @@ -283,19 +296,28 @@ async def edit_image( images = await image_edits( request=__request__, - form_data=EditImageForm(prompt=prompt, image=image_url), + form_data=EditImageForm(prompt=prompt, image=image_urls), user=user, ) + # Prepare file entries for the images + image_files = [{"type": "image", "url": img["url"]} for img in images] + + # Persist files to DB if chat context is available + if __chat_id__ and __message_id__ and images: + image_files = Chats.add_message_files_by_id_and_message_id( + __chat_id__, + __message_id__, + image_files, + ) + # Emit the images to the UI if event emitter is available - if __event_emitter__ and images: + if __event_emitter__ and image_files: await __event_emitter__( { - "type": "files", + "type": "chat:message:files", "data": { - "files": [ - {"type": "image", "url": img["url"]} for img in images - ] + "files": image_files, }, } ) diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index b7b7dc18b..c379467fc 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -1198,6 +1198,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): "__oauth_token__": oauth_token, "__request__": request, "__model__": model, + "__chat_id__": metadata.get("chat_id"), + "__message_id__": metadata.get("message_id"), } # Initialize events to store additional event to be sent to the client # Initialize contexts and citation diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 8793feae3..4027a0a83 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -402,6 +402,8 @@ def get_builtin_tools( "__request__": request, "__user__": extra_params.get("__user__", {}), "__event_emitter__": extra_params.get("__event_emitter__"), + "__chat_id__": extra_params.get("__chat_id__"), + "__message_id__": extra_params.get("__message_id__"), }, )