feat: ADRF async views (phase1) + native async serializers (phase2) + async cache infra

This commit is contained in:
async-upgrade
2026-09-06 14:26:17 +08:00
parent 9a6577f71e
commit 8f488fcaaa
55 changed files with 2224 additions and 1513 deletions
+48 -30
View File
@@ -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
View File
@@ -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)