diff --git a/backend/open_webui/routers/files.py b/backend/open_webui/routers/files.py index 8f1ee13f7..70e0f468f 100644 --- a/backend/open_webui/routers/files.py +++ b/backend/open_webui/routers/files.py @@ -22,7 +22,7 @@ from fastapi import ( from fastapi.responses import FileResponse, StreamingResponse from sqlalchemy.ext.asyncio import AsyncSession -from open_webui.internal.db import get_async_session, SessionLocal +from open_webui.internal.db import get_async_session, get_async_db_context from open_webui.constants import ERROR_MESSAGES from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT @@ -113,7 +113,7 @@ async def process_uploaded_file( file_path_processed = Storage.get_file(file_path) result = transcribe(request, file_path_processed, file_metadata, user) - process_file( + await process_file( request, ProcessFileForm(file_id=file_item.id, content=result.get('text', '')), user=user, @@ -122,7 +122,7 @@ async def process_uploaded_file( elif (not content_type.startswith(('image/', 'video/'))) or ( request.app.state.config.CONTENT_EXTRACTION_ENGINE == 'external' ): - process_file( + await process_file( request, ProcessFileForm(file_id=file_item.id), user=user, @@ -132,7 +132,7 @@ async def process_uploaded_file( raise Exception(f'File type {content_type} is not supported for processing') else: log.info(f'File type {file.content_type} is not provided, but trying to process anyway') - process_file( + await process_file( request, ProcessFileForm(file_id=file_item.id), user=user, @@ -151,10 +151,10 @@ async def process_uploaded_file( ) if db: - _process_handler(db) + await _process_handler(db) else: - with SessionLocal() as db_session: - _process_handler(db_session) + async with get_async_db_context() as db_session: + await _process_handler(db_session) @router.post('/', response_model=FileModelResponse) @@ -540,7 +540,7 @@ async def update_file_data_content_by_id( if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'write', user, db=db): try: - process_file( + await process_file( request, ProcessFileForm(file_id=id, content=form_data.content), user=user, @@ -560,7 +560,7 @@ async def update_file_data_content_by_id( # Remove old embeddings for this file from the KB collection VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': id}) # Re-add from the now-updated file-{file_id} collection - process_file( + await process_file( request, ProcessFileForm(file_id=id, collection_name=knowledge.id), user=user, diff --git a/backend/open_webui/routers/images.py b/backend/open_webui/routers/images.py index fb3bc1cec..0d534db7f 100644 --- a/backend/open_webui/routers/images.py +++ b/backend/open_webui/routers/images.py @@ -832,7 +832,7 @@ async def image_edits( except Exception as e: raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(e)) - async def get_image_file_item(base64_string, param_name='image'): + def get_image_file_item(base64_string, param_name='image'): data = base64_string header, encoded = data.split(',', 1) mime_type = header.split(';')[0].lstrip('data:') diff --git a/backend/open_webui/routers/knowledge.py b/backend/open_webui/routers/knowledge.py index f6c3416c8..3022763b4 100644 --- a/backend/open_webui/routers/knowledge.py +++ b/backend/open_webui/routers/knowledge.py @@ -2,7 +2,7 @@ from typing import List, Optional from pydantic import BaseModel from fastapi import APIRouter, Depends, HTTPException, status, Request, Query from fastapi.responses import StreamingResponse -from fastapi.concurrency import run_in_threadpool + import logging import io import zipfile @@ -319,8 +319,7 @@ async def reindex_knowledge_files( failed_files = [] for file in files: try: - await run_in_threadpool( - process_file, + await process_file( request, ProcessFileForm(file_id=file.id, collection_name=knowledge_base.id), user=user, @@ -543,7 +542,7 @@ async def update_knowledge_access_by_id( await AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db) return KnowledgeFilesResponse( - **await Knowledges.get_knowledge_by_id(id=id, db=db).model_dump(), + **(await Knowledges.get_knowledge_by_id(id=id, db=db)).model_dump(), files=await Knowledges.get_file_metadatas_by_id(id, db=db), ) @@ -659,7 +658,7 @@ async def add_file_to_knowledge_by_id( # Add content to the vector database try: - process_file( + await process_file( request, ProcessFileForm(file_id=form_data.file_id, collection_name=id), user=user, @@ -737,7 +736,7 @@ async def update_file_from_knowledge_by_id( # Add content to the vector database try: - process_file( + await process_file( request, ProcessFileForm(file_id=form_data.file_id, collection_name=id), user=user, @@ -962,7 +961,7 @@ async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: As log.debug(e) pass - knowledge = Knowledges.reset_knowledge_by_id(id=id, db=db) + knowledge = await Knowledges.reset_knowledge_by_id(id=id, db=db) return knowledge diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index 58af93437..25d1f8441 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -1213,7 +1213,7 @@ async def search_channel_messages( end_ts = end_timestamp * 1_000_000_000 if end_timestamp else None # Search messages using the model method - matching_messages = Messages.search_messages_by_channel_ids( + matching_messages = await Messages.search_messages_by_channel_ids( channel_ids=channel_ids, query=query, start_timestamp=start_ts, @@ -1274,7 +1274,7 @@ async def view_channel_message( try: user_id = __user__.get('id') - message = Messages.get_message_by_id(message_id) + message = await Messages.get_message_by_id(message_id) if not message: return json.dumps({'error': 'Message not found'}) @@ -1336,7 +1336,7 @@ async def view_channel_thread( user_id = __user__.get('id') # Get the parent message - parent_message = Messages.get_message_by_id(parent_message_id) + parent_message = await Messages.get_message_by_id(parent_message_id) if not parent_message: return json.dumps({'error': 'Message not found'}) @@ -1353,7 +1353,7 @@ async def view_channel_thread( return json.dumps({'error': 'Access denied'}) # Get all thread replies - thread_replies = Messages.get_thread_replies_by_message_id(parent_message_id) + thread_replies = await Messages.get_thread_replies_by_message_id(parent_message_id) # Build the response messages = []