feat: ADRF async views (phase1) + native async serializers (phase2) + async cache infra
This commit is contained in:
+51
-44
@@ -1,18 +1,20 @@
|
||||
from rest_framework import serializers
|
||||
from adrf.serializers import (
|
||||
ModelSerializer, Serializer, IntegerField, CharField, SerializerMethodField,
|
||||
)
|
||||
from .models import FriendRequest, Friendship, Conversation, ConversationParticipant, ChatMessage
|
||||
from django.conf import settings
|
||||
|
||||
|
||||
class UserBriefSerializer(serializers.Serializer):
|
||||
id = serializers.IntegerField()
|
||||
username = serializers.CharField()
|
||||
nickname = serializers.SerializerMethodField()
|
||||
avatar = serializers.SerializerMethodField()
|
||||
class UserBriefSerializer(Serializer):
|
||||
id = IntegerField()
|
||||
username = CharField()
|
||||
nickname = SerializerMethodField()
|
||||
avatar = SerializerMethodField()
|
||||
|
||||
def get_nickname(self, obj):
|
||||
async def get_nickname(self, obj):
|
||||
return getattr(obj, 'nickname', '') or obj.username
|
||||
|
||||
def get_avatar(self, obj):
|
||||
async def get_avatar(self, obj):
|
||||
avatar = getattr(obj, 'avatar', None)
|
||||
if avatar and hasattr(avatar, 'url'):
|
||||
request = self.context.get('request')
|
||||
@@ -22,91 +24,96 @@ class UserBriefSerializer(serializers.Serializer):
|
||||
return ''
|
||||
|
||||
|
||||
class FriendRequestSerializer(serializers.ModelSerializer):
|
||||
class FriendRequestSerializer(ModelSerializer):
|
||||
from_user = UserBriefSerializer(read_only=True)
|
||||
to_user = UserBriefSerializer(read_only=True)
|
||||
from_user_id = serializers.IntegerField(write_only=True, required=False)
|
||||
to_user_id = serializers.IntegerField(write_only=True, required=False)
|
||||
from_user_id = IntegerField(write_only=True, required=False)
|
||||
to_user_id = IntegerField(write_only=True, required=False)
|
||||
|
||||
class Meta:
|
||||
model = FriendRequest
|
||||
fields = ['id', 'from_user', 'to_user', 'from_user_id', 'to_user_id', 'status', 'message', 'created_at', 'updated_at']
|
||||
|
||||
|
||||
class FriendshipSerializer(serializers.ModelSerializer):
|
||||
friend = serializers.SerializerMethodField()
|
||||
class FriendshipSerializer(ModelSerializer):
|
||||
friend = SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = Friendship
|
||||
fields = ['id', 'friend', 'created_at']
|
||||
|
||||
def get_friend(self, obj):
|
||||
async def get_friend(self, obj):
|
||||
request = self.context.get('request')
|
||||
if not request:
|
||||
return None
|
||||
friend_user = obj.user2 if obj.user1 == request.user else obj.user1
|
||||
return UserBriefSerializer(friend_user, context=self.context).data
|
||||
friend_user = obj.user2 if obj.user1_id == request.user.id else obj.user1
|
||||
return await UserBriefSerializer(friend_user, context=self.context).adata
|
||||
|
||||
|
||||
class ConversationSerializer(serializers.ModelSerializer):
|
||||
other_user = serializers.SerializerMethodField()
|
||||
last_message = serializers.SerializerMethodField()
|
||||
unread_count = serializers.SerializerMethodField()
|
||||
class ConversationSerializer(ModelSerializer):
|
||||
other_user = SerializerMethodField()
|
||||
last_message = SerializerMethodField()
|
||||
unread_count = SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = Conversation
|
||||
fields = ['id', 'type', 'other_user', 'last_message', 'unread_count', 'created_at']
|
||||
|
||||
def get_other_user(self, obj):
|
||||
async def get_other_user(self, obj):
|
||||
request = self.context.get('request')
|
||||
if not request:
|
||||
return None
|
||||
participant = obj.participants.exclude(user=request.user).first()
|
||||
participant = await obj.participants.exclude(user=request.user).afirst()
|
||||
if participant:
|
||||
return UserBriefSerializer(participant.user, context=self.context).data
|
||||
return await UserBriefSerializer(participant.user, context=self.context).adata
|
||||
return None
|
||||
|
||||
def get_last_message(self, obj):
|
||||
last_msg = obj.messages.order_by('-created_at').first()
|
||||
async def get_last_message(self, obj):
|
||||
last_msg = await obj.messages.order_by('-created_at').select_related('sender', 'reply_to__sender').afirst()
|
||||
if last_msg:
|
||||
return ChatMessageSerializer(last_msg, context=self.context).data
|
||||
return await ChatMessageSerializer(last_msg, context=self.context).adata
|
||||
return None
|
||||
|
||||
def get_unread_count(self, obj):
|
||||
async def get_unread_count(self, obj):
|
||||
request = self.context.get('request')
|
||||
if not request:
|
||||
return 0
|
||||
try:
|
||||
participant = ConversationParticipant.objects.get(conversation=obj, user=request.user)
|
||||
participant = await ConversationParticipant.objects.aget(conversation=obj, user=request.user)
|
||||
if participant.last_read_at:
|
||||
return obj.messages.filter(created_at__gt=participant.last_read_at).exclude(sender=request.user).count()
|
||||
return obj.messages.exclude(sender=request.user).count()
|
||||
return await obj.messages.filter(created_at__gt=participant.last_read_at).exclude(sender=request.user).acount()
|
||||
return await obj.messages.exclude(sender=request.user).acount()
|
||||
except ConversationParticipant.DoesNotExist:
|
||||
return 0
|
||||
|
||||
|
||||
class ChatMessageSerializer(serializers.ModelSerializer):
|
||||
class ChatMessageSerializer(ModelSerializer):
|
||||
sender_info = UserBriefSerializer(source='sender', read_only=True)
|
||||
is_own = serializers.SerializerMethodField()
|
||||
reply_to_message = serializers.SerializerMethodField()
|
||||
is_own = SerializerMethodField()
|
||||
reply_to_message = SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = ChatMessage
|
||||
fields = ['id', 'conversation', 'sender', 'sender_info', 'content', 'msg_type', 'file_url', 'reply_to', 'reply_to_message', 'is_recalled', 'is_own', 'created_at']
|
||||
read_only_fields = ['sender', 'conversation']
|
||||
|
||||
def get_is_own(self, obj):
|
||||
async def get_is_own(self, obj):
|
||||
request = self.context.get('request')
|
||||
if request and hasattr(request, 'user'):
|
||||
return obj.sender == request.user
|
||||
# 用 id 比较,避免 sender 外键懒加载触发同步查询
|
||||
return obj.sender_id == request.user.id
|
||||
return False
|
||||
|
||||
def get_reply_to_message(self, obj):
|
||||
if obj.reply_to and not obj.reply_to.is_recalled:
|
||||
return {
|
||||
'id': obj.reply_to.id,
|
||||
'content': obj.reply_to.content[:50],
|
||||
'sender_name': getattr(obj.reply_to.sender, 'nickname', '') or obj.reply_to.sender.username,
|
||||
'msg_type': obj.reply_to.msg_type,
|
||||
}
|
||||
return None
|
||||
async def get_reply_to_message(self, obj):
|
||||
if not obj.reply_to_id:
|
||||
return None
|
||||
# 异步加载被回复消息及其发送者,避免外键懒加载触发同步查询
|
||||
reply = await ChatMessage.objects.select_related('sender').aget(pk=obj.reply_to_id)
|
||||
if reply.is_recalled:
|
||||
return None
|
||||
return {
|
||||
'id': reply.id,
|
||||
'content': reply.content[:50],
|
||||
'sender_name': getattr(reply.sender, 'nickname', '') or reply.sender.username,
|
||||
'msg_type': reply.msg_type,
|
||||
}
|
||||
|
||||
+118
-104
@@ -1,4 +1,4 @@
|
||||
from rest_framework.views import APIView
|
||||
from adrf.views import APIView
|
||||
from rest_framework.permissions import IsAuthenticated
|
||||
from rest_framework.parsers import JSONParser, MultiPartParser, FormParser
|
||||
from rest_framework.response import Response
|
||||
@@ -10,6 +10,8 @@ from django.core.files.base import ContentFile
|
||||
from datetime import timedelta
|
||||
import uuid
|
||||
import os
|
||||
import aiohttp
|
||||
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
|
||||
@@ -37,14 +39,16 @@ class FriendRequestListView(APIView):
|
||||
],
|
||||
responses={200: success_response, 401: unauthorized_response},
|
||||
)
|
||||
def get(self, request):
|
||||
async def get(self, request):
|
||||
direction = request.query_params.get('direction', 'received')
|
||||
req_status = request.query_params.get('status', 'pending')
|
||||
if direction == 'sent':
|
||||
queryset = FriendRequest.objects.filter(from_user=request.user, status=req_status)
|
||||
else:
|
||||
queryset = FriendRequest.objects.filter(to_user=request.user, status=req_status)
|
||||
serializer = FriendRequestSerializer(queryset, many=True, context={'request': request})
|
||||
queryset = queryset.select_related('from_user', 'to_user')
|
||||
requests_list = [r async for r in queryset]
|
||||
serializer = FriendRequestSerializer(requests_list, many=True, context={'request': request})
|
||||
return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS)
|
||||
|
||||
@swagger_auto_schema(
|
||||
@@ -59,7 +63,7 @@ class FriendRequestListView(APIView):
|
||||
),
|
||||
responses={201: success_response, 400: error_response, 401: unauthorized_response, 404: not_found_response},
|
||||
)
|
||||
def post(self, request):
|
||||
async def post(self, request):
|
||||
to_user_id = request.data.get('to_user_id')
|
||||
message = request.data.get('message', '')
|
||||
if not to_user_id:
|
||||
@@ -67,24 +71,24 @@ class FriendRequestListView(APIView):
|
||||
if int(to_user_id) == request.user.id:
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='不能向自己发送好友请求', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
try:
|
||||
to_user = FUser.objects.get(pk=to_user_id)
|
||||
to_user = await FUser.objects.aget(pk=to_user_id)
|
||||
except FUser.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='用户不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
if Friendship.objects.filter(
|
||||
if await Friendship.objects.filter(
|
||||
(Q(user1=request.user) & Q(user2=to_user)) | (Q(user1=to_user) & Q(user2=request.user))
|
||||
).exists():
|
||||
).aexists():
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='已经是好友关系', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
existing = FriendRequest.objects.filter(from_user=request.user, to_user=to_user, status='pending').first()
|
||||
existing = await FriendRequest.objects.filter(from_user=request.user, to_user=to_user, status='pending').afirst()
|
||||
if existing:
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='已发送过好友请求', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
reverse = FriendRequest.objects.filter(from_user=to_user, to_user=request.user, status='pending').first()
|
||||
reverse = await FriendRequest.objects.filter(from_user=to_user, to_user=request.user, status='pending').afirst()
|
||||
if reverse:
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='对方已向你发送好友请求,请直接接受', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
friend_request = FriendRequest.objects.create(from_user=request.user, to_user=to_user, message=message)
|
||||
friend_request = await FriendRequest.objects.acreate(from_user=request.user, to_user=to_user, message=message)
|
||||
serializer = FriendRequestSerializer(friend_request, context={'request': request})
|
||||
return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS, status_code=status.HTTP_201_CREATED)
|
||||
|
||||
@@ -100,17 +104,17 @@ class FriendRequestAcceptView(APIView):
|
||||
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},
|
||||
)
|
||||
def post(self, request, pk):
|
||||
async def post(self, request, pk):
|
||||
try:
|
||||
friend_request = FriendRequest.objects.get(pk=pk, to_user=request.user, status='pending')
|
||||
friend_request = await FriendRequest.objects.aget(pk=pk, to_user_id=request.user.id, status='pending')
|
||||
except FriendRequest.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='好友请求不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
friend_request.status = 'accepted'
|
||||
friend_request.save()
|
||||
await friend_request.asave()
|
||||
|
||||
u1, u2 = sorted([friend_request.from_user, friend_request.to_user], key=lambda u: u.id)
|
||||
Friendship.objects.get_or_create(user1=u1, user2=u2)
|
||||
u1, u2 = sorted([friend_request.from_user_id, friend_request.to_user_id])
|
||||
await Friendship.objects.aget_or_create(user1_id=u1, user2_id=u2)
|
||||
|
||||
return create_standardized_response(data={'status': 'accepted'}, code=ResponseCode.SUCCESS)
|
||||
|
||||
@@ -126,14 +130,14 @@ class FriendRequestRejectView(APIView):
|
||||
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},
|
||||
)
|
||||
def post(self, request, pk):
|
||||
async def post(self, request, pk):
|
||||
try:
|
||||
friend_request = FriendRequest.objects.get(pk=pk, to_user=request.user, status='pending')
|
||||
friend_request = await FriendRequest.objects.aget(pk=pk, to_user_id=request.user.id, status='pending')
|
||||
except FriendRequest.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='好友请求不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
friend_request.status = 'rejected'
|
||||
friend_request.save()
|
||||
await friend_request.asave()
|
||||
return create_standardized_response(data={'status': 'rejected'}, code=ResponseCode.SUCCESS)
|
||||
|
||||
|
||||
@@ -148,14 +152,14 @@ class FriendRequestCancelView(APIView):
|
||||
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},
|
||||
)
|
||||
def post(self, request, pk):
|
||||
async def post(self, request, pk):
|
||||
try:
|
||||
friend_request = FriendRequest.objects.get(pk=pk, from_user=request.user, status='pending')
|
||||
friend_request = await FriendRequest.objects.aget(pk=pk, from_user_id=request.user.id, status='pending')
|
||||
except FriendRequest.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='好友请求不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
friend_request.status = 'cancelled'
|
||||
friend_request.save()
|
||||
await friend_request.asave()
|
||||
return create_standardized_response(data={'status': 'cancelled'}, code=ResponseCode.SUCCESS)
|
||||
|
||||
|
||||
@@ -169,11 +173,12 @@ class FriendListView(APIView):
|
||||
operation_description='获取当前用户的好友列表',
|
||||
responses={200: success_response, 401: unauthorized_response},
|
||||
)
|
||||
def get(self, request):
|
||||
async def get(self, request):
|
||||
friendships = Friendship.objects.filter(
|
||||
Q(user1=request.user) | Q(user2=request.user)
|
||||
).select_related('user1', 'user2')
|
||||
serializer = FriendshipSerializer(friendships, many=True, context={'request': request})
|
||||
friendships_list = [f async for f in friendships]
|
||||
serializer = FriendshipSerializer(friendships_list, many=True, context={'request': request})
|
||||
return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS)
|
||||
|
||||
@swagger_auto_schema(
|
||||
@@ -183,14 +188,14 @@ class FriendListView(APIView):
|
||||
manual_parameters=[openapi.Parameter('user_id', openapi.IN_PATH, description='目标用户ID', type=openapi.TYPE_INTEGER, required=True)],
|
||||
responses={204: success_response, 401: unauthorized_response, 404: not_found_response},
|
||||
)
|
||||
def delete(self, request, user_id):
|
||||
async def delete(self, request, user_id):
|
||||
try:
|
||||
target_user = FUser.objects.get(pk=user_id)
|
||||
target_user = await FUser.objects.aget(pk=user_id)
|
||||
except FUser.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='用户不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
u1, u2 = sorted([request.user, target_user], key=lambda u: u.id)
|
||||
deleted, _ = Friendship.objects.filter(user1=u1, user2=u2).delete()
|
||||
u1, u2 = sorted([request.user.id, target_user.id])
|
||||
deleted, _ = await Friendship.objects.filter(user1_id=u1, user2_id=u2).adelete()
|
||||
if deleted:
|
||||
return create_standardized_response(code=ResponseCode.SUCCESS, status_code=status.HTTP_204_NO_CONTENT)
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='好友关系不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
@@ -207,19 +212,19 @@ class FriendCheckView(APIView):
|
||||
manual_parameters=[openapi.Parameter('user_id', openapi.IN_PATH, description='目标用户ID', type=openapi.TYPE_INTEGER, required=True)],
|
||||
responses={200: success_response, 401: unauthorized_response, 404: not_found_response},
|
||||
)
|
||||
def get(self, request, user_id):
|
||||
async def get(self, request, user_id):
|
||||
try:
|
||||
target_user = FUser.objects.get(pk=user_id)
|
||||
target_user = await FUser.objects.aget(pk=user_id)
|
||||
except FUser.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='用户不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
u1, u2 = sorted([request.user, target_user], key=lambda u: u.id)
|
||||
is_friend = Friendship.objects.filter(user1=u1, user2=u2).exists()
|
||||
u1, u2 = sorted([request.user.id, target_user.id])
|
||||
is_friend = await Friendship.objects.filter(user1_id=u1, user2_id=u2).aexists()
|
||||
|
||||
pending_request = FriendRequest.objects.filter(
|
||||
(Q(from_user=request.user, to_user=target_user) | Q(from_user=target_user, to_user=request.user)),
|
||||
pending_request = await FriendRequest.objects.filter(
|
||||
(Q(from_user_id=request.user.id, to_user_id=target_user.id) | Q(from_user_id=target_user.id, to_user_id=request.user.id)),
|
||||
status='pending'
|
||||
).first()
|
||||
).afirst()
|
||||
|
||||
return create_standardized_response(data={
|
||||
'is_friend': is_friend,
|
||||
@@ -238,7 +243,7 @@ class UserSearchView(APIView):
|
||||
manual_parameters=[openapi.Parameter('q', openapi.IN_QUERY, description='搜索关键词', type=openapi.TYPE_STRING)],
|
||||
responses={200: success_response, 401: unauthorized_response},
|
||||
)
|
||||
def get(self, request):
|
||||
async def get(self, request):
|
||||
q = request.query_params.get('q', '').strip()
|
||||
if not q:
|
||||
return create_standardized_response(data=[], code=ResponseCode.SUCCESS)
|
||||
@@ -249,13 +254,13 @@ class UserSearchView(APIView):
|
||||
|
||||
friend_ids = set()
|
||||
friendships = Friendship.objects.filter(
|
||||
Q(user1=request.user) | Q(user2=request.user)
|
||||
)
|
||||
for f in friendships:
|
||||
friend_ids.add(f.user2_id if f.user1 == request.user else f.user1_id)
|
||||
Q(user1_id=request.user.id) | Q(user2_id=request.user.id)
|
||||
).values_list('user1_id', 'user2_id')
|
||||
async for user1_id, user2_id in friendships:
|
||||
friend_ids.add(user2_id if user1_id == request.user.id else user1_id)
|
||||
|
||||
results = []
|
||||
for user in users:
|
||||
async for user in users:
|
||||
avatar_url = ''
|
||||
if user.avatar and hasattr(user.avatar, 'url'):
|
||||
avatar_url = request.build_absolute_uri(user.avatar.url)
|
||||
@@ -280,13 +285,15 @@ class ConversationListView(APIView):
|
||||
operation_description='获取当前用户参与的所有会话',
|
||||
responses={200: success_response, 401: unauthorized_response},
|
||||
)
|
||||
def get(self, request):
|
||||
async def get(self, request):
|
||||
participations = ConversationParticipant.objects.filter(
|
||||
user=request.user
|
||||
).select_related('conversation').order_by('-conversation__created_at')
|
||||
conversations = [p.conversation for p in participations]
|
||||
conversations = [p.conversation async for p in participations]
|
||||
serializer = ConversationSerializer(conversations, many=True, context={'request': request})
|
||||
return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS)
|
||||
# 兜底:ConversationSerializer 的 SerializerMethodField 内部有同步 ORM 查询
|
||||
data = await sync_to_async(lambda: serializer.data)()
|
||||
return create_standardized_response(data=data, code=ResponseCode.SUCCESS)
|
||||
|
||||
@swagger_auto_schema(
|
||||
tags=['AI对话'],
|
||||
@@ -295,7 +302,7 @@ class ConversationListView(APIView):
|
||||
request_body=openapi.Schema(type=openapi.TYPE_OBJECT, properties={'user_id': openapi.Schema(type=openapi.TYPE_INTEGER, description='目标用户ID')}),
|
||||
responses={201: success_response, 400: error_response, 401: unauthorized_response, 404: not_found_response},
|
||||
)
|
||||
def post(self, request):
|
||||
async def post(self, request):
|
||||
user_id = request.data.get('user_id')
|
||||
if not user_id:
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='user_id 不能为空', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
@@ -303,30 +310,32 @@ class ConversationListView(APIView):
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='不能和自己聊天', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
try:
|
||||
target_user = FUser.objects.get(pk=user_id)
|
||||
target_user = await FUser.objects.aget(pk=user_id)
|
||||
except FUser.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='用户不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
u1, u2 = sorted([request.user, target_user], key=lambda u: u.id)
|
||||
if not Friendship.objects.filter(user1=u1, user2=u2).exists():
|
||||
u1, u2 = sorted([request.user.id, target_user.id])
|
||||
if not await Friendship.objects.filter(user1_id=u1, user2_id=u2).aexists():
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='只能与好友聊天', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
my_participations = ConversationParticipant.objects.filter(
|
||||
user=request.user, conversation__type='private'
|
||||
).values_list('conversation_id', flat=True)
|
||||
existing = ConversationParticipant.objects.filter(
|
||||
existing = await ConversationParticipant.objects.filter(
|
||||
user=target_user, conversation_id__in=my_participations, conversation__type='private'
|
||||
).first()
|
||||
).select_related('conversation').afirst()
|
||||
|
||||
if existing:
|
||||
conversation = existing.conversation
|
||||
else:
|
||||
conversation = Conversation.objects.create(type='private')
|
||||
ConversationParticipant.objects.create(conversation=conversation, user=request.user)
|
||||
ConversationParticipant.objects.create(conversation=conversation, user=target_user)
|
||||
conversation = await Conversation.objects.acreate(type='private')
|
||||
await ConversationParticipant.objects.acreate(conversation=conversation, user=request.user)
|
||||
await ConversationParticipant.objects.acreate(conversation=conversation, user=target_user)
|
||||
|
||||
serializer = ConversationSerializer(conversation, context={'request': request})
|
||||
return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS, status_code=status.HTTP_201_CREATED)
|
||||
# 兜底:ConversationSerializer 的 SerializerMethodField 内部有同步 ORM 查询
|
||||
data = await sync_to_async(lambda: serializer.data)()
|
||||
return create_standardized_response(data=data, code=ResponseCode.SUCCESS, status_code=status.HTTP_201_CREATED)
|
||||
|
||||
|
||||
class ConversationMessageView(APIView):
|
||||
@@ -344,13 +353,13 @@ class ConversationMessageView(APIView):
|
||||
],
|
||||
responses={200: success_response, 401: unauthorized_response, 404: not_found_response},
|
||||
)
|
||||
def get(self, request, pk):
|
||||
async def get(self, request, pk):
|
||||
try:
|
||||
conversation = Conversation.objects.get(pk=pk)
|
||||
conversation = await Conversation.objects.aget(pk=pk)
|
||||
except Conversation.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='会话不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
if not ConversationParticipant.objects.filter(conversation=conversation, user=request.user).exists():
|
||||
if not await ConversationParticipant.objects.filter(conversation=conversation, user=request.user).aexists():
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='你不是该会话的参与者', status_code=status.HTTP_403_FORBIDDEN)
|
||||
|
||||
page = int(request.query_params.get('page', 1))
|
||||
@@ -358,12 +367,14 @@ class ConversationMessageView(APIView):
|
||||
offset = (page - 1) * page_size
|
||||
|
||||
messages = ChatMessage.objects.filter(conversation=conversation).select_related('sender', 'reply_to', 'reply_to__sender')
|
||||
total = messages.count()
|
||||
messages = messages[offset:offset + page_size]
|
||||
total = await messages.acount()
|
||||
page_messages = [m async for m in messages[offset:offset + page_size]]
|
||||
|
||||
serializer = ChatMessageSerializer(messages, many=True, context={'request': request})
|
||||
serializer = ChatMessageSerializer(page_messages, many=True, context={'request': request})
|
||||
# 兜底:'conversation' 外键未预加载,serializer 取值会触发同步 ORM
|
||||
data = await sync_to_async(lambda: serializer.data)()
|
||||
return create_standardized_response(data={
|
||||
'results': serializer.data,
|
||||
'results': data,
|
||||
'total': total,
|
||||
'page': page,
|
||||
'page_size': page_size,
|
||||
@@ -385,13 +396,13 @@ class ConversationMessageView(APIView):
|
||||
),
|
||||
responses={201: success_response, 400: error_response, 401: unauthorized_response, 404: not_found_response},
|
||||
)
|
||||
def post(self, request, pk):
|
||||
async def post(self, request, pk):
|
||||
try:
|
||||
conversation = Conversation.objects.get(pk=pk)
|
||||
conversation = await Conversation.objects.aget(pk=pk)
|
||||
except Conversation.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='会话不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
if not ConversationParticipant.objects.filter(conversation=conversation, user=request.user).exists():
|
||||
if not await ConversationParticipant.objects.filter(conversation=conversation, user=request.user).aexists():
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='你不是该会话的参与者', status_code=status.HTTP_403_FORBIDDEN)
|
||||
|
||||
content = request.data.get('content', '').strip()
|
||||
@@ -405,11 +416,11 @@ class ConversationMessageView(APIView):
|
||||
reply_to = None
|
||||
if reply_to_id:
|
||||
try:
|
||||
reply_to = ChatMessage.objects.get(pk=reply_to_id, conversation=conversation)
|
||||
reply_to = await ChatMessage.objects.aget(pk=reply_to_id, conversation=conversation)
|
||||
except ChatMessage.DoesNotExist:
|
||||
pass
|
||||
|
||||
message = ChatMessage.objects.create(
|
||||
message = await ChatMessage.objects.acreate(
|
||||
conversation=conversation,
|
||||
sender=request.user,
|
||||
content=content,
|
||||
@@ -433,17 +444,17 @@ class ConversationClearView(APIView):
|
||||
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},
|
||||
)
|
||||
def delete(self, request, pk):
|
||||
async def delete(self, request, pk):
|
||||
try:
|
||||
conversation = Conversation.objects.get(pk=pk)
|
||||
conversation = await Conversation.objects.aget(pk=pk)
|
||||
except Conversation.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='会话不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
if not ConversationParticipant.objects.filter(conversation=conversation, user=request.user).exists():
|
||||
if not await ConversationParticipant.objects.filter(conversation=conversation, user=request.user).aexists():
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='你不是该会话的参与者', status_code=status.HTTP_403_FORBIDDEN)
|
||||
|
||||
# 物理删除该会话下所有消息
|
||||
deleted_count, _ = ChatMessage.objects.filter(conversation=conversation).delete()
|
||||
deleted_count, _ = await ChatMessage.objects.filter(conversation=conversation).adelete()
|
||||
|
||||
return create_standardized_response(data={'deleted': deleted_count}, code=ResponseCode.SUCCESS)
|
||||
permission_classes = [IsAuthenticated]
|
||||
@@ -456,14 +467,14 @@ class ConversationClearView(APIView):
|
||||
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},
|
||||
)
|
||||
def post(self, request, pk):
|
||||
async def post(self, request, pk):
|
||||
try:
|
||||
participant = ConversationParticipant.objects.get(conversation_id=pk, user=request.user)
|
||||
participant = await ConversationParticipant.objects.aget(conversation_id=pk, user=request.user)
|
||||
except ConversationParticipant.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='会话不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
participant.last_read_at = timezone.now()
|
||||
participant.save(update_fields=['last_read_at'])
|
||||
await participant.asave(update_fields=['last_read_at'])
|
||||
return create_standardized_response(data={'read': True}, code=ResponseCode.SUCCESS)
|
||||
|
||||
|
||||
@@ -478,13 +489,13 @@ class MessageRecallView(APIView):
|
||||
manual_parameters=[openapi.Parameter('pk', openapi.IN_PATH, description='消息ID', type=openapi.TYPE_INTEGER, required=True)],
|
||||
responses={200: success_response, 400: error_response, 401: unauthorized_response, 404: not_found_response},
|
||||
)
|
||||
def post(self, request, pk):
|
||||
async def post(self, request, pk):
|
||||
try:
|
||||
message = ChatMessage.objects.get(pk=pk)
|
||||
message = await ChatMessage.objects.aget(pk=pk)
|
||||
except ChatMessage.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='消息不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
if message.sender != request.user:
|
||||
if message.sender_id != request.user.id:
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='只能撤回自己发送的消息', status_code=status.HTTP_403_FORBIDDEN)
|
||||
|
||||
if timezone.now() - message.created_at > timedelta(minutes=2):
|
||||
@@ -492,7 +503,7 @@ class MessageRecallView(APIView):
|
||||
|
||||
message.is_recalled = True
|
||||
message.content = '你撤回了一条消息'
|
||||
message.save(update_fields=['is_recalled', 'content'])
|
||||
await message.asave(update_fields=['is_recalled', 'content'])
|
||||
|
||||
return create_standardized_response(data={'recalled': True}, code=ResponseCode.SUCCESS)
|
||||
|
||||
@@ -506,18 +517,18 @@ class MessageDeleteView(APIView):
|
||||
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, 403: error_response, 404: not_found_response},
|
||||
responses={200: success_response, 400: error_response, 401: unauthorized_response, 403: error_response, 404: not_found_response},
|
||||
)
|
||||
def delete(self, request, pk):
|
||||
async def delete(self, request, pk):
|
||||
try:
|
||||
message = ChatMessage.objects.get(pk=pk)
|
||||
message = await ChatMessage.objects.aget(pk=pk)
|
||||
except ChatMessage.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='消息不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
if message.sender != request.user:
|
||||
if message.sender_id != request.user.id:
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='只能删除自己发送的消息', status_code=status.HTTP_403_FORBIDDEN)
|
||||
|
||||
message.delete()
|
||||
await message.adelete()
|
||||
return create_standardized_response(data={'deleted': True}, code=ResponseCode.SUCCESS)
|
||||
|
||||
|
||||
@@ -532,7 +543,7 @@ class FileUploadView(APIView):
|
||||
manual_parameters=[openapi.Parameter('file', openapi.IN_FORM, description='文件', type=openapi.TYPE_FILE, required=True)],
|
||||
responses={200: success_response, 400: error_response, 401: unauthorized_response},
|
||||
)
|
||||
def post(self, request):
|
||||
async def post(self, request):
|
||||
file = request.FILES.get('file')
|
||||
if not file:
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='请选择文件', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
@@ -543,7 +554,7 @@ class FileUploadView(APIView):
|
||||
filepath = f'chat_files/{date_path}/{filename}'
|
||||
|
||||
from django.core.files.storage import default_storage
|
||||
saved_path = default_storage.save(filepath, file)
|
||||
saved_path = await sync_to_async(default_storage.save)(filepath, file)
|
||||
url = request.build_absolute_uri(settings.MEDIA_URL + saved_path)
|
||||
|
||||
return create_standardized_response(data={
|
||||
@@ -564,7 +575,7 @@ class StickerUploadView(APIView):
|
||||
manual_parameters=[openapi.Parameter('file', openapi.IN_FORM, description='图片文件', type=openapi.TYPE_FILE, required=True)],
|
||||
responses={200: success_response, 400: error_response, 401: unauthorized_response},
|
||||
)
|
||||
def post(self, request):
|
||||
async def post(self, request):
|
||||
file = request.FILES.get('file')
|
||||
if not file:
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='请选择文件', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
@@ -577,9 +588,9 @@ class StickerUploadView(APIView):
|
||||
filepath = f'chat_stickers/{request.user.id}/{filename}'
|
||||
|
||||
from django.core.files.storage import default_storage
|
||||
saved_path = default_storage.save(filepath, file)
|
||||
saved_path = await sync_to_async(default_storage.save)(filepath, file)
|
||||
|
||||
sticker = FavoriteSticker.objects.create(
|
||||
sticker = await FavoriteSticker.objects.acreate(
|
||||
user=request.user,
|
||||
image=saved_path,
|
||||
)
|
||||
@@ -600,10 +611,10 @@ class FavoriteStickerView(APIView):
|
||||
operation_summary='获取收藏表情列表',
|
||||
responses={200: success_response, 401: unauthorized_response},
|
||||
)
|
||||
def get(self, request):
|
||||
async def get(self, request):
|
||||
stickers = FavoriteSticker.objects.filter(user=request.user)
|
||||
data = []
|
||||
for s in stickers:
|
||||
async for s in stickers:
|
||||
data.append({
|
||||
'id': s.id,
|
||||
'url': request.build_absolute_uri(s.image.url) if s.image else '',
|
||||
@@ -616,7 +627,7 @@ class FavoriteStickerView(APIView):
|
||||
manual_parameters=[openapi.Parameter('image', openapi.IN_FORM, description='图片文件', type=openapi.TYPE_FILE, required=True)],
|
||||
responses={200: success_response, 400: error_response, 401: unauthorized_response},
|
||||
)
|
||||
def post(self, request):
|
||||
async def post(self, request):
|
||||
file = request.FILES.get('image')
|
||||
if not file:
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='请选择图片', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
@@ -629,9 +640,9 @@ class FavoriteStickerView(APIView):
|
||||
filepath = f'chat_stickers/{request.user.id}/{filename}'
|
||||
|
||||
from django.core.files.storage import default_storage
|
||||
saved_path = default_storage.save(filepath, file)
|
||||
saved_path = await sync_to_async(default_storage.save)(filepath, file)
|
||||
|
||||
sticker = FavoriteSticker.objects.create(
|
||||
sticker = await FavoriteSticker.objects.acreate(
|
||||
user=request.user,
|
||||
image=saved_path,
|
||||
)
|
||||
@@ -648,10 +659,10 @@ class FavoriteStickerView(APIView):
|
||||
manual_parameters=[openapi.Parameter('pk', openapi.IN_PATH, description='表情ID', type=openapi.TYPE_INTEGER, required=True)],
|
||||
responses={204: success_response, 401: unauthorized_response, 404: not_found_response},
|
||||
)
|
||||
def delete(self, request, pk):
|
||||
async def delete(self, request, pk):
|
||||
try:
|
||||
sticker = FavoriteSticker.objects.get(pk=pk, user=request.user)
|
||||
sticker.delete()
|
||||
sticker = await FavoriteSticker.objects.aget(pk=pk, user=request.user)
|
||||
await sticker.adelete()
|
||||
return Response(status=status.HTTP_204_NO_CONTENT)
|
||||
except FavoriteSticker.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='收藏不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
@@ -673,13 +684,13 @@ class FavoriteStickerFromMessageView(APIView):
|
||||
),
|
||||
responses={200: success_response, 400: error_response, 401: unauthorized_response, 404: not_found_response},
|
||||
)
|
||||
def post(self, request):
|
||||
async def post(self, request):
|
||||
message_id = request.data.get('message_id')
|
||||
if not message_id:
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='message_id 不能为空', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
try:
|
||||
message = ChatMessage.objects.get(pk=message_id)
|
||||
message = await ChatMessage.objects.aget(pk=message_id)
|
||||
except ChatMessage.DoesNotExist:
|
||||
return create_standardized_error_response(code=ResponseCode.NOT_FOUND, message='消息不存在', status_code=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
@@ -688,15 +699,18 @@ class FavoriteStickerFromMessageView(APIView):
|
||||
return create_standardized_error_response(code=ResponseCode.VALIDATION_ERROR, message='该消息没有可收藏的图片', status_code=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
try:
|
||||
import urllib.request
|
||||
if file_url.startswith('http'):
|
||||
req = urllib.request.Request(file_url, headers={'User-Agent': 'Mozilla/5.0'})
|
||||
response = urllib.request.urlopen(req, timeout=10)
|
||||
image_data = response.read()
|
||||
async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=10)) as session:
|
||||
async with session.get(file_url, headers={'User-Agent': 'Mozilla/5.0'}) as resp:
|
||||
image_data = await resp.read()
|
||||
else:
|
||||
local_path = os.path.join(settings.MEDIA_ROOT, file_url.replace(settings.MEDIA_URL, ''))
|
||||
with open(local_path, 'rb') as f:
|
||||
image_data = f.read()
|
||||
|
||||
def _read_local(p):
|
||||
with open(p, 'rb') as f:
|
||||
return f.read()
|
||||
|
||||
image_data = await sync_to_async(_read_local)(local_path)
|
||||
|
||||
ext = '.png'
|
||||
if '.' in file_url.split('/')[-1]:
|
||||
@@ -705,9 +719,9 @@ class FavoriteStickerFromMessageView(APIView):
|
||||
filepath = f'chat_stickers/{request.user.id}/{filename}'
|
||||
|
||||
from django.core.files.storage import default_storage
|
||||
saved_path = default_storage.save(filepath, ContentFile(image_data))
|
||||
saved_path = await sync_to_async(default_storage.save)(filepath, ContentFile(image_data))
|
||||
|
||||
sticker = FavoriteSticker.objects.create(
|
||||
sticker = await FavoriteSticker.objects.acreate(
|
||||
user=request.user,
|
||||
image=saved_path,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user