299 lines
12 KiB
Python
299 lines
12 KiB
Python
from rest_framework import status
|
||
from adrf.views import APIView
|
||
from rest_framework.permissions import IsAuthenticated
|
||
from rest_framework.parsers import JSONParser
|
||
from rest_framework.pagination import PageNumberPagination
|
||
from rest_framework.response import Response
|
||
from django.db.models import Count, Q
|
||
from asgiref.sync import sync_to_async
|
||
from drf_yasg.utils import swagger_auto_schema
|
||
from drf_yasg import openapi
|
||
from chunyu_project.common_schemas import success_response, error_response, unauthorized_response, not_found_response
|
||
|
||
from utils.response_codes import ResponseCode, create_standardized_response, create_standardized_error_response
|
||
from .models import Message, SystemMessage, SystemMessageRead
|
||
from .serializers import MessageListSerializer, SystemMessageSerializer, SystemMessageDetailSerializer
|
||
|
||
|
||
class MessagePagination(PageNumberPagination):
|
||
page_size = 20
|
||
page_size_query_param = 'page_size'
|
||
max_page_size = 100
|
||
|
||
|
||
class MessageListView(APIView):
|
||
permission_classes = [IsAuthenticated]
|
||
parser_classes = [JSONParser]
|
||
|
||
@swagger_auto_schema(
|
||
tags=['消息'],
|
||
operation_summary='获取消息列表',
|
||
operation_description='获取当前用户的消息列表,支持按类型筛选: reply/at_me/like/system',
|
||
manual_parameters=[
|
||
openapi.Parameter('type', openapi.IN_QUERY, description='消息类型: reply/at_me/like/system', type=openapi.TYPE_STRING),
|
||
openapi.Parameter('page', openapi.IN_QUERY, description='页码', type=openapi.TYPE_INTEGER),
|
||
openapi.Parameter('page_size', openapi.IN_QUERY, description='每页数量', type=openapi.TYPE_INTEGER),
|
||
],
|
||
responses={200: success_response, 401: unauthorized_response}
|
||
)
|
||
async def get(self, request):
|
||
msg_type = request.query_params.get('type', '')
|
||
|
||
if msg_type == 'system':
|
||
queryset = SystemMessage.objects.all()
|
||
|
||
def _paginate_system():
|
||
paginator = MessagePagination()
|
||
page = paginator.paginate_queryset(queryset, request)
|
||
read_ids = set(
|
||
SystemMessageRead.objects.filter(
|
||
user=request.user,
|
||
system_message__in=page
|
||
).values_list('system_message_id', flat=True)
|
||
)
|
||
serializer = SystemMessageSerializer(page, many=True, context={'request': request, 'read_ids': read_ids})
|
||
return paginator.get_paginated_response(serializer.data).data
|
||
|
||
# 兜底:DRF 分页器与 SystemMessageSerializer 内部为同步调用
|
||
data = await sync_to_async(_paginate_system)()
|
||
return create_standardized_response(data=data, code=ResponseCode.SUCCESS)
|
||
|
||
valid_types = ['reply', 'at_me', 'like']
|
||
if msg_type and msg_type in valid_types:
|
||
queryset = Message.objects.filter(recipient=request.user, msg_type=msg_type)
|
||
else:
|
||
queryset = Message.objects.filter(recipient=request.user)
|
||
|
||
def _paginate_messages():
|
||
paginator = MessagePagination()
|
||
page = paginator.paginate_queryset(queryset, request)
|
||
serializer = MessageListSerializer(page, many=True, context={'request': request})
|
||
return paginator.get_paginated_response(serializer.data).data
|
||
|
||
# 兜底:DRF 分页器与 MessageListSerializer(obj.sender 外键访问)内部为同步调用
|
||
data = await sync_to_async(_paginate_messages)()
|
||
return create_standardized_response(data=data, code=ResponseCode.SUCCESS)
|
||
|
||
|
||
class UnreadCountView(APIView):
|
||
permission_classes = [IsAuthenticated]
|
||
parser_classes = [JSONParser]
|
||
|
||
@swagger_auto_schema(
|
||
tags=['消息'],
|
||
operation_summary='获取未读消息数量',
|
||
operation_description='获取当前用户各类型未读消息的数量统计',
|
||
responses={200: success_response, 401: unauthorized_response}
|
||
)
|
||
async def get(self, request):
|
||
user = request.user
|
||
message_counts = Message.objects.filter(
|
||
recipient=user, is_read=False
|
||
).values('msg_type').annotate(count=Count('id'))
|
||
|
||
counts = {'reply': 0, 'at_me': 0, 'like': 0}
|
||
async for item in message_counts:
|
||
if item['msg_type'] in counts:
|
||
counts[item['msg_type']] = item['count']
|
||
|
||
read_system_ids = [rid async for rid in SystemMessageRead.objects.filter(user=user).values_list('system_message_id', flat=True)]
|
||
system_unread = await SystemMessage.objects.filter(is_global=True).exclude(id__in=read_system_ids).acount()
|
||
|
||
counts['system'] = system_unread
|
||
counts['total'] = counts['reply'] + counts['at_me'] + counts['like'] + counts['system']
|
||
|
||
return create_standardized_response(data=counts, code=ResponseCode.SUCCESS)
|
||
|
||
|
||
class MessageReadView(APIView):
|
||
permission_classes = [IsAuthenticated]
|
||
parser_classes = [JSONParser]
|
||
|
||
@swagger_auto_schema(
|
||
tags=['消息'],
|
||
operation_summary='标记消息为已读',
|
||
operation_description='将指定消息标记为已读,支持普通消息和系统消息',
|
||
manual_parameters=[
|
||
openapi.Parameter('pk', openapi.IN_PATH, description='消息ID', type=openapi.TYPE_INTEGER, required=True),
|
||
],
|
||
responses={200: success_response, 401: unauthorized_response, 404: not_found_response}
|
||
)
|
||
async def post(self, request, pk):
|
||
try:
|
||
message = await Message.objects.aget(pk=pk, recipient=request.user)
|
||
message.is_read = True
|
||
await message.asave(update_fields=['is_read'])
|
||
return create_standardized_response(data={'is_read': True}, code=ResponseCode.SUCCESS)
|
||
except Message.DoesNotExist:
|
||
pass
|
||
|
||
try:
|
||
system_message = await SystemMessage.objects.aget(pk=pk)
|
||
await SystemMessageRead.objects.aget_or_create(
|
||
user=request.user,
|
||
system_message=system_message
|
||
)
|
||
return create_standardized_response(data={'is_read': True}, code=ResponseCode.SUCCESS)
|
||
except SystemMessage.DoesNotExist:
|
||
return create_standardized_error_response(
|
||
code=ResponseCode.NOT_FOUND,
|
||
message='消息不存在',
|
||
status_code=status.HTTP_404_NOT_FOUND
|
||
)
|
||
|
||
|
||
class MessageReadAllView(APIView):
|
||
permission_classes = [IsAuthenticated]
|
||
parser_classes = [JSONParser]
|
||
|
||
@swagger_auto_schema(
|
||
tags=['消息'],
|
||
operation_summary='全部标记为已读',
|
||
operation_description='将指定类型或所有消息标记为已读',
|
||
request_body=openapi.Schema(
|
||
type=openapi.TYPE_OBJECT,
|
||
properties={
|
||
'type': openapi.Schema(type=openapi.TYPE_STRING, description='消息类型: reply/at_me/like/system,不传则标记全部'),
|
||
}
|
||
),
|
||
responses={200: success_response, 401: unauthorized_response}
|
||
)
|
||
async def post(self, request):
|
||
msg_type = request.data.get('type', '')
|
||
user = request.user
|
||
affected = 0
|
||
|
||
if msg_type == 'system':
|
||
system_messages = SystemMessage.objects.filter(is_global=True)
|
||
async for sm in system_messages:
|
||
_, created = await SystemMessageRead.objects.aget_or_create(
|
||
user=user,
|
||
system_message=sm
|
||
)
|
||
if created:
|
||
affected += 1
|
||
elif msg_type and msg_type in ['reply', 'at_me', 'like']:
|
||
affected = await Message.objects.filter(
|
||
recipient=user, msg_type=msg_type, is_read=False
|
||
).aupdate(is_read=True)
|
||
else:
|
||
affected = await Message.objects.filter(
|
||
recipient=user, is_read=False
|
||
).aupdate(is_read=True)
|
||
|
||
system_messages = SystemMessage.objects.filter(is_global=True)
|
||
async for sm in system_messages:
|
||
_, created = await SystemMessageRead.objects.aget_or_create(
|
||
user=user,
|
||
system_message=sm
|
||
)
|
||
if created:
|
||
affected += 1
|
||
|
||
return create_standardized_response(data={'affected': affected}, code=ResponseCode.SUCCESS)
|
||
|
||
|
||
class MessageDeleteView(APIView):
|
||
permission_classes = [IsAuthenticated]
|
||
parser_classes = [JSONParser]
|
||
|
||
@swagger_auto_schema(
|
||
tags=['消息'],
|
||
operation_summary='删除消息',
|
||
operation_description='删除指定消息,支持普通消息和系统消息',
|
||
manual_parameters=[
|
||
openapi.Parameter('pk', openapi.IN_PATH, description='消息ID', type=openapi.TYPE_INTEGER, required=True),
|
||
],
|
||
responses={204: '删除成功', 401: unauthorized_response, 404: not_found_response}
|
||
)
|
||
async def delete(self, request, pk):
|
||
try:
|
||
message = await Message.objects.aget(pk=pk, recipient=request.user)
|
||
await message.adelete()
|
||
return Response(status=status.HTTP_204_NO_CONTENT)
|
||
except Message.DoesNotExist:
|
||
pass
|
||
|
||
try:
|
||
system_message = await SystemMessage.objects.aget(pk=pk)
|
||
await SystemMessageRead.objects.aget_or_create(
|
||
user=request.user,
|
||
system_message=system_message
|
||
)
|
||
return Response(status=status.HTTP_204_NO_CONTENT)
|
||
except SystemMessage.DoesNotExist:
|
||
return create_standardized_error_response(
|
||
code=ResponseCode.NOT_FOUND,
|
||
message='消息不存在',
|
||
status_code=status.HTTP_404_NOT_FOUND
|
||
)
|
||
|
||
|
||
class MessageClearAllView(APIView):
|
||
permission_classes = [IsAuthenticated]
|
||
parser_classes = [JSONParser]
|
||
|
||
@swagger_auto_schema(
|
||
tags=['消息'],
|
||
operation_summary='清空所有消息',
|
||
operation_description='删除指定类型或全部消息,可选类型: reply/at_me/like/system',
|
||
request_body=openapi.Schema(
|
||
type=openapi.TYPE_OBJECT,
|
||
properties={
|
||
'type': openapi.Schema(type=openapi.TYPE_STRING, description='消息类型: reply/at_me/like/system,不传则清空全部'),
|
||
}
|
||
),
|
||
responses={200: success_response, 401: unauthorized_response}
|
||
)
|
||
async def delete(self, request):
|
||
msg_type = request.data.get('type', '')
|
||
user = request.user
|
||
deleted_count = 0
|
||
|
||
if msg_type == 'system':
|
||
# 删除系统消息的已读记录(相当于清空系统消息)
|
||
deleted_count = (await SystemMessageRead.objects.filter(user=user).adelete())[0]
|
||
elif msg_type and msg_type in ['reply', 'at_me', 'like']:
|
||
deleted_count = (await Message.objects.filter(recipient=user, msg_type=msg_type).adelete())[0]
|
||
else:
|
||
# 删除所有普通消息
|
||
deleted_count = (await Message.objects.filter(recipient=user).adelete())[0]
|
||
# 同时删除所有系统消息已读记录
|
||
deleted_count += (await SystemMessageRead.objects.filter(user=user).adelete())[0]
|
||
|
||
return create_standardized_response(data={'deleted': deleted_count}, code=ResponseCode.SUCCESS)
|
||
|
||
|
||
class SystemMessageDetailView(APIView):
|
||
permission_classes = [IsAuthenticated]
|
||
parser_classes = [JSONParser]
|
||
|
||
@swagger_auto_schema(
|
||
tags=['消息'],
|
||
operation_summary='获取系统消息详情',
|
||
operation_description='获取指定系统消息的详细内容,并自动标记为已读',
|
||
manual_parameters=[
|
||
openapi.Parameter('pk', openapi.IN_PATH, description='系统消息ID', type=openapi.TYPE_INTEGER, required=True),
|
||
],
|
||
responses={200: success_response, 401: unauthorized_response, 404: not_found_response}
|
||
)
|
||
async def get(self, request, pk):
|
||
try:
|
||
system_message = await SystemMessage.objects.aget(pk=pk)
|
||
except SystemMessage.DoesNotExist:
|
||
return create_standardized_error_response(
|
||
code=ResponseCode.NOT_FOUND,
|
||
message='系统消息不存在',
|
||
status_code=status.HTTP_404_NOT_FOUND
|
||
)
|
||
|
||
await SystemMessageRead.objects.aget_or_create(
|
||
user=request.user,
|
||
system_message=system_message
|
||
)
|
||
|
||
serializer = SystemMessageDetailSerializer(system_message, context={'request': request})
|
||
# 兜底:SystemMessageDetailSerializer.get_is_read 内部有同步 ORM exists() 查询
|
||
data = await sync_to_async(lambda: serializer.data)()
|
||
return create_standardized_response(data=data, code=ResponseCode.SUCCESS)
|