mirror of
https://github.com/open-webui/open-webui.git
synced 2026-09-28 02:07:05 -04:00
refac
This commit is contained in:
@@ -275,7 +275,8 @@ async def emit_chat_list_event(metadata: dict, chat_id: str):
|
||||
|
||||
event_emitter = await get_event_emitter(metadata, update_db=False)
|
||||
if event_emitter:
|
||||
await event_emitter({'type': 'chat:list', 'data': {'chat_id': chat_id}})
|
||||
folder_id = metadata.get('folder_id') or await Chats.get_chat_folder_id(chat_id, metadata.get('user_id'))
|
||||
await event_emitter({'type': 'chat:list', 'data': {'chat_id': chat_id, 'folder_id': folder_id}})
|
||||
|
||||
|
||||
class SPAStaticFiles(StaticFiles):
|
||||
@@ -1642,7 +1643,17 @@ async def chat_completion(
|
||||
event_emitter = await get_event_emitter(metadata, update_db=False)
|
||||
if event_emitter:
|
||||
try:
|
||||
await asyncio.shield(event_emitter({'type': 'chat:active', 'data': {'active': False}}))
|
||||
folder_id = metadata.get('folder_id') or await Chats.get_chat_folder_id(
|
||||
chat_id, user.id
|
||||
)
|
||||
await asyncio.shield(
|
||||
event_emitter(
|
||||
{
|
||||
'type': 'chat:active',
|
||||
'data': {'active': False, 'folder_id': folder_id},
|
||||
}
|
||||
)
|
||||
)
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
except Exception:
|
||||
@@ -1748,7 +1759,8 @@ async def chat_completion(
|
||||
update_db=False,
|
||||
)
|
||||
if event_emitter:
|
||||
await event_emitter({'type': 'chat:active', 'data': {'active': True}})
|
||||
folder_id = metadata.get('folder_id') or await Chats.get_chat_folder_id(chat_id, user.id)
|
||||
await event_emitter({'type': 'chat:active', 'data': {'active': True, 'folder_id': folder_id}})
|
||||
|
||||
return {
|
||||
'status': True,
|
||||
|
||||
@@ -36,13 +36,37 @@ from sqlalchemy import (
|
||||
)
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
from sqlalchemy.sql import exists
|
||||
from sqlalchemy.sql import case, exists
|
||||
from sqlalchemy.sql.expression import bindparam
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
ACTIVE_CHAT_GAP_SECONDS = 30 * 60
|
||||
|
||||
|
||||
def chat_list_order(sort_by: str = 'updated_at', sort_dir: str = 'desc', user_id: str | None = None):
|
||||
if sort_by != 'unread_updated_at':
|
||||
sort_column = Chat.title if sort_by == 'title' else Chat.updated_at
|
||||
order_clause = sort_column.asc() if sort_dir == 'asc' else sort_column.desc()
|
||||
return order_clause, Chat.id
|
||||
|
||||
unfinished_assistant = (
|
||||
select(ChatMessage.id)
|
||||
.where(ChatMessage.chat_id == Chat.id)
|
||||
.where(ChatMessage.role == 'assistant')
|
||||
.where(ChatMessage.done.is_(False))
|
||||
.exists()
|
||||
)
|
||||
conditions = [Chat.updated_at > func.coalesce(Chat.last_read_at, 0), ~unfinished_assistant]
|
||||
if user_id is not None:
|
||||
conditions.append(Chat.user_id == user_id)
|
||||
|
||||
unread = case(
|
||||
(and_(*conditions), 1),
|
||||
else_=0,
|
||||
)
|
||||
return unread.desc(), Chat.updated_at.desc(), Chat.id
|
||||
|
||||
|
||||
class Chat(Base): # database table mapping for chat entity
|
||||
__tablename__ = 'chat'
|
||||
|
||||
@@ -614,19 +638,62 @@ class ChatTable:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def update_chat_last_read_at_by_id(self, id: str, user_id: str, db: AsyncSession | None = None) -> int | None:
|
||||
async def update_chat_last_read_at_by_id(
|
||||
self, id: str, user_id: str, db: AsyncSession | None = None
|
||||
) -> tuple[int, bool] | None:
|
||||
try:
|
||||
async with get_async_db_context(db) as session:
|
||||
chat = await session.get(Chat, id)
|
||||
if chat and chat.user_id == user_id:
|
||||
last_read_at = int(time.time())
|
||||
was_unread = chat.last_read_at is None or chat.updated_at > chat.last_read_at
|
||||
chat.last_read_at = last_read_at
|
||||
await session.commit()
|
||||
return last_read_at
|
||||
return last_read_at, was_unread
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def mark_chat_unread_by_id(
|
||||
self, id: str, user_id: str, db: AsyncSession | None = None
|
||||
) -> ChatTitleIdResponse | None:
|
||||
try:
|
||||
async with get_async_db_context(db) as session:
|
||||
chat = await session.get(Chat, id)
|
||||
if chat and chat.user_id == user_id:
|
||||
chat.last_read_at = 0
|
||||
await session.commit()
|
||||
return ChatTitleIdResponse(
|
||||
id=chat.id,
|
||||
title=chat.title,
|
||||
updated_at=chat.updated_at,
|
||||
created_at=chat.created_at,
|
||||
last_read_at=chat.last_read_at,
|
||||
)
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def mark_chats_read_by_folder_ids(
|
||||
self, user_id: str, folder_ids: list[str], db: AsyncSession | None = None
|
||||
) -> int:
|
||||
if not folder_ids:
|
||||
return 0
|
||||
|
||||
async with get_async_db_context(db) as session:
|
||||
result = await session.execute(
|
||||
update(Chat)
|
||||
.where(
|
||||
Chat.user_id == user_id,
|
||||
Chat.folder_id.in_(folder_ids),
|
||||
Chat.archived == False,
|
||||
Chat.meta['internal'].as_boolean().is_not(True),
|
||||
)
|
||||
.values(last_read_at=Chat.updated_at)
|
||||
)
|
||||
await session.commit()
|
||||
return result.rowcount or 0
|
||||
|
||||
async def update_chat_title_by_id(self, id: str, title: str) -> ChatModel | None:
|
||||
try:
|
||||
async with get_async_db_context() as session:
|
||||
@@ -1212,6 +1279,8 @@ class ChatTable:
|
||||
include_archived: bool = False,
|
||||
include_folders: bool = False,
|
||||
include_pinned: bool = False,
|
||||
sort_by: str = 'updated_at',
|
||||
sort_dir: str = 'desc',
|
||||
skip: int | None = None,
|
||||
limit: int | None = None,
|
||||
db: AsyncSession | None = None,
|
||||
@@ -1231,7 +1300,7 @@ class ChatTable:
|
||||
if not include_archived:
|
||||
stmt = stmt.filter_by(archived=False)
|
||||
|
||||
stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id)
|
||||
stmt = stmt.order_by(*chat_list_order(sort_by, sort_dir))
|
||||
|
||||
if skip:
|
||||
stmt = stmt.offset(skip)
|
||||
@@ -1802,6 +1871,8 @@ class ChatTable:
|
||||
user_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 60,
|
||||
sort_by: str = 'updated_at',
|
||||
sort_dir: str = 'desc',
|
||||
db: AsyncSession | None = None,
|
||||
) -> list[ChatTitleIdResponse]:
|
||||
async with get_async_db_context(db) as session:
|
||||
@@ -1811,8 +1882,8 @@ class ChatTable:
|
||||
.filter(or_(Chat.pinned == False, Chat.pinned == None))
|
||||
.filter_by(archived=False)
|
||||
.where(Chat.meta['internal'].as_boolean().is_not(True))
|
||||
.order_by(Chat.updated_at.desc(), Chat.id)
|
||||
)
|
||||
stmt = stmt.order_by(*chat_list_order(sort_by, sort_dir))
|
||||
|
||||
if skip:
|
||||
stmt = stmt.offset(skip)
|
||||
@@ -1841,20 +1912,19 @@ class ChatTable:
|
||||
limit: int = 60,
|
||||
sort_by: str = 'updated_at',
|
||||
sort_dir: str = 'desc',
|
||||
unread_for_user_id: str | None = None,
|
||||
db: AsyncSession | None = None,
|
||||
) -> list[dict]:
|
||||
"""Get chats in a folder across ALL users. Returns dicts with user_id."""
|
||||
async with get_async_db_context(db) as session:
|
||||
sort_column = Chat.title if sort_by == 'title' else Chat.updated_at
|
||||
order_clause = sort_column.asc() if sort_dir == 'asc' else sort_column.desc()
|
||||
stmt = (
|
||||
select(Chat.id, Chat.title, Chat.user_id, Chat.updated_at, Chat.created_at, Chat.last_read_at)
|
||||
.filter_by(folder_id=folder_id)
|
||||
.filter(or_(Chat.pinned == False, Chat.pinned == None))
|
||||
.filter_by(archived=False)
|
||||
.where(Chat.meta['internal'].as_boolean().is_not(True))
|
||||
.order_by(order_clause, Chat.id)
|
||||
)
|
||||
stmt = stmt.order_by(*chat_list_order(sort_by, sort_dir, unread_for_user_id))
|
||||
|
||||
if skip:
|
||||
stmt = stmt.offset(skip)
|
||||
|
||||
@@ -115,6 +115,24 @@ async def add_active_state_to_chat_list(
|
||||
return chat_list
|
||||
|
||||
|
||||
async def get_folder_unread_counts(user_id: str, db: AsyncSession | None = None) -> dict[str, int]:
|
||||
user_folders = await Folders.get_folders_by_user_id(user_id, db=db)
|
||||
parent_by_id = {folder.id: folder.parent_id for folder in user_folders}
|
||||
unread_counts = dict.fromkeys(parent_by_id.keys(), 0)
|
||||
direct_unread_counts = await Chats.count_unread_by_folder_ids(user_id, list(parent_by_id.keys()), db=db)
|
||||
|
||||
for unread_folder_id, unread_count in direct_unread_counts.items():
|
||||
current_id = unread_folder_id
|
||||
seen = set()
|
||||
while current_id and current_id not in seen:
|
||||
seen.add(current_id)
|
||||
if current_id in unread_counts:
|
||||
unread_counts[current_id] += unread_count
|
||||
current_id = parent_by_id.get(current_id)
|
||||
|
||||
return unread_counts
|
||||
|
||||
|
||||
class ChatConfigForm(BaseModel):
|
||||
ENABLE_CONTEXT_COMPACTION: bool
|
||||
CONTEXT_COMPACTION_TOKEN_THRESHOLD: int
|
||||
@@ -203,6 +221,8 @@ async def get_session_user_chat_list(
|
||||
page: int | None = None,
|
||||
include_pinned: bool | None = False,
|
||||
include_folders: bool | None = False,
|
||||
sort_by: str = 'updated_at',
|
||||
sort_dir: str = 'desc',
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
try:
|
||||
@@ -214,6 +234,8 @@ async def get_session_user_chat_list(
|
||||
user.id,
|
||||
include_folders=include_folders,
|
||||
include_pinned=include_pinned,
|
||||
sort_by=sort_by,
|
||||
sort_dir=sort_dir,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
db=db,
|
||||
@@ -223,6 +245,8 @@ async def get_session_user_chat_list(
|
||||
user.id,
|
||||
include_folders=include_folders,
|
||||
include_pinned=include_pinned,
|
||||
sort_by=sort_by,
|
||||
sort_dir=sort_dir,
|
||||
db=db,
|
||||
)
|
||||
return await add_active_state_to_chat_list(request, chats)
|
||||
@@ -869,6 +893,8 @@ async def get_chat_list_by_folder_id(
|
||||
request: Request,
|
||||
folder_id: str,
|
||||
page: int | None = 1,
|
||||
sort_by: str = 'unread_updated_at',
|
||||
sort_dir: str = 'desc',
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
@@ -876,7 +902,15 @@ async def get_chat_list_by_folder_id(
|
||||
limit = 10
|
||||
skip = (page - 1) * limit
|
||||
|
||||
chats = await Chats.get_chats_by_folder_id_and_user_id(folder_id, user.id, skip=skip, limit=limit, db=db)
|
||||
chats = await Chats.get_chats_by_folder_id_and_user_id(
|
||||
folder_id,
|
||||
user.id,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
sort_by=sort_by,
|
||||
sort_dir=sort_dir,
|
||||
db=db,
|
||||
)
|
||||
return await add_active_state_to_chat_list(request, chats)
|
||||
|
||||
except Exception as e:
|
||||
@@ -2033,6 +2067,28 @@ class ChatFolderIdForm(BaseModel):
|
||||
folder_id: str | None = None
|
||||
|
||||
|
||||
@router.post('/{id}/unread')
|
||||
async def mark_chat_unread_by_id(
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
chat = await Chats.mark_chat_unread_by_id(id, user.id, db=db)
|
||||
if not chat:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
folder_id = await Chats.get_chat_folder_id(id, user.id, db=db)
|
||||
return {
|
||||
'chat_id': id,
|
||||
'last_read_at': chat.last_read_at,
|
||||
'folder_id': folder_id,
|
||||
'folder_unread_counts': await get_folder_unread_counts(user.id, db=db),
|
||||
}
|
||||
|
||||
|
||||
@router.post('/{id}/folder', response_model=ChatResponse | None)
|
||||
async def update_chat_folder_id_by_id(
|
||||
request: Request,
|
||||
|
||||
@@ -45,6 +45,24 @@ router = APIRouter()
|
||||
from open_webui.utils.access_control.folders import has_folder_access as _has_folder_access
|
||||
|
||||
|
||||
async def get_folder_unread_counts(user_id: str, db: AsyncSession | None = None) -> dict[str, int]:
|
||||
folders = await Folders.get_folders_by_user_id(user_id, db=db)
|
||||
parent_by_id = {folder.id: folder.parent_id for folder in folders}
|
||||
unread_counts = dict.fromkeys(parent_by_id.keys(), 0)
|
||||
direct_unread_counts = await Chats.count_unread_by_folder_ids(user_id, list(parent_by_id.keys()), db=db)
|
||||
|
||||
for unread_folder_id, unread_count in direct_unread_counts.items():
|
||||
current_id = unread_folder_id
|
||||
seen = set()
|
||||
while current_id and current_id not in seen:
|
||||
seen.add(current_id)
|
||||
if current_id in unread_counts:
|
||||
unread_counts[current_id] += unread_count
|
||||
current_id = parent_by_id.get(current_id)
|
||||
|
||||
return unread_counts
|
||||
|
||||
|
||||
async def check_folders_permission(request: Request, user, db=None):
|
||||
"""Verify the folders feature is enabled and the user has permission."""
|
||||
config = await Config.get_many('folders.enable', 'user.permissions')
|
||||
@@ -97,19 +115,7 @@ async def get_folders(
|
||||
|
||||
folder_list.append(folder)
|
||||
|
||||
direct_unread_counts = await Chats.count_unread_by_folder_ids(
|
||||
user.id, [folder.id for folder in folder_list], db=db
|
||||
)
|
||||
parent_by_id = {folder.id: folder.parent_id for folder in folder_list}
|
||||
unread_counts = dict.fromkeys(parent_by_id.keys(), 0)
|
||||
for unread_folder_id, unread_count in direct_unread_counts.items():
|
||||
current_id = unread_folder_id
|
||||
seen = set()
|
||||
while current_id and current_id not in seen:
|
||||
seen.add(current_id)
|
||||
if current_id in unread_counts:
|
||||
unread_counts[current_id] += unread_count
|
||||
current_id = parent_by_id.get(current_id)
|
||||
unread_counts = await get_folder_unread_counts(user.id, db=db)
|
||||
|
||||
return [
|
||||
FolderNameIdResponse(**folder.model_dump(), unread_count=unread_counts.get(folder.id, 0))
|
||||
@@ -504,7 +510,7 @@ async def get_shared_folder_chats(
|
||||
request: Request,
|
||||
id: str,
|
||||
page: int | None = Query(None, ge=1),
|
||||
sort_by: str = Query('updated_at'),
|
||||
sort_by: str = Query('unread_updated_at'),
|
||||
sort_dir: str = Query('desc'),
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
@@ -537,6 +543,7 @@ async def get_shared_folder_chats(
|
||||
limit=limit if page is not None else 60,
|
||||
sort_by=sort_by,
|
||||
sort_dir=sort_dir,
|
||||
unread_for_user_id=user.id,
|
||||
db=db,
|
||||
)
|
||||
total = await Chats.count_all_chats_by_folder_id(id, db=db) if page is not None else len(chats)
|
||||
@@ -564,6 +571,44 @@ async def get_shared_folder_chats(
|
||||
return response
|
||||
|
||||
|
||||
@router.post('/{id}/read')
|
||||
async def mark_folder_chats_read_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
await check_folders_permission(request, user, db=db)
|
||||
folder = await Folders.get_folder_by_id(id, db=db)
|
||||
if not folder:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
is_owner = user.id == folder.user_id
|
||||
is_admin = user.role == 'admin'
|
||||
if not (is_owner or is_admin or await _has_folder_access(user.id, folder, 'read', db)):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
folder_ids = (
|
||||
await Folders.get_folder_ids_by_id_and_user_id_in_subtree(id, folder.user_id, db=db)
|
||||
if is_owner or is_admin
|
||||
else [id]
|
||||
)
|
||||
updated_count = await Chats.mark_chats_read_by_folder_ids(user.id, folder_ids, db=db)
|
||||
|
||||
return {
|
||||
'folder_id': id,
|
||||
'folder_ids': folder_ids,
|
||||
'updated_count': updated_count,
|
||||
'folder_unread_counts': await get_folder_unread_counts(user.id, db=db),
|
||||
}
|
||||
|
||||
|
||||
############################
|
||||
# Delete Folder By Id
|
||||
############################
|
||||
|
||||
@@ -545,20 +545,24 @@ async def chat_events(sid, data):
|
||||
event_type = event_data.get('type')
|
||||
|
||||
if event_type == 'last_read_at':
|
||||
last_read_at = await Chats.update_chat_last_read_at_by_id(data['chat_id'], user['id'])
|
||||
if not last_read_at:
|
||||
read_update = await Chats.update_chat_last_read_at_by_id(data['chat_id'], user['id'])
|
||||
if not read_update:
|
||||
return
|
||||
last_read_at, was_unread = read_update
|
||||
response_data = {
|
||||
'chat_id': data['chat_id'],
|
||||
'last_read_at': last_read_at,
|
||||
}
|
||||
if was_unread:
|
||||
response_data['folder_unread_counts'] = await get_folder_unread_counts(user['id'])
|
||||
|
||||
await sio.emit(
|
||||
'events',
|
||||
{
|
||||
'chat_id': data['chat_id'],
|
||||
'data': {
|
||||
'type': 'chat:list',
|
||||
'data': {
|
||||
'chat_id': data['chat_id'],
|
||||
'last_read_at': last_read_at,
|
||||
'folder_unread_counts': await get_folder_unread_counts(user['id']),
|
||||
},
|
||||
'data': response_data,
|
||||
},
|
||||
},
|
||||
room=f'user:{user["id"]}',
|
||||
|
||||
@@ -167,7 +167,8 @@ async def publish_chat_finished_event(
|
||||
)
|
||||
event_emitter = await get_event_emitter(metadata, update_db=False)
|
||||
if event_emitter:
|
||||
await event_emitter({'type': 'chat:list', 'data': {'chat_id': chat_id}})
|
||||
folder_id = metadata.get('folder_id') or await Chats.get_chat_folder_id(chat_id, metadata.get('user_id'))
|
||||
await event_emitter({'type': 'chat:list', 'data': {'chat_id': chat_id, 'folder_id': folder_id}})
|
||||
|
||||
|
||||
# We believe in one maker of all models, seen and unseen,
|
||||
|
||||
Reference in New Issue
Block a user