refac/enh: db session sharing
This commit is contained in:
@@ -22,6 +22,8 @@ from fastapi import (
|
||||
)
|
||||
|
||||
from fastapi.responses import FileResponse, StreamingResponse
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import get_session, SessionLocal
|
||||
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
|
||||
@@ -62,9 +64,12 @@ router = APIRouter()
|
||||
|
||||
# TODO: Optimize this function to use the knowledge_file table for faster lookups.
|
||||
def has_access_to_file(
|
||||
file_id: Optional[str], access_type: str, user=Depends(get_verified_user)
|
||||
file_id: Optional[str],
|
||||
access_type: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: Optional[Session] = None,
|
||||
) -> bool:
|
||||
file = Files.get_file_by_id(file_id)
|
||||
file = Files.get_file_by_id(file_id, db=db)
|
||||
log.debug(f"Checking if user has {access_type} access to file")
|
||||
if not file:
|
||||
raise HTTPException(
|
||||
@@ -73,31 +78,33 @@ def has_access_to_file(
|
||||
)
|
||||
|
||||
# Check if the file is associated with any knowledge bases the user has access to
|
||||
knowledge_bases = Knowledges.get_knowledges_by_file_id(file_id)
|
||||
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)}
|
||||
knowledge_bases = Knowledges.get_knowledges_by_file_id(file_id, db=db)
|
||||
user_group_ids = {
|
||||
group.id for group in Groups.get_groups_by_member_id(user.id, db=db)
|
||||
}
|
||||
for knowledge_base in knowledge_bases:
|
||||
if knowledge_base.user_id == user.id or has_access(
|
||||
user.id, access_type, knowledge_base.access_control, user_group_ids
|
||||
user.id, access_type, knowledge_base.access_control, user_group_ids, db=db
|
||||
):
|
||||
return True
|
||||
|
||||
knowledge_base_id = file.meta.get("collection_name") if file.meta else None
|
||||
if knowledge_base_id:
|
||||
knowledge_bases = Knowledges.get_knowledge_bases_by_user_id(
|
||||
user.id, access_type
|
||||
user.id, access_type, db=db
|
||||
)
|
||||
for knowledge_base in knowledge_bases:
|
||||
if knowledge_base.id == knowledge_base_id:
|
||||
return True
|
||||
|
||||
# Check if the file is associated with any channels the user has access to
|
||||
channels = Channels.get_channels_by_file_id_and_user_id(file_id, user.id)
|
||||
channels = Channels.get_channels_by_file_id_and_user_id(file_id, user.id, db=db)
|
||||
if access_type == "read" and channels:
|
||||
return True
|
||||
|
||||
# Check if the file is associated with any chats the user has access to
|
||||
# TODO: Granular access control for chats
|
||||
chats = Chats.get_shared_chats_by_file_id(file_id)
|
||||
chats = Chats.get_shared_chats_by_file_id(file_id, db=db)
|
||||
if chats:
|
||||
return True
|
||||
|
||||
@@ -109,47 +116,78 @@ def has_access_to_file(
|
||||
############################
|
||||
|
||||
|
||||
def process_uploaded_file(request, file, file_path, file_item, file_metadata, user):
|
||||
try:
|
||||
if file.content_type:
|
||||
stt_supported_content_types = getattr(
|
||||
request.app.state.config, "STT_SUPPORTED_CONTENT_TYPES", []
|
||||
)
|
||||
def process_uploaded_file(
|
||||
request,
|
||||
file,
|
||||
file_path,
|
||||
file_item,
|
||||
file_metadata,
|
||||
user,
|
||||
db: Optional[Session] = None,
|
||||
):
|
||||
def _process_handler(db_session):
|
||||
try:
|
||||
if file.content_type:
|
||||
stt_supported_content_types = getattr(
|
||||
request.app.state.config, "STT_SUPPORTED_CONTENT_TYPES", []
|
||||
)
|
||||
|
||||
if strict_match_mime_type(stt_supported_content_types, file.content_type):
|
||||
file_path = Storage.get_file(file_path)
|
||||
result = transcribe(request, file_path, file_metadata, user)
|
||||
if strict_match_mime_type(
|
||||
stt_supported_content_types, file.content_type
|
||||
):
|
||||
file_path_processed = Storage.get_file(file_path)
|
||||
result = transcribe(
|
||||
request, file_path_processed, file_metadata, user
|
||||
)
|
||||
|
||||
process_file(
|
||||
request,
|
||||
ProcessFileForm(
|
||||
file_id=file_item.id, content=result.get("text", "")
|
||||
),
|
||||
user=user,
|
||||
db=db_session,
|
||||
)
|
||||
elif (not file.content_type.startswith(("image/", "video/"))) or (
|
||||
request.app.state.config.CONTENT_EXTRACTION_ENGINE == "external"
|
||||
):
|
||||
process_file(
|
||||
request,
|
||||
ProcessFileForm(file_id=file_item.id),
|
||||
user=user,
|
||||
db=db_session,
|
||||
)
|
||||
else:
|
||||
raise Exception(
|
||||
f"File type {file.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(
|
||||
request,
|
||||
ProcessFileForm(
|
||||
file_id=file_item.id, content=result.get("text", "")
|
||||
),
|
||||
ProcessFileForm(file_id=file_item.id),
|
||||
user=user,
|
||||
db=db_session,
|
||||
)
|
||||
elif (not file.content_type.startswith(("image/", "video/"))) or (
|
||||
request.app.state.config.CONTENT_EXTRACTION_ENGINE == "external"
|
||||
):
|
||||
process_file(request, ProcessFileForm(file_id=file_item.id), user=user)
|
||||
else:
|
||||
raise Exception(
|
||||
f"File type {file.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(request, ProcessFileForm(file_id=file_item.id), user=user)
|
||||
|
||||
except Exception as e:
|
||||
log.error(f"Error processing file: {file_item.id}")
|
||||
Files.update_file_data_by_id(
|
||||
file_item.id,
|
||||
{
|
||||
"status": "failed",
|
||||
"error": str(e.detail) if hasattr(e, "detail") else str(e),
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
log.error(f"Error processing file: {file_item.id}")
|
||||
Files.update_file_data_by_id(
|
||||
file_item.id,
|
||||
{
|
||||
"status": "failed",
|
||||
"error": str(e.detail) if hasattr(e, "detail") else str(e),
|
||||
},
|
||||
db=db_session,
|
||||
)
|
||||
|
||||
if db:
|
||||
_process_handler(db)
|
||||
else:
|
||||
with SessionLocal() as db_session:
|
||||
_process_handler(db_session)
|
||||
|
||||
|
||||
@router.post("/", response_model=FileModelResponse)
|
||||
@@ -161,6 +199,7 @@ def upload_file(
|
||||
process: bool = Query(True),
|
||||
process_in_background: bool = Query(True),
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
return upload_file_handler(
|
||||
request,
|
||||
@@ -170,6 +209,7 @@ def upload_file(
|
||||
process_in_background=process_in_background,
|
||||
user=user,
|
||||
background_tasks=background_tasks,
|
||||
db=db,
|
||||
)
|
||||
|
||||
|
||||
@@ -181,6 +221,7 @@ def upload_file_handler(
|
||||
process_in_background: bool = Query(True),
|
||||
user=Depends(get_verified_user),
|
||||
background_tasks: Optional[BackgroundTasks] = None,
|
||||
db: Optional[Session] = None,
|
||||
):
|
||||
log.info(f"file.content_type: {file.content_type} {process}")
|
||||
|
||||
@@ -248,14 +289,17 @@ def upload_file_handler(
|
||||
},
|
||||
}
|
||||
),
|
||||
db=db,
|
||||
)
|
||||
|
||||
if "channel_id" in file_metadata:
|
||||
channel = Channels.get_channel_by_id_and_user_id(
|
||||
file_metadata["channel_id"], user.id
|
||||
file_metadata["channel_id"], user.id, db=db
|
||||
)
|
||||
if channel:
|
||||
Channels.add_file_to_channel_by_id(channel.id, file_item.id, user.id)
|
||||
Channels.add_file_to_channel_by_id(
|
||||
channel.id, file_item.id, user.id, db=db
|
||||
)
|
||||
|
||||
if process:
|
||||
if background_tasks and process_in_background:
|
||||
@@ -277,6 +321,7 @@ def upload_file_handler(
|
||||
file_item,
|
||||
file_metadata,
|
||||
user,
|
||||
db=db,
|
||||
)
|
||||
return {"status": True, **file_item.model_dump()}
|
||||
else:
|
||||
@@ -302,11 +347,15 @@ def upload_file_handler(
|
||||
|
||||
|
||||
@router.get("/", response_model=list[FileModelResponse])
|
||||
async def list_files(user=Depends(get_verified_user), content: bool = Query(True)):
|
||||
async def list_files(
|
||||
user=Depends(get_verified_user),
|
||||
content: bool = Query(True),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if user.role == "admin":
|
||||
files = Files.get_files()
|
||||
files = Files.get_files(db=db)
|
||||
else:
|
||||
files = Files.get_files_by_user_id(user.id)
|
||||
files = Files.get_files_by_user_id(user.id, db=db)
|
||||
|
||||
if not content:
|
||||
for file in files:
|
||||
@@ -329,15 +378,16 @@ async def search_files(
|
||||
),
|
||||
content: bool = Query(True),
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
"""
|
||||
Search for files by filename with support for wildcard patterns.
|
||||
"""
|
||||
# Get files according to user role
|
||||
if user.role == "admin":
|
||||
files = Files.get_files()
|
||||
files = Files.get_files(db=db)
|
||||
else:
|
||||
files = Files.get_files_by_user_id(user.id)
|
||||
files = Files.get_files_by_user_id(user.id, db=db)
|
||||
|
||||
# Get matching files
|
||||
matching_files = [
|
||||
@@ -364,8 +414,10 @@ async def search_files(
|
||||
|
||||
|
||||
@router.delete("/all")
|
||||
async def delete_all_files(user=Depends(get_admin_user)):
|
||||
result = Files.delete_all_files()
|
||||
async def delete_all_files(
|
||||
user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
result = Files.delete_all_files(db=db)
|
||||
if result:
|
||||
try:
|
||||
Storage.delete_all_files()
|
||||
@@ -391,8 +443,10 @@ async def delete_all_files(user=Depends(get_admin_user)):
|
||||
|
||||
|
||||
@router.get("/{id}", response_model=Optional[FileModel])
|
||||
async def get_file_by_id(id: str, user=Depends(get_verified_user)):
|
||||
file = Files.get_file_by_id(id)
|
||||
async def get_file_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
file = Files.get_file_by_id(id, db=db)
|
||||
|
||||
if not file:
|
||||
raise HTTPException(
|
||||
@@ -403,7 +457,7 @@ async def get_file_by_id(id: str, user=Depends(get_verified_user)):
|
||||
if (
|
||||
file.user_id == user.id
|
||||
or user.role == "admin"
|
||||
or has_access_to_file(id, "read", user)
|
||||
or has_access_to_file(id, "read", user, db=db)
|
||||
):
|
||||
return file
|
||||
else:
|
||||
@@ -415,9 +469,12 @@ async def get_file_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
@router.get("/{id}/process/status")
|
||||
async def get_file_process_status(
|
||||
id: str, stream: bool = Query(False), user=Depends(get_verified_user)
|
||||
id: str,
|
||||
stream: bool = Query(False),
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
file = Files.get_file_by_id(id)
|
||||
file = Files.get_file_by_id(id, db=db)
|
||||
|
||||
if not file:
|
||||
raise HTTPException(
|
||||
@@ -428,7 +485,7 @@ async def get_file_process_status(
|
||||
if (
|
||||
file.user_id == user.id
|
||||
or user.role == "admin"
|
||||
or has_access_to_file(id, "read", user)
|
||||
or has_access_to_file(id, "read", user, db=db)
|
||||
):
|
||||
if stream:
|
||||
MAX_FILE_PROCESSING_DURATION = 3600 * 2
|
||||
@@ -436,7 +493,7 @@ async def get_file_process_status(
|
||||
async def event_stream(file_item):
|
||||
if file_item:
|
||||
for _ in range(MAX_FILE_PROCESSING_DURATION):
|
||||
file_item = Files.get_file_by_id(file_item.id)
|
||||
file_item = Files.get_file_by_id(file_item.id, db=db)
|
||||
if file_item:
|
||||
data = file_item.model_dump().get("data", {})
|
||||
status = data.get("status")
|
||||
@@ -476,8 +533,10 @@ async def get_file_process_status(
|
||||
|
||||
|
||||
@router.get("/{id}/data/content")
|
||||
async def get_file_data_content_by_id(id: str, user=Depends(get_verified_user)):
|
||||
file = Files.get_file_by_id(id)
|
||||
async def get_file_data_content_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
file = Files.get_file_by_id(id, db=db)
|
||||
|
||||
if not file:
|
||||
raise HTTPException(
|
||||
@@ -488,7 +547,7 @@ async def get_file_data_content_by_id(id: str, user=Depends(get_verified_user)):
|
||||
if (
|
||||
file.user_id == user.id
|
||||
or user.role == "admin"
|
||||
or has_access_to_file(id, "read", user)
|
||||
or has_access_to_file(id, "read", user, db=db)
|
||||
):
|
||||
return {"content": file.data.get("content", "")}
|
||||
else:
|
||||
@@ -509,9 +568,13 @@ class ContentForm(BaseModel):
|
||||
|
||||
@router.post("/{id}/data/content/update")
|
||||
async def update_file_data_content_by_id(
|
||||
request: Request, id: str, form_data: ContentForm, user=Depends(get_verified_user)
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: ContentForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
file = Files.get_file_by_id(id)
|
||||
file = Files.get_file_by_id(id, db=db)
|
||||
|
||||
if not file:
|
||||
raise HTTPException(
|
||||
@@ -522,7 +585,7 @@ async def update_file_data_content_by_id(
|
||||
if (
|
||||
file.user_id == user.id
|
||||
or user.role == "admin"
|
||||
or has_access_to_file(id, "write", user)
|
||||
or has_access_to_file(id, "write", user, db=db)
|
||||
):
|
||||
try:
|
||||
process_file(
|
||||
@@ -530,7 +593,7 @@ async def update_file_data_content_by_id(
|
||||
ProcessFileForm(file_id=id, content=form_data.content),
|
||||
user=user,
|
||||
)
|
||||
file = Files.get_file_by_id(id=id)
|
||||
file = Files.get_file_by_id(id=id, db=db)
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
log.error(f"Error processing file: {file.id}")
|
||||
@@ -550,9 +613,12 @@ async def update_file_data_content_by_id(
|
||||
|
||||
@router.get("/{id}/content")
|
||||
async def get_file_content_by_id(
|
||||
id: str, user=Depends(get_verified_user), attachment: bool = Query(False)
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
attachment: bool = Query(False),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
file = Files.get_file_by_id(id)
|
||||
file = Files.get_file_by_id(id, db=db)
|
||||
|
||||
if not file:
|
||||
raise HTTPException(
|
||||
@@ -563,7 +629,7 @@ async def get_file_content_by_id(
|
||||
if (
|
||||
file.user_id == user.id
|
||||
or user.role == "admin"
|
||||
or has_access_to_file(id, "read", user)
|
||||
or has_access_to_file(id, "read", user, db=db)
|
||||
):
|
||||
try:
|
||||
file_path = Storage.get_file(file.path)
|
||||
@@ -619,8 +685,10 @@ async def get_file_content_by_id(
|
||||
|
||||
|
||||
@router.get("/{id}/content/html")
|
||||
async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user)):
|
||||
file = Files.get_file_by_id(id)
|
||||
async def get_html_file_content_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
file = Files.get_file_by_id(id, db=db)
|
||||
|
||||
if not file:
|
||||
raise HTTPException(
|
||||
@@ -628,7 +696,7 @@ async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user)):
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
file_user = Users.get_user_by_id(file.user_id)
|
||||
file_user = Users.get_user_by_id(file.user_id, db=db)
|
||||
if not file_user.role == "admin":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
@@ -638,7 +706,7 @@ async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user)):
|
||||
if (
|
||||
file.user_id == user.id
|
||||
or user.role == "admin"
|
||||
or has_access_to_file(id, "read", user)
|
||||
or has_access_to_file(id, "read", user, db=db)
|
||||
):
|
||||
try:
|
||||
file_path = Storage.get_file(file.path)
|
||||
@@ -668,8 +736,10 @@ async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.get("/{id}/content/{file_name}")
|
||||
async def get_file_content_by_id(id: str, user=Depends(get_verified_user)):
|
||||
file = Files.get_file_by_id(id)
|
||||
async def get_file_content_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
file = Files.get_file_by_id(id, db=db)
|
||||
|
||||
if not file:
|
||||
raise HTTPException(
|
||||
@@ -680,7 +750,7 @@ async def get_file_content_by_id(id: str, user=Depends(get_verified_user)):
|
||||
if (
|
||||
file.user_id == user.id
|
||||
or user.role == "admin"
|
||||
or has_access_to_file(id, "read", user)
|
||||
or has_access_to_file(id, "read", user, db=db)
|
||||
):
|
||||
file_path = file.path
|
||||
|
||||
@@ -730,8 +800,10 @@ async def get_file_content_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.delete("/{id}")
|
||||
async def delete_file_by_id(id: str, user=Depends(get_verified_user)):
|
||||
file = Files.get_file_by_id(id)
|
||||
async def delete_file_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
file = Files.get_file_by_id(id, db=db)
|
||||
|
||||
if not file:
|
||||
raise HTTPException(
|
||||
@@ -742,10 +814,10 @@ async def delete_file_by_id(id: str, user=Depends(get_verified_user)):
|
||||
if (
|
||||
file.user_id == user.id
|
||||
or user.role == "admin"
|
||||
or has_access_to_file(id, "write", user)
|
||||
or has_access_to_file(id, "write", user, db=db)
|
||||
):
|
||||
|
||||
result = Files.delete_file_by_id(id)
|
||||
result = Files.delete_file_by_id(id, db=db)
|
||||
if result:
|
||||
try:
|
||||
Storage.delete_file(file.path)
|
||||
|
||||
Reference in New Issue
Block a user