refac
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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:')
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
Reference in New Issue
Block a user