refac: shared chat
This commit is contained in:
@@ -17,13 +17,14 @@ from open_webui.models.chats import (
|
||||
ChatResponse,
|
||||
Chats,
|
||||
ChatTitleIdResponse,
|
||||
SharedChatResponse,
|
||||
ChatStatsExport,
|
||||
AggregateChatStats,
|
||||
ChatBody,
|
||||
ChatHistoryStats,
|
||||
MessageStats,
|
||||
)
|
||||
from open_webui.models.shared_chats import SharedChats, SharedChatResponse
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.tags import TagModel, Tags
|
||||
from open_webui.models.folders import Folders
|
||||
from open_webui.internal.db import get_async_session
|
||||
@@ -35,7 +36,7 @@ from pydantic import BaseModel
|
||||
|
||||
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.access_control import has_permission, filter_allowed_access_grants
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -807,7 +808,7 @@ async def get_shared_session_user_chat_list(
|
||||
if direction:
|
||||
filter['direction'] = direction
|
||||
|
||||
return await Chats.get_shared_chat_list_by_user_id(
|
||||
return await SharedChats.get_by_user_id(
|
||||
user.id,
|
||||
filter=filter,
|
||||
skip=skip,
|
||||
@@ -828,17 +829,32 @@ async def get_shared_chat_by_id(
|
||||
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 = await Chats.get_chat_by_share_id(share_id, db=db)
|
||||
elif user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS:
|
||||
if user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS:
|
||||
chat = await Chats.get_chat_by_id(share_id, db=db)
|
||||
|
||||
if chat:
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
else:
|
||||
chat = await Chats.get_chat_by_share_id(share_id, db=db)
|
||||
|
||||
if not chat:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
# Look up the original chat_id to check access grants
|
||||
shared = await SharedChats.get_by_id(share_id, db=db)
|
||||
if shared:
|
||||
has_grant = await AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type='shared_chat',
|
||||
resource_id=shared.chat_id,
|
||||
permission='read',
|
||||
db=db,
|
||||
)
|
||||
if not has_grant:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
|
||||
############################
|
||||
# GetChatsByTags
|
||||
@@ -878,11 +894,25 @@ async def get_user_chat_list_by_tag_name(
|
||||
async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
|
||||
if not chat:
|
||||
# Check if user has access via access grants (shared_chat grants)
|
||||
if user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS:
|
||||
chat = await Chats.get_chat_by_id(id, db=db)
|
||||
else:
|
||||
has_grant = await AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type='shared_chat',
|
||||
resource_id=id,
|
||||
permission='read',
|
||||
db=db,
|
||||
)
|
||||
if has_grant:
|
||||
chat = await Chats.get_chat_by_id(id, db=db)
|
||||
|
||||
if chat:
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
else:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
|
||||
############################
|
||||
@@ -1158,39 +1188,58 @@ async def clone_shared_chat_by_id(
|
||||
else:
|
||||
chat = await Chats.get_chat_by_share_id(id, db=db)
|
||||
|
||||
if chat:
|
||||
updated_chat = {
|
||||
**chat.chat,
|
||||
'originalChatId': chat.id,
|
||||
'branchPointMessageId': chat.chat['history']['currentId'],
|
||||
'title': f'Clone of {chat.title}',
|
||||
}
|
||||
|
||||
chats = await Chats.import_chats(
|
||||
user.id,
|
||||
[
|
||||
ChatImportForm(
|
||||
**{
|
||||
'chat': updated_chat,
|
||||
'meta': chat.meta,
|
||||
'pinned': chat.pinned,
|
||||
'folder_id': chat.folder_id,
|
||||
}
|
||||
)
|
||||
],
|
||||
db=db,
|
||||
if not chat:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if chats:
|
||||
chat = chats[0]
|
||||
return ChatResponse(**chat.model_dump())
|
||||
else:
|
||||
# Enforce access grants
|
||||
shared = await SharedChats.get_by_id(id, db=db)
|
||||
if shared and user.role != 'admin':
|
||||
has_grant = await AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type='shared_chat',
|
||||
resource_id=shared.chat_id,
|
||||
permission='read',
|
||||
db=db,
|
||||
)
|
||||
if not has_grant:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=ERROR_MESSAGES.DEFAULT(),
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
updated_chat = {
|
||||
**chat.chat,
|
||||
'originalChatId': chat.id,
|
||||
'branchPointMessageId': chat.chat['history']['currentId'],
|
||||
'title': f'Clone of {chat.title}',
|
||||
}
|
||||
|
||||
chats = await Chats.import_chats(
|
||||
user.id,
|
||||
[
|
||||
ChatImportForm(
|
||||
**{
|
||||
'chat': updated_chat,
|
||||
'meta': chat.meta,
|
||||
'pinned': chat.pinned,
|
||||
'folder_id': chat.folder_id,
|
||||
}
|
||||
)
|
||||
],
|
||||
db=db,
|
||||
)
|
||||
|
||||
if chats:
|
||||
chat = chats[0]
|
||||
return ChatResponse(**chat.model_dump())
|
||||
else:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT())
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=ERROR_MESSAGES.DEFAULT(),
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
@@ -1241,16 +1290,28 @@ async def share_chat_by_id(
|
||||
|
||||
if chat:
|
||||
if chat.share_id:
|
||||
shared_chat = await Chats.update_shared_chat_by_chat_id(chat.id, db=db)
|
||||
return ChatResponse(**shared_chat.model_dump())
|
||||
# Re-snapshot existing share
|
||||
shared = await SharedChats.update(chat.share_id, db=db)
|
||||
if shared:
|
||||
# Re-fetch the original chat to return
|
||||
chat = await Chats.get_chat_by_id(id, db=db)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
shared_chat = await Chats.insert_shared_chat_by_chat_id(chat.id, db=db)
|
||||
if not shared_chat:
|
||||
# Create new share
|
||||
shared = await SharedChats.create(id, user.id, db=db)
|
||||
if not shared:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=ERROR_MESSAGES.DEFAULT(),
|
||||
)
|
||||
return ChatResponse(**shared_chat.model_dump())
|
||||
# Set share_id on the original chat
|
||||
chat = await Chats.update_chat_share_id_by_id(id, shared.id, db=db)
|
||||
if not chat:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=ERROR_MESSAGES.DEFAULT(),
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -1260,7 +1321,7 @@ async def share_chat_by_id(
|
||||
|
||||
|
||||
############################
|
||||
# DeletedSharedChatById
|
||||
# DeleteSharedChatById
|
||||
############################
|
||||
|
||||
|
||||
@@ -1273,10 +1334,13 @@ async def delete_shared_chat_by_id(
|
||||
if not chat.share_id:
|
||||
return False
|
||||
|
||||
result = await Chats.delete_shared_chat_by_chat_id(id, db=db)
|
||||
update_result = await Chats.update_chat_share_id_by_id(id, None, db=db)
|
||||
await SharedChats.delete_by_chat_id(id, db=db)
|
||||
await Chats.update_chat_share_id_by_id(id, None, db=db)
|
||||
|
||||
return result and update_result != None
|
||||
# Revoke all access grants for this shared chat
|
||||
await AccessGrants.set_access_grants('shared_chat', id, [], db=db)
|
||||
|
||||
return True
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -1284,6 +1348,85 @@ async def delete_shared_chat_by_id(
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
# UpdateSharedChatAccessById
|
||||
############################
|
||||
|
||||
|
||||
class ChatAccessGrantsForm(BaseModel):
|
||||
access_grants: list[dict]
|
||||
|
||||
|
||||
@router.post('/shared/{id}/access/update', response_model=Optional[ChatResponse])
|
||||
async def update_shared_chat_access_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: ChatAccessGrantsForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if not chat:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if chat.user_id != user.id and user.role != 'admin':
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
'sharing.public_chats',
|
||||
)
|
||||
|
||||
await AccessGrants.set_access_grants('shared_chat', id, form_data.access_grants, db=db)
|
||||
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
|
||||
############################
|
||||
# GetSharedChatAccessById
|
||||
############################
|
||||
|
||||
|
||||
@router.get('/shared/{id}/access', response_model=list)
|
||||
async def get_shared_chat_access_by_id(
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if not chat:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if chat.user_id != user.id and user.role != 'admin':
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
grants = await AccessGrants.get_grants_by_resource('shared_chat', id, db=db)
|
||||
return [
|
||||
{
|
||||
'id': g.id,
|
||||
'principal_type': g.principal_type,
|
||||
'principal_id': g.principal_id,
|
||||
'permission': g.permission,
|
||||
}
|
||||
for g in grants
|
||||
]
|
||||
|
||||
|
||||
############################
|
||||
# UpdateChatFolderIdById
|
||||
############################
|
||||
|
||||
Reference in New Issue
Block a user