refac/enh: db session sharing
This commit is contained in:
+234
-111
@@ -1,6 +1,7 @@
|
||||
import json
|
||||
import logging
|
||||
from typing import Optional
|
||||
from sqlalchemy.orm import Session
|
||||
import asyncio
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
@@ -23,6 +24,7 @@ from open_webui.models.chats import (
|
||||
)
|
||||
from open_webui.models.tags import TagModel, Tags
|
||||
from open_webui.models.folders import Folders
|
||||
from open_webui.internal.db import get_session
|
||||
|
||||
from open_webui.config import ENABLE_ADMIN_CHAT_ACCESS, ENABLE_ADMIN_EXPORT
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
@@ -49,6 +51,7 @@ def get_session_user_chat_list(
|
||||
page: Optional[int] = None,
|
||||
include_pinned: Optional[bool] = False,
|
||||
include_folders: Optional[bool] = False,
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
try:
|
||||
if page is not None:
|
||||
@@ -61,10 +64,14 @@ def get_session_user_chat_list(
|
||||
include_pinned=include_pinned,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
db=db,
|
||||
)
|
||||
else:
|
||||
return Chats.get_chat_title_id_list_by_user_id(
|
||||
user.id, include_folders=include_folders, include_pinned=include_pinned
|
||||
user.id,
|
||||
include_folders=include_folders,
|
||||
include_pinned=include_pinned,
|
||||
db=db,
|
||||
)
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -84,12 +91,13 @@ def get_session_user_chat_usage_stats(
|
||||
items_per_page: Optional[int] = 50,
|
||||
page: Optional[int] = 1,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
try:
|
||||
limit = items_per_page
|
||||
skip = (page - 1) * limit
|
||||
|
||||
result = Chats.get_chats_by_user_id(user.id, skip=skip, limit=limit)
|
||||
result = Chats.get_chats_by_user_id(user.id, skip=skip, limit=limit, db=db)
|
||||
|
||||
chats = result.items
|
||||
total = result.total
|
||||
@@ -216,6 +224,7 @@ class ChatStatsExportList(BaseModel):
|
||||
|
||||
def _process_chat_for_export(chat) -> Optional[ChatStatsExport]:
|
||||
try:
|
||||
|
||||
def get_message_content_length(message):
|
||||
content = message.get("content", "")
|
||||
if isinstance(content, str):
|
||||
@@ -348,7 +357,9 @@ def _process_chat_for_export(chat) -> Optional[ChatStatsExport]:
|
||||
return None
|
||||
|
||||
|
||||
def calculate_chat_stats(user_id, skip=0, limit=10, filter=None):
|
||||
def calculate_chat_stats(
|
||||
user_id, skip=0, limit=10, filter=None, db: Optional[Session] = None
|
||||
):
|
||||
if filter is None:
|
||||
filter = {}
|
||||
|
||||
@@ -357,6 +368,7 @@ def calculate_chat_stats(user_id, skip=0, limit=10, filter=None):
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
filter=filter,
|
||||
db=db,
|
||||
)
|
||||
|
||||
chat_stats_export_list = []
|
||||
@@ -368,14 +380,21 @@ def calculate_chat_stats(user_id, skip=0, limit=10, filter=None):
|
||||
return chat_stats_export_list, result.total
|
||||
|
||||
|
||||
async def generate_chat_stats_jsonl_generator(user_id, filter):
|
||||
async def generate_chat_stats_jsonl_generator(
|
||||
user_id, filter, db: Optional[Session] = None
|
||||
):
|
||||
skip = 0
|
||||
limit = CHAT_EXPORT_PAGE_ITEM_COUNT
|
||||
|
||||
while True:
|
||||
# Use asyncio.to_thread to make the blocking DB call non-blocking
|
||||
result = await asyncio.to_thread(
|
||||
Chats.get_chats_by_user_id, user_id, filter=filter, skip=skip, limit=limit
|
||||
Chats.get_chats_by_user_id,
|
||||
user_id,
|
||||
filter=filter,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
db=db,
|
||||
)
|
||||
if not result.items:
|
||||
break
|
||||
@@ -386,7 +405,7 @@ async def generate_chat_stats_jsonl_generator(user_id, filter):
|
||||
if chat_stat:
|
||||
yield chat_stat.model_dump_json() + "\n"
|
||||
except Exception as e:
|
||||
log.exception(f"Error processing chat {chat.id}: {e}")
|
||||
log.exception(f"Error processing chat {chat.id}: {e}")
|
||||
|
||||
skip += limit
|
||||
|
||||
@@ -400,6 +419,7 @@ async def export_chat_stats(
|
||||
page: Optional[int] = 1,
|
||||
stream: bool = False,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
# Check if the user has permission to share/export chats
|
||||
if (user.role != "admin") and (
|
||||
@@ -415,7 +435,7 @@ async def export_chat_stats(
|
||||
filter = {"order_by": "created_at", "direction": "asc"}
|
||||
|
||||
if chat_id:
|
||||
chat = Chats.get_chat_by_id(chat_id)
|
||||
chat = Chats.get_chat_by_id(chat_id, db=db)
|
||||
if chat:
|
||||
filter["start_time"] = chat.created_at
|
||||
|
||||
@@ -426,7 +446,7 @@ async def export_chat_stats(
|
||||
|
||||
if stream:
|
||||
return StreamingResponse(
|
||||
generate_chat_stats_jsonl_generator(user.id, filter),
|
||||
generate_chat_stats_jsonl_generator(user.id, filter, db=db),
|
||||
media_type="application/x-ndjson",
|
||||
headers={
|
||||
"Content-Disposition": f"attachment; filename=chat-stats-export-{user.id}.jsonl"
|
||||
@@ -437,7 +457,7 @@ async def export_chat_stats(
|
||||
skip = (page - 1) * limit
|
||||
|
||||
chat_stats_export_list, total = await asyncio.to_thread(
|
||||
calculate_chat_stats, user.id, skip, limit, filter
|
||||
calculate_chat_stats, user.id, skip, limit, filter, db=db
|
||||
)
|
||||
|
||||
return ChatStatsExportList(
|
||||
@@ -452,7 +472,11 @@ async def export_chat_stats(
|
||||
|
||||
|
||||
@router.delete("/", response_model=bool)
|
||||
async def delete_all_user_chats(request: Request, user=Depends(get_verified_user)):
|
||||
async def delete_all_user_chats(
|
||||
request: Request,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
|
||||
if user.role == "user" and not has_permission(
|
||||
user.id, "chat.delete", request.app.state.config.USER_PERMISSIONS
|
||||
@@ -462,7 +486,7 @@ async def delete_all_user_chats(request: Request, user=Depends(get_verified_user
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
result = Chats.delete_chats_by_user_id(user.id)
|
||||
result = Chats.delete_chats_by_user_id(user.id, db=db)
|
||||
return result
|
||||
|
||||
|
||||
@@ -479,6 +503,7 @@ async def get_user_chat_list_by_user_id(
|
||||
order_by: Optional[str] = None,
|
||||
direction: Optional[str] = None,
|
||||
user=Depends(get_admin_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if not ENABLE_ADMIN_CHAT_ACCESS:
|
||||
raise HTTPException(
|
||||
@@ -501,7 +526,7 @@ async def get_user_chat_list_by_user_id(
|
||||
filter["direction"] = direction
|
||||
|
||||
return Chats.get_chat_list_by_user_id(
|
||||
user_id, include_archived=True, filter=filter, skip=skip, limit=limit
|
||||
user_id, include_archived=True, filter=filter, skip=skip, limit=limit, db=db
|
||||
)
|
||||
|
||||
|
||||
@@ -511,9 +536,13 @@ async def get_user_chat_list_by_user_id(
|
||||
|
||||
|
||||
@router.post("/new", response_model=Optional[ChatResponse])
|
||||
async def create_new_chat(form_data: ChatForm, user=Depends(get_verified_user)):
|
||||
async def create_new_chat(
|
||||
form_data: ChatForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
try:
|
||||
chat = Chats.insert_new_chat(user.id, form_data)
|
||||
chat = Chats.insert_new_chat(user.id, form_data, db=db)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -528,9 +557,13 @@ async def create_new_chat(form_data: ChatForm, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.post("/import", response_model=list[ChatResponse])
|
||||
async def import_chats(form_data: ChatsImportForm, user=Depends(get_verified_user)):
|
||||
async def import_chats(
|
||||
form_data: ChatsImportForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
try:
|
||||
chats = Chats.import_chats(user.id, form_data.chats)
|
||||
chats = Chats.import_chats(user.id, form_data.chats, db=db)
|
||||
return chats
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -546,7 +579,10 @@ async def import_chats(form_data: ChatsImportForm, user=Depends(get_verified_use
|
||||
|
||||
@router.get("/search", response_model=list[ChatTitleIdResponse])
|
||||
def search_user_chats(
|
||||
text: str, page: Optional[int] = None, user=Depends(get_verified_user)
|
||||
text: str,
|
||||
page: Optional[int] = None,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if page is None:
|
||||
page = 1
|
||||
@@ -557,7 +593,7 @@ def search_user_chats(
|
||||
chat_list = [
|
||||
ChatTitleIdResponse(**chat.model_dump())
|
||||
for chat in Chats.get_chats_by_user_id_and_search_text(
|
||||
user.id, text, skip=skip, limit=limit
|
||||
user.id, text, skip=skip, limit=limit, db=db
|
||||
)
|
||||
]
|
||||
|
||||
@@ -566,9 +602,9 @@ def search_user_chats(
|
||||
if page == 1 and len(words) == 1 and words[0].startswith("tag:"):
|
||||
tag_id = words[0].replace("tag:", "")
|
||||
if len(chat_list) == 0:
|
||||
if Tags.get_tag_by_name_and_user_id(tag_id, user.id):
|
||||
if Tags.get_tag_by_name_and_user_id(tag_id, user.id, db=db):
|
||||
log.debug(f"deleting tag: {tag_id}")
|
||||
Tags.delete_tag_by_name_and_user_id(tag_id, user.id)
|
||||
Tags.delete_tag_by_name_and_user_id(tag_id, user.id, db=db)
|
||||
|
||||
return chat_list
|
||||
|
||||
@@ -579,23 +615,30 @@ def search_user_chats(
|
||||
|
||||
|
||||
@router.get("/folder/{folder_id}", response_model=list[ChatResponse])
|
||||
async def get_chats_by_folder_id(folder_id: str, user=Depends(get_verified_user)):
|
||||
async def get_chats_by_folder_id(
|
||||
folder_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
folder_ids = [folder_id]
|
||||
children_folders = Folders.get_children_folders_by_id_and_user_id(
|
||||
folder_id, user.id
|
||||
folder_id, user.id, db=db
|
||||
)
|
||||
if children_folders:
|
||||
folder_ids.extend([folder.id for folder in children_folders])
|
||||
|
||||
return [
|
||||
ChatResponse(**chat.model_dump())
|
||||
for chat in Chats.get_chats_by_folder_ids_and_user_id(folder_ids, user.id)
|
||||
for chat in Chats.get_chats_by_folder_ids_and_user_id(
|
||||
folder_ids, user.id, db=db
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
@router.get("/folder/{folder_id}/list")
|
||||
async def get_chat_list_by_folder_id(
|
||||
folder_id: str, page: Optional[int] = 1, user=Depends(get_verified_user)
|
||||
folder_id: str,
|
||||
page: Optional[int] = 1,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
try:
|
||||
limit = 10
|
||||
@@ -604,7 +647,7 @@ async def get_chat_list_by_folder_id(
|
||||
return [
|
||||
{"title": chat.title, "id": chat.id, "updated_at": chat.updated_at}
|
||||
for chat in Chats.get_chats_by_folder_id_and_user_id(
|
||||
folder_id, user.id, skip=skip, limit=limit
|
||||
folder_id, user.id, skip=skip, limit=limit, db=db
|
||||
)
|
||||
]
|
||||
|
||||
@@ -621,10 +664,12 @@ async def get_chat_list_by_folder_id(
|
||||
|
||||
|
||||
@router.get("/pinned", response_model=list[ChatTitleIdResponse])
|
||||
async def get_user_pinned_chats(user=Depends(get_verified_user)):
|
||||
async def get_user_pinned_chats(
|
||||
user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
return [
|
||||
ChatTitleIdResponse(**chat.model_dump())
|
||||
for chat in Chats.get_pinned_chats_by_user_id(user.id)
|
||||
for chat in Chats.get_pinned_chats_by_user_id(user.id, db=db)
|
||||
]
|
||||
|
||||
|
||||
@@ -634,10 +679,12 @@ async def get_user_pinned_chats(user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.get("/all", response_model=list[ChatResponse])
|
||||
async def get_user_chats(user=Depends(get_verified_user)):
|
||||
async def get_user_chats(
|
||||
user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
return [
|
||||
ChatResponse(**chat.model_dump())
|
||||
for chat in Chats.get_chats_by_user_id(user.id)
|
||||
for chat in Chats.get_chats_by_user_id(user.id, db=db)
|
||||
]
|
||||
|
||||
|
||||
@@ -647,10 +694,12 @@ async def get_user_chats(user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.get("/all/archived", response_model=list[ChatResponse])
|
||||
async def get_user_archived_chats(user=Depends(get_verified_user)):
|
||||
async def get_user_archived_chats(
|
||||
user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
return [
|
||||
ChatResponse(**chat.model_dump())
|
||||
for chat in Chats.get_archived_chats_by_user_id(user.id)
|
||||
for chat in Chats.get_archived_chats_by_user_id(user.id, db=db)
|
||||
]
|
||||
|
||||
|
||||
@@ -660,9 +709,11 @@ async def get_user_archived_chats(user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.get("/all/tags", response_model=list[TagModel])
|
||||
async def get_all_user_tags(user=Depends(get_verified_user)):
|
||||
async def get_all_user_tags(
|
||||
user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
try:
|
||||
tags = Tags.get_tags_by_user_id(user.id)
|
||||
tags = Tags.get_tags_by_user_id(user.id, db=db)
|
||||
return tags
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -677,13 +728,15 @@ async def get_all_user_tags(user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.get("/all/db", response_model=list[ChatResponse])
|
||||
async def get_all_user_chats_in_db(user=Depends(get_admin_user)):
|
||||
async def get_all_user_chats_in_db(
|
||||
user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
if not ENABLE_ADMIN_EXPORT:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
return [ChatResponse(**chat.model_dump()) for chat in Chats.get_chats()]
|
||||
return [ChatResponse(**chat.model_dump()) for chat in Chats.get_chats(db=db)]
|
||||
|
||||
|
||||
############################
|
||||
@@ -698,6 +751,7 @@ async def get_archived_session_user_chat_list(
|
||||
order_by: Optional[str] = None,
|
||||
direction: Optional[str] = None,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if page is None:
|
||||
page = 1
|
||||
@@ -720,6 +774,7 @@ async def get_archived_session_user_chat_list(
|
||||
filter=filter,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
db=db,
|
||||
)
|
||||
]
|
||||
|
||||
@@ -732,8 +787,10 @@ async def get_archived_session_user_chat_list(
|
||||
|
||||
|
||||
@router.post("/archive/all", response_model=bool)
|
||||
async def archive_all_chats(user=Depends(get_verified_user)):
|
||||
return Chats.archive_all_chats_by_user_id(user.id)
|
||||
async def archive_all_chats(
|
||||
user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
return Chats.archive_all_chats_by_user_id(user.id, db=db)
|
||||
|
||||
|
||||
############################
|
||||
@@ -742,8 +799,10 @@ async def archive_all_chats(user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.post("/unarchive/all", response_model=bool)
|
||||
async def unarchive_all_chats(user=Depends(get_verified_user)):
|
||||
return Chats.unarchive_all_chats_by_user_id(user.id)
|
||||
async def unarchive_all_chats(
|
||||
user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
return Chats.unarchive_all_chats_by_user_id(user.id, db=db)
|
||||
|
||||
|
||||
############################
|
||||
@@ -752,16 +811,18 @@ async def unarchive_all_chats(user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.get("/share/{share_id}", response_model=Optional[ChatResponse])
|
||||
async def get_shared_chat_by_id(share_id: str, user=Depends(get_verified_user)):
|
||||
async def get_shared_chat_by_id(
|
||||
share_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
if user.role == "pending":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND
|
||||
)
|
||||
|
||||
if user.role == "user" or (user.role == "admin" and not ENABLE_ADMIN_CHAT_ACCESS):
|
||||
chat = Chats.get_chat_by_share_id(share_id)
|
||||
chat = Chats.get_chat_by_share_id(share_id, db=db)
|
||||
elif user.role == "admin" and ENABLE_ADMIN_CHAT_ACCESS:
|
||||
chat = Chats.get_chat_by_id(share_id)
|
||||
chat = Chats.get_chat_by_id(share_id, db=db)
|
||||
|
||||
if chat:
|
||||
return ChatResponse(**chat.model_dump())
|
||||
@@ -788,13 +849,15 @@ class TagFilterForm(TagForm):
|
||||
|
||||
@router.post("/tags", response_model=list[ChatTitleIdResponse])
|
||||
async def get_user_chat_list_by_tag_name(
|
||||
form_data: TagFilterForm, user=Depends(get_verified_user)
|
||||
form_data: TagFilterForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
chats = Chats.get_chat_list_by_user_id_and_tag_name(
|
||||
user.id, form_data.name, form_data.skip, form_data.limit
|
||||
user.id, form_data.name, form_data.skip, form_data.limit, db=db
|
||||
)
|
||||
if len(chats) == 0:
|
||||
Tags.delete_tag_by_name_and_user_id(form_data.name, user.id)
|
||||
Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db)
|
||||
|
||||
return chats
|
||||
|
||||
@@ -805,8 +868,10 @@ async def get_user_chat_list_by_tag_name(
|
||||
|
||||
|
||||
@router.get("/{id}", response_model=Optional[ChatResponse])
|
||||
async def get_chat_by_id(id: str, user=Depends(get_verified_user)):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
async def get_chat_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
|
||||
if chat:
|
||||
return ChatResponse(**chat.model_dump())
|
||||
@@ -824,12 +889,15 @@ async def get_chat_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
@router.post("/{id}", response_model=Optional[ChatResponse])
|
||||
async def update_chat_by_id(
|
||||
id: str, form_data: ChatForm, user=Depends(get_verified_user)
|
||||
id: str,
|
||||
form_data: ChatForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
updated_chat = {**chat.chat, **form_data.chat}
|
||||
chat = Chats.update_chat_by_id(id, updated_chat)
|
||||
chat = Chats.update_chat_by_id(id, updated_chat, db=db)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -847,9 +915,13 @@ class MessageForm(BaseModel):
|
||||
|
||||
@router.post("/{id}/messages/{message_id}", response_model=Optional[ChatResponse])
|
||||
async def update_chat_message_by_id(
|
||||
id: str, message_id: str, form_data: MessageForm, user=Depends(get_verified_user)
|
||||
id: str,
|
||||
message_id: str,
|
||||
form_data: MessageForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
chat = Chats.get_chat_by_id(id)
|
||||
chat = Chats.get_chat_by_id(id, db=db)
|
||||
|
||||
if not chat:
|
||||
raise HTTPException(
|
||||
@@ -869,6 +941,7 @@ async def update_chat_message_by_id(
|
||||
{
|
||||
"content": form_data.content,
|
||||
},
|
||||
db=db,
|
||||
)
|
||||
|
||||
event_emitter = get_event_emitter(
|
||||
@@ -905,9 +978,13 @@ class EventForm(BaseModel):
|
||||
|
||||
@router.post("/{id}/messages/{message_id}/event", response_model=Optional[bool])
|
||||
async def send_chat_message_event_by_id(
|
||||
id: str, message_id: str, form_data: EventForm, user=Depends(get_verified_user)
|
||||
id: str,
|
||||
message_id: str,
|
||||
form_data: EventForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
chat = Chats.get_chat_by_id(id)
|
||||
chat = Chats.get_chat_by_id(id, db=db)
|
||||
|
||||
if not chat:
|
||||
raise HTTPException(
|
||||
@@ -945,14 +1022,19 @@ async def send_chat_message_event_by_id(
|
||||
|
||||
|
||||
@router.delete("/{id}", response_model=bool)
|
||||
async def delete_chat_by_id(request: Request, id: str, user=Depends(get_verified_user)):
|
||||
async def delete_chat_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if user.role == "admin":
|
||||
chat = Chats.get_chat_by_id(id)
|
||||
chat = Chats.get_chat_by_id(id, db=db)
|
||||
for tag in chat.meta.get("tags", []):
|
||||
if Chats.count_chats_by_tag_name_and_user_id(tag, user.id) == 1:
|
||||
Tags.delete_tag_by_name_and_user_id(tag, user.id)
|
||||
if Chats.count_chats_by_tag_name_and_user_id(tag, user.id, db=db) == 1:
|
||||
Tags.delete_tag_by_name_and_user_id(tag, user.id, db=db)
|
||||
|
||||
result = Chats.delete_chat_by_id(id)
|
||||
result = Chats.delete_chat_by_id(id, db=db)
|
||||
|
||||
return result
|
||||
else:
|
||||
@@ -964,12 +1046,12 @@ async def delete_chat_by_id(request: Request, id: str, user=Depends(get_verified
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
chat = Chats.get_chat_by_id(id)
|
||||
chat = Chats.get_chat_by_id(id, db=db)
|
||||
for tag in chat.meta.get("tags", []):
|
||||
if Chats.count_chats_by_tag_name_and_user_id(tag, user.id) == 1:
|
||||
Tags.delete_tag_by_name_and_user_id(tag, user.id)
|
||||
if Chats.count_chats_by_tag_name_and_user_id(tag, user.id, db=db) == 1:
|
||||
Tags.delete_tag_by_name_and_user_id(tag, user.id, db=db)
|
||||
|
||||
result = Chats.delete_chat_by_id_and_user_id(id, user.id)
|
||||
result = Chats.delete_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
return result
|
||||
|
||||
|
||||
@@ -979,8 +1061,10 @@ async def delete_chat_by_id(request: Request, id: str, user=Depends(get_verified
|
||||
|
||||
|
||||
@router.get("/{id}/pinned", response_model=Optional[bool])
|
||||
async def get_pinned_status_by_id(id: str, user=Depends(get_verified_user)):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
async def get_pinned_status_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
return chat.pinned
|
||||
else:
|
||||
@@ -995,10 +1079,12 @@ async def get_pinned_status_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.post("/{id}/pin", response_model=Optional[ChatResponse])
|
||||
async def pin_chat_by_id(id: str, user=Depends(get_verified_user)):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
async def pin_chat_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
chat = Chats.toggle_chat_pinned_by_id(id)
|
||||
chat = Chats.toggle_chat_pinned_by_id(id, db=db)
|
||||
return chat
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -1017,9 +1103,12 @@ class CloneForm(BaseModel):
|
||||
|
||||
@router.post("/{id}/clone", response_model=Optional[ChatResponse])
|
||||
async def clone_chat_by_id(
|
||||
form_data: CloneForm, id: str, user=Depends(get_verified_user)
|
||||
form_data: CloneForm,
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
updated_chat = {
|
||||
**chat.chat,
|
||||
@@ -1040,6 +1129,7 @@ async def clone_chat_by_id(
|
||||
}
|
||||
)
|
||||
],
|
||||
db=db,
|
||||
)
|
||||
|
||||
if chats:
|
||||
@@ -1062,12 +1152,14 @@ async def clone_chat_by_id(
|
||||
|
||||
|
||||
@router.post("/{id}/clone/shared", response_model=Optional[ChatResponse])
|
||||
async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user)):
|
||||
async def clone_shared_chat_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
|
||||
if user.role == "admin":
|
||||
chat = Chats.get_chat_by_id(id)
|
||||
chat = Chats.get_chat_by_id(id, db=db)
|
||||
else:
|
||||
chat = Chats.get_chat_by_share_id(id)
|
||||
chat = Chats.get_chat_by_share_id(id, db=db)
|
||||
|
||||
if chat:
|
||||
updated_chat = {
|
||||
@@ -1089,6 +1181,7 @@ async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user)):
|
||||
}
|
||||
)
|
||||
],
|
||||
db=db,
|
||||
)
|
||||
|
||||
if chats:
|
||||
@@ -1111,23 +1204,28 @@ async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.post("/{id}/archive", response_model=Optional[ChatResponse])
|
||||
async def archive_chat_by_id(id: str, user=Depends(get_verified_user)):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
async def archive_chat_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
chat = Chats.toggle_chat_archive_by_id(id)
|
||||
chat = Chats.toggle_chat_archive_by_id(id, db=db)
|
||||
|
||||
# Delete tags if chat is archived
|
||||
if chat.archived:
|
||||
for tag_id in chat.meta.get("tags", []):
|
||||
if Chats.count_chats_by_tag_name_and_user_id(tag_id, user.id) == 0:
|
||||
if (
|
||||
Chats.count_chats_by_tag_name_and_user_id(tag_id, user.id, db=db)
|
||||
== 0
|
||||
):
|
||||
log.debug(f"deleting tag: {tag_id}")
|
||||
Tags.delete_tag_by_name_and_user_id(tag_id, user.id)
|
||||
Tags.delete_tag_by_name_and_user_id(tag_id, user.id, db=db)
|
||||
else:
|
||||
for tag_id in chat.meta.get("tags", []):
|
||||
tag = Tags.get_tag_by_name_and_user_id(tag_id, user.id)
|
||||
tag = Tags.get_tag_by_name_and_user_id(tag_id, user.id, db=db)
|
||||
if tag is None:
|
||||
log.debug(f"inserting tag: {tag_id}")
|
||||
tag = Tags.insert_new_tag(tag_id, user.id)
|
||||
tag = Tags.insert_new_tag(tag_id, user.id, db=db)
|
||||
|
||||
return ChatResponse(**chat.model_dump())
|
||||
else:
|
||||
@@ -1142,7 +1240,12 @@ async def archive_chat_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.post("/{id}/share", response_model=Optional[ChatResponse])
|
||||
async def share_chat_by_id(request: Request, id: str, user=Depends(get_verified_user)):
|
||||
async def share_chat_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if (user.role != "admin") and (
|
||||
not has_permission(
|
||||
user.id, "chat.share", request.app.state.config.USER_PERMISSIONS
|
||||
@@ -1153,14 +1256,14 @@ async def share_chat_by_id(request: Request, id: str, user=Depends(get_verified_
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
|
||||
if chat:
|
||||
if chat.share_id:
|
||||
shared_chat = Chats.update_shared_chat_by_chat_id(chat.id)
|
||||
shared_chat = Chats.update_shared_chat_by_chat_id(chat.id, db=db)
|
||||
return ChatResponse(**shared_chat.model_dump())
|
||||
|
||||
shared_chat = Chats.insert_shared_chat_by_chat_id(chat.id)
|
||||
shared_chat = Chats.insert_shared_chat_by_chat_id(chat.id, db=db)
|
||||
if not shared_chat:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
@@ -1181,14 +1284,16 @@ async def share_chat_by_id(request: Request, id: str, user=Depends(get_verified_
|
||||
|
||||
|
||||
@router.delete("/{id}/share", response_model=Optional[bool])
|
||||
async def delete_shared_chat_by_id(id: str, user=Depends(get_verified_user)):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
async def delete_shared_chat_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
if not chat.share_id:
|
||||
return False
|
||||
|
||||
result = Chats.delete_shared_chat_by_chat_id(id)
|
||||
update_result = Chats.update_chat_share_id_by_id(id, None)
|
||||
result = Chats.delete_shared_chat_by_chat_id(id, db=db)
|
||||
update_result = Chats.update_chat_share_id_by_id(id, None, db=db)
|
||||
|
||||
return result and update_result != None
|
||||
else:
|
||||
@@ -1209,12 +1314,15 @@ class ChatFolderIdForm(BaseModel):
|
||||
|
||||
@router.post("/{id}/folder", response_model=Optional[ChatResponse])
|
||||
async def update_chat_folder_id_by_id(
|
||||
id: str, form_data: ChatFolderIdForm, user=Depends(get_verified_user)
|
||||
id: str,
|
||||
form_data: ChatFolderIdForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
chat = Chats.update_chat_folder_id_by_id_and_user_id(
|
||||
id, user.id, form_data.folder_id
|
||||
id, user.id, form_data.folder_id, db=db
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
else:
|
||||
@@ -1229,11 +1337,13 @@ async def update_chat_folder_id_by_id(
|
||||
|
||||
|
||||
@router.get("/{id}/tags", response_model=list[TagModel])
|
||||
async def get_chat_tags_by_id(id: str, user=Depends(get_verified_user)):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
async def get_chat_tags_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
tags = chat.meta.get("tags", [])
|
||||
return Tags.get_tags_by_ids_and_user_id(tags, user.id)
|
||||
return Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND
|
||||
@@ -1247,9 +1357,12 @@ async def get_chat_tags_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
@router.post("/{id}/tags", response_model=list[TagModel])
|
||||
async def add_tag_by_id_and_tag_name(
|
||||
id: str, form_data: TagForm, user=Depends(get_verified_user)
|
||||
id: str,
|
||||
form_data: TagForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
tags = chat.meta.get("tags", [])
|
||||
tag_id = form_data.name.replace(" ", "_").lower()
|
||||
@@ -1262,12 +1375,12 @@ async def add_tag_by_id_and_tag_name(
|
||||
|
||||
if tag_id not in tags:
|
||||
Chats.add_chat_tag_by_id_and_user_id_and_tag_name(
|
||||
id, user.id, form_data.name
|
||||
id, user.id, form_data.name, db=db
|
||||
)
|
||||
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
tags = chat.meta.get("tags", [])
|
||||
return Tags.get_tags_by_ids_and_user_id(tags, user.id)
|
||||
return Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT()
|
||||
@@ -1281,18 +1394,26 @@ async def add_tag_by_id_and_tag_name(
|
||||
|
||||
@router.delete("/{id}/tags", response_model=list[TagModel])
|
||||
async def delete_tag_by_id_and_tag_name(
|
||||
id: str, form_data: TagForm, user=Depends(get_verified_user)
|
||||
id: str,
|
||||
form_data: TagForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
Chats.delete_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name)
|
||||
Chats.delete_tag_by_id_and_user_id_and_tag_name(
|
||||
id, user.id, form_data.name, db=db
|
||||
)
|
||||
|
||||
if Chats.count_chats_by_tag_name_and_user_id(form_data.name, user.id) == 0:
|
||||
Tags.delete_tag_by_name_and_user_id(form_data.name, user.id)
|
||||
if (
|
||||
Chats.count_chats_by_tag_name_and_user_id(form_data.name, user.id, db=db)
|
||||
== 0
|
||||
):
|
||||
Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db)
|
||||
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
tags = chat.meta.get("tags", [])
|
||||
return Tags.get_tags_by_ids_and_user_id(tags, user.id)
|
||||
return Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND
|
||||
@@ -1305,14 +1426,16 @@ async def delete_tag_by_id_and_tag_name(
|
||||
|
||||
|
||||
@router.delete("/{id}/tags/all", response_model=Optional[bool])
|
||||
async def delete_all_tags_by_id(id: str, user=Depends(get_verified_user)):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id)
|
||||
async def delete_all_tags_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
Chats.delete_all_tags_by_id_and_user_id(id, user.id)
|
||||
Chats.delete_all_tags_by_id_and_user_id(id, user.id, db=db)
|
||||
|
||||
for tag in chat.meta.get("tags", []):
|
||||
if Chats.count_chats_by_tag_name_and_user_id(tag, user.id) == 0:
|
||||
Tags.delete_tag_by_name_and_user_id(tag, user.id)
|
||||
if Chats.count_chats_by_tag_name_and_user_id(tag, user.id, db=db) == 0:
|
||||
Tags.delete_tag_by_name_and_user_id(tag, user.id, db=db)
|
||||
|
||||
return True
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user