feat: ADRF async views (phase1) + native async serializers (phase2) + async cache infra
This commit is contained in:
+48
-30
@@ -1,17 +1,33 @@
|
||||
from rest_framework import serializers
|
||||
from adrf.serializers import (
|
||||
ModelSerializer, CharField, SerializerMethodField,
|
||||
)
|
||||
from django.contrib.auth import get_user_model
|
||||
from django.core.exceptions import SynchronousOnlyOperation
|
||||
from .models import Message, SystemMessage, SystemMessageRead
|
||||
|
||||
User = get_user_model()
|
||||
|
||||
class MessageListSerializer(serializers.ModelSerializer):
|
||||
sender_name = serializers.SerializerMethodField()
|
||||
sender_avatar = serializers.SerializerMethodField()
|
||||
time = serializers.SerializerMethodField()
|
||||
is_comment = serializers.SerializerMethodField()
|
||||
comment_id = serializers.SerializerMethodField()
|
||||
article_id = serializers.SerializerMethodField()
|
||||
type = serializers.CharField(source='msg_type', read_only=True)
|
||||
extra = serializers.CharField(source='extra_info', read_only=True)
|
||||
link = serializers.CharField(source='target_link', read_only=True)
|
||||
|
||||
async def _load_sender(obj):
|
||||
"""优先命中 select_related 预加载缓存;未预加载时异步回源,避免外键懒加载"""
|
||||
try:
|
||||
return obj.sender
|
||||
except SynchronousOnlyOperation:
|
||||
if not obj.sender_id:
|
||||
return None
|
||||
return await User.objects.aget(pk=obj.sender_id)
|
||||
|
||||
|
||||
class MessageListSerializer(ModelSerializer):
|
||||
sender_name = SerializerMethodField()
|
||||
sender_avatar = SerializerMethodField()
|
||||
time = SerializerMethodField()
|
||||
is_comment = SerializerMethodField()
|
||||
comment_id = SerializerMethodField()
|
||||
article_id = SerializerMethodField()
|
||||
type = CharField(source='msg_type', read_only=True)
|
||||
extra = CharField(source='extra_info', read_only=True)
|
||||
link = CharField(source='target_link', read_only=True)
|
||||
|
||||
class Meta:
|
||||
model = Message
|
||||
@@ -21,31 +37,33 @@ class MessageListSerializer(serializers.ModelSerializer):
|
||||
'is_comment', 'comment_id', 'article_id',
|
||||
]
|
||||
|
||||
def get_sender_name(self, obj):
|
||||
if obj.sender:
|
||||
return obj.sender.nickname or obj.sender.username
|
||||
async def get_sender_name(self, obj):
|
||||
sender = await _load_sender(obj)
|
||||
if sender:
|
||||
return sender.nickname or sender.username
|
||||
return '系统'
|
||||
|
||||
def get_sender_avatar(self, obj):
|
||||
if obj.sender and obj.sender.avatar:
|
||||
async def get_sender_avatar(self, obj):
|
||||
sender = await _load_sender(obj)
|
||||
if sender and sender.avatar:
|
||||
request = self.context.get('request')
|
||||
if request:
|
||||
return request.build_absolute_uri(obj.sender.avatar.url)
|
||||
return obj.sender.avatar.url
|
||||
return request.build_absolute_uri(sender.avatar.url)
|
||||
return sender.avatar.url
|
||||
return ''
|
||||
|
||||
def get_time(self, obj):
|
||||
async def get_time(self, obj):
|
||||
return obj.created_at.strftime('%Y-%m-%d %H:%M')
|
||||
|
||||
def get_is_comment(self, obj):
|
||||
async def get_is_comment(self, obj):
|
||||
return obj.target_type == 'comment'
|
||||
|
||||
def get_comment_id(self, obj):
|
||||
async def get_comment_id(self, obj):
|
||||
if obj.target_type == 'comment' and obj.target_id:
|
||||
return str(obj.target_id)
|
||||
return ''
|
||||
|
||||
def get_article_id(self, obj):
|
||||
async def get_article_id(self, obj):
|
||||
if obj.target_link:
|
||||
import re
|
||||
match = re.search(r'/article/(\d+)', obj.target_link)
|
||||
@@ -59,30 +77,30 @@ class MessageListSerializer(serializers.ModelSerializer):
|
||||
return None
|
||||
|
||||
|
||||
class SystemMessageSerializer(serializers.ModelSerializer):
|
||||
is_read = serializers.SerializerMethodField()
|
||||
class SystemMessageSerializer(ModelSerializer):
|
||||
is_read = SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = SystemMessage
|
||||
fields = ['id', 'title', 'content', 'category', 'icon_name', 'icon_color', 'created_at', 'is_read']
|
||||
|
||||
def get_is_read(self, obj):
|
||||
async def get_is_read(self, obj):
|
||||
read_ids = self.context.get('read_ids', set())
|
||||
return obj.id in read_ids
|
||||
|
||||
|
||||
class SystemMessageDetailSerializer(serializers.ModelSerializer):
|
||||
is_read = serializers.SerializerMethodField()
|
||||
class SystemMessageDetailSerializer(ModelSerializer):
|
||||
is_read = SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = SystemMessage
|
||||
fields = ['id', 'title', 'content', 'detail', 'category', 'icon_name', 'icon_color', 'created_at', 'is_read']
|
||||
|
||||
def get_is_read(self, obj):
|
||||
async def get_is_read(self, obj):
|
||||
request = self.context.get('request')
|
||||
if request and request.user.is_authenticated:
|
||||
return SystemMessageRead.objects.filter(
|
||||
return await SystemMessageRead.objects.filter(
|
||||
user=request.user,
|
||||
system_message=obj
|
||||
).exists()
|
||||
).aexists()
|
||||
return False
|
||||
|
||||
+62
-48
@@ -1,10 +1,11 @@
|
||||
from rest_framework import status
|
||||
from rest_framework.views import APIView
|
||||
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
|
||||
@@ -35,21 +36,27 @@ class MessageListView(APIView):
|
||||
],
|
||||
responses={200: success_response, 401: unauthorized_response}
|
||||
)
|
||||
def get(self, request):
|
||||
async def get(self, request):
|
||||
msg_type = request.query_params.get('type', '')
|
||||
|
||||
if msg_type == 'system':
|
||||
queryset = SystemMessage.objects.all()
|
||||
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 create_standardized_response(data=paginator.get_paginated_response(serializer.data).data, code=ResponseCode.SUCCESS)
|
||||
|
||||
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:
|
||||
@@ -57,10 +64,15 @@ class MessageListView(APIView):
|
||||
else:
|
||||
queryset = Message.objects.filter(recipient=request.user)
|
||||
|
||||
paginator = MessagePagination()
|
||||
page = paginator.paginate_queryset(queryset, request)
|
||||
serializer = MessageListSerializer(page, many=True, context={'request': request})
|
||||
return create_standardized_response(data=paginator.get_paginated_response(serializer.data).data, code=ResponseCode.SUCCESS)
|
||||
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):
|
||||
@@ -73,19 +85,19 @@ class UnreadCountView(APIView):
|
||||
operation_description='获取当前用户各类型未读消息的数量统计',
|
||||
responses={200: success_response, 401: unauthorized_response}
|
||||
)
|
||||
def get(self, request):
|
||||
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}
|
||||
for item in message_counts:
|
||||
async for item in message_counts:
|
||||
if item['msg_type'] in counts:
|
||||
counts[item['msg_type']] = item['count']
|
||||
|
||||
read_system_ids = SystemMessageRead.objects.filter(user=user).values_list('system_message_id', flat=True)
|
||||
system_unread = SystemMessage.objects.filter(is_global=True).exclude(id__in=read_system_ids).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']
|
||||
@@ -106,18 +118,18 @@ class MessageReadView(APIView):
|
||||
],
|
||||
responses={200: success_response, 401: unauthorized_response, 404: not_found_response}
|
||||
)
|
||||
def post(self, request, pk):
|
||||
async def post(self, request, pk):
|
||||
try:
|
||||
message = Message.objects.get(pk=pk, recipient=request.user)
|
||||
message = await Message.objects.aget(pk=pk, recipient=request.user)
|
||||
message.is_read = True
|
||||
message.save(update_fields=['is_read'])
|
||||
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 = SystemMessage.objects.get(pk=pk)
|
||||
SystemMessageRead.objects.get_or_create(
|
||||
system_message = await SystemMessage.objects.aget(pk=pk)
|
||||
await SystemMessageRead.objects.aget_or_create(
|
||||
user=request.user,
|
||||
system_message=system_message
|
||||
)
|
||||
@@ -146,32 +158,32 @@ class MessageReadAllView(APIView):
|
||||
),
|
||||
responses={200: success_response, 401: unauthorized_response}
|
||||
)
|
||||
def post(self, request):
|
||||
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)
|
||||
for sm in system_messages:
|
||||
_, created = SystemMessageRead.objects.get_or_create(
|
||||
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 = Message.objects.filter(
|
||||
affected = await Message.objects.filter(
|
||||
recipient=user, msg_type=msg_type, is_read=False
|
||||
).update(is_read=True)
|
||||
).aupdate(is_read=True)
|
||||
else:
|
||||
affected = Message.objects.filter(
|
||||
affected = await Message.objects.filter(
|
||||
recipient=user, is_read=False
|
||||
).update(is_read=True)
|
||||
).aupdate(is_read=True)
|
||||
|
||||
system_messages = SystemMessage.objects.filter(is_global=True)
|
||||
for sm in system_messages:
|
||||
_, created = SystemMessageRead.objects.get_or_create(
|
||||
async for sm in system_messages:
|
||||
_, created = await SystemMessageRead.objects.aget_or_create(
|
||||
user=user,
|
||||
system_message=sm
|
||||
)
|
||||
@@ -194,17 +206,17 @@ class MessageDeleteView(APIView):
|
||||
],
|
||||
responses={204: '删除成功', 401: unauthorized_response, 404: not_found_response}
|
||||
)
|
||||
def delete(self, request, pk):
|
||||
async def delete(self, request, pk):
|
||||
try:
|
||||
message = Message.objects.get(pk=pk, recipient=request.user)
|
||||
message.delete()
|
||||
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 = SystemMessage.objects.get(pk=pk)
|
||||
SystemMessageRead.objects.get_or_create(
|
||||
system_message = await SystemMessage.objects.aget(pk=pk)
|
||||
await SystemMessageRead.objects.aget_or_create(
|
||||
user=request.user,
|
||||
system_message=system_message
|
||||
)
|
||||
@@ -233,21 +245,21 @@ class MessageClearAllView(APIView):
|
||||
),
|
||||
responses={200: success_response, 401: unauthorized_response}
|
||||
)
|
||||
def delete(self, request):
|
||||
async def delete(self, request):
|
||||
msg_type = request.data.get('type', '')
|
||||
user = request.user
|
||||
deleted_count = 0
|
||||
|
||||
if msg_type == 'system':
|
||||
# 删除系统消息的已读记录(相当于清空系统消息)
|
||||
deleted_count = SystemMessageRead.objects.filter(user=user).delete()[0]
|
||||
deleted_count = (await SystemMessageRead.objects.filter(user=user).adelete())[0]
|
||||
elif msg_type and msg_type in ['reply', 'at_me', 'like']:
|
||||
deleted_count = Message.objects.filter(recipient=user, msg_type=msg_type).delete()[0]
|
||||
deleted_count = (await Message.objects.filter(recipient=user, msg_type=msg_type).adelete())[0]
|
||||
else:
|
||||
# 删除所有普通消息
|
||||
deleted_count = Message.objects.filter(recipient=user).delete()[0]
|
||||
deleted_count = (await Message.objects.filter(recipient=user).adelete())[0]
|
||||
# 同时删除所有系统消息已读记录
|
||||
deleted_count += SystemMessageRead.objects.filter(user=user).delete()[0]
|
||||
deleted_count += (await SystemMessageRead.objects.filter(user=user).adelete())[0]
|
||||
|
||||
return create_standardized_response(data={'deleted': deleted_count}, code=ResponseCode.SUCCESS)
|
||||
|
||||
@@ -265,9 +277,9 @@ class SystemMessageDetailView(APIView):
|
||||
],
|
||||
responses={200: success_response, 401: unauthorized_response, 404: not_found_response}
|
||||
)
|
||||
def get(self, request, pk):
|
||||
async def get(self, request, pk):
|
||||
try:
|
||||
system_message = SystemMessage.objects.get(pk=pk)
|
||||
system_message = await SystemMessage.objects.aget(pk=pk)
|
||||
except SystemMessage.DoesNotExist:
|
||||
return create_standardized_error_response(
|
||||
code=ResponseCode.NOT_FOUND,
|
||||
@@ -275,10 +287,12 @@ class SystemMessageDetailView(APIView):
|
||||
status_code=status.HTTP_404_NOT_FOUND
|
||||
)
|
||||
|
||||
SystemMessageRead.objects.get_or_create(
|
||||
await SystemMessageRead.objects.aget_or_create(
|
||||
user=request.user,
|
||||
system_message=system_message
|
||||
)
|
||||
|
||||
serializer = SystemMessageDetailSerializer(system_message, context={'request': request})
|
||||
return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS)
|
||||
# 兜底:SystemMessageDetailSerializer.get_is_read 内部有同步 ORM exists() 查询
|
||||
data = await sync_to_async(lambda: serializer.data)()
|
||||
return create_standardized_response(data=data, code=ResponseCode.SUCCESS)
|
||||
|
||||
Reference in New Issue
Block a user