refac/enh: db session sharing

This commit is contained in:
Timothy Jaeryang Baek
2025-12-29 00:21:18 +04:00
parent 6dd0f99b90
commit b1d0f00d8c
23 changed files with 1173 additions and 663 deletions
+234 -111
View File
@@ -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: