From 8f488fcaaac71ed16e2171bd3c9a8bfc899bb96c Mon Sep 17 00:00:00 2001 From: async-upgrade Date: Sun, 6 Sep 2026 14:26:17 +0800 Subject: [PATCH] feat: ADRF async views (phase1) + native async serializers (phase2) + async cache infra --- Dockerfile | 8 +- air_quality/views.py | 107 ++++++----- api/views/BaiduFanyiView.py | 62 +++---- api/views/GetIPDataView.py | 48 +++-- apidirectory/serializers.py | 32 ++-- apidirectory/views.py | 25 +-- app/views.py | 13 +- article/serializers.py | 182 ++++++++++--------- article/views.py | 211 ++++++++++++---------- bug/serializers.py | 82 +++++---- bug/views.py | 49 +++-- chat/serializers.py | 95 +++++----- chat/views.py | 222 ++++++++++++----------- chunyu_project/settings.py | 71 ++++++-- currency/views.py | 113 +++++++----- docs/ADRF_CONVERSION_GUIDE.md | 74 ++++++++ docs/ADRF_SERIALIZER_GUIDE.md | 57 ++++++ docs/ARCHITECTURE.md | 151 ++++++++++++++++ history/serializers.py | 16 +- history/views.py | 35 ++-- learn/serializers.py | 96 +++++----- learn/views.py | 239 ++++++++++++++----------- logs/views.py | 55 +++--- message/serializers.py | 78 ++++---- message/views.py | 110 +++++++----- requirements.txt | 6 + search/views.py | 26 +-- shorturl/views.py | 28 +-- tool/serializers.py | 57 +++--- tool/views/color_history_view.py | 46 +++-- tool/views/compression_history_view.py | 28 +-- tool/views/image_compress_view.py | 2 +- tool/views/text_diff_view.py | 2 +- tool/views/tool_favorite_view.py | 107 ++++++----- tool/views/tool_manage_view.py | 57 +++--- user/serializers/region_serializers.py | 17 +- user/serializers/user_serializers.py | 224 ++++++++++++----------- user/views/activities.py | 10 +- user/views/blacklist.py | 29 +-- user/views/captcha.py | 2 +- user/views/email.py | 34 ++-- user/views/favorites.py | 22 +-- user/views/index.py | 2 +- user/views/invitation.py | 16 +- user/views/login_record.py | 19 +- user/views/phone.py | 17 +- user/views/qr_login.py | 32 ++-- user/views/region.py | 7 +- user/views/settings.py | 36 ++-- user/views/slider_captcha.py | 2 +- user/views/tasks.py | 134 +++++++------- user/views/user.py | 218 +++++++++++----------- user/views/wallet.py | 100 +++++++---- utils/async_cache.py | 114 ++++++++++++ weather/views.py | 112 ++++++------ 55 files changed, 2224 insertions(+), 1513 deletions(-) create mode 100644 docs/ADRF_CONVERSION_GUIDE.md create mode 100644 docs/ADRF_SERIALIZER_GUIDE.md create mode 100644 docs/ARCHITECTURE.md create mode 100644 utils/async_cache.py diff --git a/Dockerfile b/Dockerfile index e29a1a9..f5695ab 100644 --- a/Dockerfile +++ b/Dockerfile @@ -21,7 +21,11 @@ COPY . . # 暴露端口 EXPOSE 8000 -# 启动命令 - 运行迁移并启动服务 +# 启动命令 - 运行迁移、收集静态文件,并使用 Granian(Rust) 以 ASGI 接口启动 +# Granian 性能远超 Daphne,原生支持 ASGI/WebSocket(Channels ProtocolTypeRouter) +# 注意:ASGI 接口下 blocking-threads 必须为 1(Granian 限制),并发由事件循环承担 CMD python manage.py migrate --noinput && \ python manage.py collectstatic --noinput && \ - daphne -b 0.0.0.0 -p 8000 chunyu_project.asgi:application + granian --interface asgi --host 0.0.0.0 --port 8000 \ + --workers 2 --backlog 2048 \ + --log-level info chunyu_project.asgi:application diff --git a/air_quality/views.py b/air_quality/views.py index 254bb69..ff42753 100644 --- a/air_quality/views.py +++ b/air_quality/views.py @@ -1,8 +1,10 @@ -import requests +import asyncio from datetime import datetime +import aiohttp + from rest_framework.permissions import AllowAny -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from drf_yasg.utils import swagger_auto_schema @@ -58,7 +60,7 @@ class AirQualityView(APIView): ], responses={200: success_response, 400: error_response, 404: error_response, 502: error_response} ) - def get(self, request): + async def get(self, request): city = request.GET.get('city', '').strip() lang = request.GET.get('lang', 'zh_cn').strip() @@ -68,7 +70,7 @@ class AirQualityView(APIView): status=status.HTTP_400_BAD_REQUEST ) - return self._fetch_air_quality(city, lang) + return await self._fetch_air_quality(city, lang) @swagger_auto_schema( tags=['Air Quality'], @@ -84,7 +86,7 @@ class AirQualityView(APIView): ), responses={200: success_response, 400: error_response, 404: error_response, 502: error_response} ) - def post(self, request): + async def post(self, request): city = request.data.get('city', '').strip() if isinstance(request.data, dict) else '' lang = request.data.get('lang', 'zh_cn').strip() if isinstance(request.data, dict) else 'zh_cn' @@ -94,60 +96,60 @@ class AirQualityView(APIView): status=status.HTTP_400_BAD_REQUEST ) - return self._fetch_air_quality(city, lang) + return await self._fetch_air_quality(city, lang) - def _fetch_air_quality(self, city, lang): + async def _fetch_air_quality(self, city, lang): """ - 调用 Open-Meteo Air Quality API 获取空气质量数据 + 调用 Open-Meteo Air Quality API 获取空气质量数据(aiohttp 全异步) """ + timeout = aiohttp.ClientTimeout(total=10) try: - # 第一步:通过地理编码API获取城市坐标 - geo_params = { - 'name': city, - 'count': 1, - 'language': 'zh' if lang == 'zh_cn' else 'en', - 'format': 'json' - } - geo_response = requests.get(GEOCODING_URL, params=geo_params, timeout=10) + async with aiohttp.ClientSession(timeout=timeout) as session: + # 第一步:通过地理编码API获取城市坐标 + geo_params = { + 'name': city, + 'count': 1, + 'language': 'zh' if lang == 'zh_cn' else 'en', + 'format': 'json' + } + async with session.get(GEOCODING_URL, params=geo_params) as geo_response: + if geo_response.status != 200: + return Response( + {"code": 502, "message": "地理编码服务异常,请稍后重试", "data": None}, + status=status.HTTP_502_BAD_GATEWAY + ) + geo_data = await geo_response.json() - if geo_response.status_code != 200: - return Response( - {"code": 502, "message": "地理编码服务异常,请稍后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY - ) + results = geo_data.get('results', []) - geo_data = geo_response.json() - results = geo_data.get('results', []) + if not results: + return Response( + {"code": 404, "message": f"未找到城市:{city},请检查城市名称是否正确", "data": None}, + status=status.HTTP_404_NOT_FOUND + ) - if not results: - return Response( - {"code": 404, "message": f"未找到城市:{city},请检查城市名称是否正确", "data": None}, - status=status.HTTP_404_NOT_FOUND - ) + location = results[0] + latitude = location.get('latitude') + longitude = location.get('longitude') + city_name = location.get('name', city) + country = location.get('country', '') - location = results[0] - latitude = location.get('latitude') - longitude = location.get('longitude') - city_name = location.get('name', city) - country = location.get('country', '') + # 第二步:获取空气质量数据 + air_quality_params = { + 'latitude': latitude, + 'longitude': longitude, + 'current': 'pm10,pm2_5,carbon_monoxide,nitrogen_dioxide,sulphur_dioxide,ozone,us_aqi', + 'timezone': 'auto' + } - # 第二步:获取空气质量数据 - air_quality_params = { - 'latitude': latitude, - 'longitude': longitude, - 'current': 'pm10,pm2_5,carbon_monoxide,nitrogen_dioxide,sulphur_dioxide,ozone,us_aqi', - 'timezone': 'auto' - } + async with session.get(AIR_QUALITY_URL, params=air_quality_params) as aq_response: + if aq_response.status != 200: + return Response( + {"code": 502, "message": "空气质量服务异常,请稍后重试", "data": None}, + status=status.HTTP_502_BAD_GATEWAY + ) + aq_data = await aq_response.json() - aq_response = requests.get(AIR_QUALITY_URL, params=air_quality_params, timeout=10) - - if aq_response.status_code != 200: - return Response( - {"code": 502, "message": "空气质量服务异常,请稍后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY - ) - - aq_data = aq_response.json() current = aq_data.get('current', {}) pm2_5 = current.get('pm2_5', 0) @@ -171,16 +173,11 @@ class AirQualityView(APIView): status=status.HTTP_200_OK ) - except requests.exceptions.Timeout: + except (aiohttp.ClientError, asyncio.TimeoutError): return Response( {"code": 502, "message": "请求超时,请稍后重试", "data": None}, status=status.HTTP_502_BAD_GATEWAY ) - except requests.exceptions.ConnectionError: - return Response( - {"code": 502, "message": "网络连接异常,请检查网络后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY - ) except Exception as e: return Response( {"code": 500, "message": f"服务器内部错误:{str(e)}", "data": None}, diff --git a/api/views/BaiduFanyiView.py b/api/views/BaiduFanyiView.py index 2ed6d1f..d2a3a25 100644 --- a/api/views/BaiduFanyiView.py +++ b/api/views/BaiduFanyiView.py @@ -1,9 +1,10 @@ import base64 +from asgiref.sync import sync_to_async from django.utils import timezone from rest_framework.parsers import MultiPartParser, FormParser from rest_framework.permissions import AllowAny -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from drf_yasg.utils import swagger_auto_schema @@ -61,7 +62,7 @@ class BaiduFanyiView(APIView): ), responses={202: success_response, 400: error_response} ) - def post(self, request): + async def post(self, request): from_lang = request.data.get('from_lang') to_lang = request.data.get('to_lang') @@ -95,15 +96,16 @@ class BaiduFanyiView(APIView): ) # 优先同步调用百度 API,成功后直接返回翻译结果,不依赖 Celery worker + # (百度 API 调用为阻塞网络 I/O,sync_to_async 兜底) try: - sync_result = baidu_translate_task(query, from_lang, to_lang) + sync_result = await sync_to_async(baidu_translate_task)(query, from_lang, to_lang) except Exception as e: sync_result = {'success': False, 'error_code': 'INNER_ERROR', 'error_msg': str(e)} if sync_result.get('success'): baidu_data = sync_result.get('data') or {} user_id = request.user.id if request.user.is_authenticated else None - record_translate_usage.delay('text', from_lang, to_lang, len(query), user_id=user_id, ip_address=get_client_ip(request)) + await sync_to_async(record_translate_usage.delay)('text', from_lang, to_lang, len(query), user_id=user_id, ip_address=get_client_ip(request)) return Response({ "code": 0, "data": { @@ -115,10 +117,10 @@ class BaiduFanyiView(APIView): # 同步失败降级为异步(若 Celery worker 可用) try: - task = submit_task(baidu_translate_task, query, from_lang, to_lang) + task = await sync_to_async(submit_task)(baidu_translate_task, query, from_lang, to_lang) task_id = task.id if task is not None else None user_id = request.user.id if request.user.is_authenticated else None - record_translate_usage.delay('text', from_lang, to_lang, len(query), user_id=user_id, ip_address=get_client_ip(request)) + await sync_to_async(record_translate_usage.delay)('text', from_lang, to_lang, len(query), user_id=user_id, ip_address=get_client_ip(request)) return Response( { "task_id": task_id, @@ -143,7 +145,7 @@ class AutoLangTypeViews(APIView): operation_description='返回支持自动识别的语言类型列表', responses={200: success_response} ) - def get(self, request): + async def get(self, request): return Response(auto_lang, status=status.HTTP_200_OK) @@ -156,7 +158,7 @@ class AllLangTypeViews(APIView): operation_description='返回所有支持的翻译语言类型列表', responses={200: success_response} ) - def get(self, request): + async def get(self, request): return Response(text_languages_flat, status=status.HTTP_200_OK) @@ -169,7 +171,7 @@ class PictureLangTypeViews(APIView): operation_description='返回图片翻译支持的源语言和目标语言类型列表', responses={200: success_response} ) - def get(self, request): + async def get(self, request): return Response({ "source_langs": p_lang_type, "target_langs": p_target_lang_type @@ -188,7 +190,7 @@ class SpeechLangTypeViews(APIView): ], responses={200: success_response, 400: error_response} ) - def get(self, request): + async def get(self, request): speech_type = request.GET.get("speech_type", "mp3") if speech_type not in y_speech_type: @@ -222,7 +224,7 @@ class RecognizeLangTypeViews(APIView): ), responses={202: success_response, 400: error_response} ) - def post(self, request): + async def post(self, request): query = request.data.get("q") if not query: @@ -231,9 +233,9 @@ class RecognizeLangTypeViews(APIView): if len(query) > 3000: return Response({"message": "请求数据长度过长已经超过3000."}, status=status.HTTP_400_BAD_REQUEST) - # 优先同步调用百度 API,避免依赖 Celery worker + # 优先同步调用百度 API,避免依赖 Celery worker(sync_to_async 兜底阻塞 I/O) try: - sync_result = baidu_recognize_lang_task(query) + sync_result = await sync_to_async(baidu_recognize_lang_task)(query) except Exception as e: sync_result = {'success': False, 'error_code': 'INNER_ERROR', 'error_msg': str(e)} @@ -251,7 +253,7 @@ class RecognizeLangTypeViews(APIView): src_lang = baidu_data.get('src') trans_result = [{'src': query, 'dst': src_lang}] if src_lang else [] user_id = request.user.id if request.user.is_authenticated else None - record_translate_usage.delay('detect', 'auto', '', len(query), user_id=user_id, ip_address=get_client_ip(request)) + await sync_to_async(record_translate_usage.delay)('detect', 'auto', '', len(query), user_id=user_id, ip_address=get_client_ip(request)) return Response({ "code": 0, "data": { @@ -262,10 +264,10 @@ class RecognizeLangTypeViews(APIView): }, status=status.HTTP_200_OK) try: - task = submit_task(baidu_recognize_lang_task, query) + task = await sync_to_async(submit_task)(baidu_recognize_lang_task, query) task_id = task.id if task is not None else None user_id = request.user.id if request.user.is_authenticated else None - record_translate_usage.delay('detect', 'auto', '', len(query), user_id=user_id, ip_address=get_client_ip(request)) + await sync_to_async(record_translate_usage.delay)('detect', 'auto', '', len(query), user_id=user_id, ip_address=get_client_ip(request)) return Response( { "task_id": task_id, @@ -297,7 +299,7 @@ class PictureRecognizeViews(APIView): ], responses={202: success_response, 400: error_response} ) - def post(self, request): + async def post(self, request): file_data = request.FILES['file'].read() file_data_base64 = base64.b64encode(file_data).decode('utf-8') @@ -311,7 +313,7 @@ class PictureRecognizeViews(APIView): return Response({"message": "翻译语种异常."}, status=status.HTTP_400_BAD_REQUEST) try: - task = submit_task(baidu_picture_translate_task, file_data_base64, from_lang, to_lang, picture_type) + task = await sync_to_async(submit_task)(baidu_picture_translate_task, file_data_base64, from_lang, to_lang, picture_type) except Exception as e: return Response( {"code": 1, "message": "图片翻译任务提交失败,请稍后重试"}, @@ -320,7 +322,7 @@ class PictureRecognizeViews(APIView): task_id = task.id if task is not None else None user_id = request.user.id if request.user.is_authenticated else None - record_translate_usage.delay('image', from_lang, to_lang, len(file_data), user_id=user_id, ip_address=get_client_ip(request)) + await sync_to_async(record_translate_usage.delay)('image', from_lang, to_lang, len(file_data), user_id=user_id, ip_address=get_client_ip(request)) return Response( { @@ -348,7 +350,7 @@ class SpeechRecognitionView(APIView): ], responses={202: success_response, 400: error_response} ) - def post(self, request): + async def post(self, request): speech_type = request.GET.get("speech_type") if speech_type not in y_speech_type: @@ -367,7 +369,7 @@ class SpeechRecognitionView(APIView): audio_data_base64 = base64.b64encode(voice).decode('utf-8') try: - task = submit_task(baidu_speech_recognize_task, audio_data_base64, from_lang, to_lang, speech_type) + task = await sync_to_async(submit_task)(baidu_speech_recognize_task, audio_data_base64, from_lang, to_lang, speech_type) except Exception as e: return Response( {"code": 1, "message": "语音识别任务提交失败,请稍后重试"}, @@ -376,7 +378,7 @@ class SpeechRecognitionView(APIView): task_id = task.id if task is not None else None user_id = request.user.id if request.user.is_authenticated else None - record_translate_usage.delay('audio', from_lang, to_lang, len(voice), user_id=user_id, ip_address=get_client_ip(request)) + await sync_to_async(record_translate_usage.delay)('audio', from_lang, to_lang, len(voice), user_id=user_id, ip_address=get_client_ip(request)) return Response( { @@ -397,16 +399,16 @@ class TaskResultView(APIView): operation_description='通过 task_id 查询异步翻译任务的执行结果', responses={200: '成功返回翻译结果'} ) - def get(self, request, task_id): + async def get(self, request, task_id): try: - task_result = AsyncResult(task_id) + task_result = await sync_to_async(AsyncResult)(task_id) except Exception: return Response( {"code": 1, "message": "任务查询失败"}, status=status.HTTP_400_BAD_REQUEST ) - if not task_result.ready(): + if not await sync_to_async(task_result.ready)(): return Response({ "code": 0, "data": { @@ -416,7 +418,7 @@ class TaskResultView(APIView): } }) - if task_result.failed(): + if await sync_to_async(task_result.failed)(): return Response({ "code": 1, "message": str(task_result.result) if task_result.result else "任务执行失败" @@ -451,7 +453,7 @@ class TranslateUsageView(APIView): operation_description='返回当前用户/IP的今日使用量、总量以及各类型使用统计', responses={200: success_response} ) - def get(self, request): + async def get(self, request): from ..models import TranslateUsage if request.user.is_authenticated: @@ -461,12 +463,12 @@ class TranslateUsageView(APIView): queryset = TranslateUsage.objects.filter(ip_address=ip_address) today_start = timezone.localdate() - today_count = queryset.filter(created_at__date=today_start).count() - total_count = queryset.count() + today_count = await queryset.filter(created_at__date=today_start).acount() + total_count = await queryset.acount() by_type = {} for type_key, _ in TranslateUsage.TRANSLATE_TYPE_CHOICES: - by_type[type_key] = queryset.filter(translate_type=type_key).count() + by_type[type_key] = await queryset.filter(translate_type=type_key).acount() return Response({ "code": 0, diff --git a/api/views/GetIPDataView.py b/api/views/GetIPDataView.py index a61b2ea..2f8a3de 100644 --- a/api/views/GetIPDataView.py +++ b/api/views/GetIPDataView.py @@ -1,7 +1,10 @@ -import requests +import asyncio + +import aiohttp +from asgiref.sync import sync_to_async from rest_framework.permissions import AllowAny -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from drf_yasg.utils import swagger_auto_schema @@ -62,7 +65,7 @@ class GetIPDataView(APIView): 502: openapi.Response(description="第三方服务异常"), } ) - def get(self, request): + async def get(self, request): # 获取要查询的IP地址 ip = request.GET.get('ip', '').strip() if not ip: @@ -74,7 +77,7 @@ class GetIPDataView(APIView): status=status.HTTP_400_BAD_REQUEST ) - return self._fetch_ip_location(ip) + return await self._fetch_ip_location(ip) def _get_client_ip(self, request): """ @@ -96,14 +99,14 @@ class GetIPDataView(APIView): return None - def _fetch_ip_location(self, ip): + async def _fetch_ip_location(self, ip): """ - 调用 ip-api.com 获取IP地理信息,支持Redis缓存 + 调用 ip-api.com 获取IP地理信息,支持Redis缓存(全异步:aiohttp + sync_to_async 缓存兜底) """ cache_key = f"ip_location:{ip}" - # 先尝试从缓存获取 - cached_data = default_cache.get(cache_key) + # 先尝试从缓存获取(django-redis 为同步客户端,用 sync_to_async 兜底) + cached_data = await sync_to_async(default_cache.get)(cache_key) if cached_data: return Response( {"code": 200, "message": "success", "data": cached_data}, @@ -112,15 +115,15 @@ class GetIPDataView(APIView): try: url = IP_API_URL.format(ip=ip) - response = requests.get(url, timeout=REQUEST_TIMEOUT) - - if response.status_code != 200: - return Response( - {"code": 502, "message": "IP定位服务异常,请稍后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY - ) - - result = response.json() + timeout = aiohttp.ClientTimeout(total=REQUEST_TIMEOUT) + async with aiohttp.ClientSession(timeout=timeout) as session: + async with session.get(url) as resp: + if resp.status != 200: + return Response( + {"code": 502, "message": "IP定位服务异常,请稍后重试", "data": None}, + status=status.HTTP_502_BAD_GATEWAY + ) + result = await resp.json() if result.get('status') != 'success': return Response( @@ -139,24 +142,19 @@ class GetIPDataView(APIView): 'timezone': result.get('timezone', ''), } - # 写入缓存 - default_cache.set(cache_key, ip_info, IP_CACHE_TIMEOUT) + # 写入缓存(sync_to_async 兜底) + await sync_to_async(default_cache.set)(cache_key, ip_info, IP_CACHE_TIMEOUT) return Response( {"code": 200, "message": "success", "data": ip_info}, status=status.HTTP_200_OK ) - except requests.exceptions.Timeout: + except (aiohttp.ClientError, asyncio.TimeoutError): return Response( {"code": 502, "message": "请求超时,请稍后重试", "data": None}, status=status.HTTP_502_BAD_GATEWAY ) - except requests.exceptions.ConnectionError: - return Response( - {"code": 502, "message": "网络连接异常,请检查网络后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY - ) except Exception as e: return Response( {"code": 500, "message": f"服务器内部错误:{str(e)}", "data": None}, diff --git a/apidirectory/serializers.py b/apidirectory/serializers.py index c19b070..7443ba0 100644 --- a/apidirectory/serializers.py +++ b/apidirectory/serializers.py @@ -1,23 +1,23 @@ -from rest_framework import serializers +from adrf.serializers import ModelSerializer, CharField, SerializerMethodField from .models import ApiCategory, ApiItem, ApiFavorite -class ApiCategorySerializer(serializers.ModelSerializer): - item_count = serializers.SerializerMethodField() +class ApiCategorySerializer(ModelSerializer): + item_count = SerializerMethodField() class Meta: model = ApiCategory fields = ['id', 'name', 'icon', 'sort_order', 'item_count', 'created_at'] read_only_fields = ['id', 'created_at'] - def get_item_count(self, obj): - return obj.items.filter(is_enabled=True).count() + async def get_item_count(self, obj): + return await obj.items.filter(is_enabled=True).acount() -class ApiItemSerializer(serializers.ModelSerializer): - category_name = serializers.CharField(source='category.name', read_only=True, default='') - favorites_count = serializers.SerializerMethodField() - is_favorited = serializers.SerializerMethodField() +class ApiItemSerializer(ModelSerializer): + category_name = CharField(source='category.name', read_only=True, default='') + favorites_count = SerializerMethodField() + is_favorited = SerializerMethodField() class Meta: model = ApiItem @@ -29,19 +29,19 @@ class ApiItemSerializer(serializers.ModelSerializer): ] read_only_fields = ['id', 'created_at', 'updated_at'] - def get_favorites_count(self, obj): - return obj.favorites.count() + async def get_favorites_count(self, obj): + return await obj.favorites.acount() - def get_is_favorited(self, obj): + async def get_is_favorited(self, obj): request = self.context.get('request') if request and request.user.is_authenticated: - return obj.favorites.filter(user=request.user).exists() + return await obj.favorites.filter(user=request.user).aexists() return False -class ApiFavoriteSerializer(serializers.ModelSerializer): - api_item_name = serializers.CharField(source='api_item.name', read_only=True) - api_item_url = serializers.CharField(source='api_item.url_path', read_only=True) +class ApiFavoriteSerializer(ModelSerializer): + api_item_name = CharField(source='api_item.name', read_only=True) + api_item_url = CharField(source='api_item.url_path', read_only=True) class Meta: model = ApiFavorite diff --git a/apidirectory/views.py b/apidirectory/views.py index 6a6f1b2..a59a56e 100644 --- a/apidirectory/views.py +++ b/apidirectory/views.py @@ -1,4 +1,4 @@ -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from rest_framework.permissions import AllowAny, IsAuthenticated @@ -20,8 +20,8 @@ class CategoryListView(APIView): tags=['API接口大全'], responses={200: '成功'} ) - def get(self, request): - categories = ApiCategory.objects.prefetch_related('items').all() + async def get(self, request): + categories = [c async for c in ApiCategory.objects.prefetch_related('items').all()] serializer = ApiCategorySerializer(categories, many=True) return JsonResponse({ 'code': 200, @@ -44,7 +44,7 @@ class ApiItemListView(APIView): ], responses={200: '成功'} ) - def get(self, request): + async def get(self, request): category_id = request.query_params.get('category_id') search = request.query_params.get('search', '') ordering = request.query_params.get('ordering', '') @@ -69,7 +69,8 @@ class ApiItemListView(APIView): else: items = items.order_by('sort_order', '-created_at') - serializer = ApiItemSerializer(items, many=True, context={'request': request}) + item_list = [x async for x in items] + serializer = ApiItemSerializer(item_list, many=True, context={'request': request}) return JsonResponse({ 'code': 200, 'message': 'success', @@ -86,9 +87,9 @@ class ApiItemDetailView(APIView): tags=['API接口大全'], responses={200: '成功'} ) - def get(self, request, pk): + async def get(self, request, pk): try: - item = ApiItem.objects.select_related('category').get(pk=pk, is_enabled=True) + item = await ApiItem.objects.select_related('category').aget(pk=pk, is_enabled=True) serializer = ApiItemSerializer(item, context={'request': request}) return JsonResponse({ 'code': 200, @@ -112,9 +113,9 @@ class FavoriteToggleView(APIView): tags=['API接口大全'], responses={200: '成功'} ) - def post(self, request, pk): + async def post(self, request, pk): try: - api_item = ApiItem.objects.get(pk=pk, is_enabled=True) + api_item = await ApiItem.objects.aget(pk=pk, is_enabled=True) except ApiItem.DoesNotExist: return JsonResponse({ 'code': 404, @@ -122,14 +123,14 @@ class FavoriteToggleView(APIView): 'data': None }, status=status.HTTP_404_NOT_FOUND) - fav, created = ApiFavorite.objects.get_or_create(user=request.user, api_item=api_item) + fav, created = await ApiFavorite.objects.aget_or_create(user=request.user, api_item=api_item) if not created: - fav.delete() + await fav.adelete() is_favorited = False else: is_favorited = True - favorites_count = ApiFavorite.objects.filter(api_item=api_item).count() + favorites_count = await ApiFavorite.objects.filter(api_item=api_item).acount() return JsonResponse({ 'code': 200, 'message': 'success', diff --git a/app/views.py b/app/views.py index 43b0536..030f34e 100644 --- a/app/views.py +++ b/app/views.py @@ -1,9 +1,10 @@ -from rest_framework import viewsets, permissions +from rest_framework import permissions from rest_framework.pagination import PageNumberPagination from rest_framework.response import Response from drf_yasg.utils import swagger_auto_schema from drf_yasg import openapi from chunyu_project.common_schemas import success_response +from adrf.viewsets import ModelViewSet from .models import Changelog from .serializers import ChangelogSerializer @@ -27,7 +28,7 @@ class ChangelogPagination(PageNumberPagination): }) -class ChangelogViewSet(viewsets.ModelViewSet): +class ChangelogViewSet(ModelViewSet): """更新日志 ViewSet""" queryset = Changelog.objects.filter(is_published=True) serializer_class = ChangelogSerializer @@ -53,8 +54,8 @@ class ChangelogViewSet(viewsets.ModelViewSet): ], responses={200: success_response} ) - def list(self, request, *args, **kwargs): - return super().list(request, *args, **kwargs) + async def list(self, request, *args, **kwargs): + return await super().list(request, *args, **kwargs) @swagger_auto_schema( tags=['更新日志'], @@ -65,5 +66,5 @@ class ChangelogViewSet(viewsets.ModelViewSet): ], responses={200: success_response} ) - def retrieve(self, request, *args, **kwargs): - return super().retrieve(request, *args, **kwargs) \ No newline at end of file + async def retrieve(self, request, *args, **kwargs): + return await super().retrieve(request, *args, **kwargs) \ No newline at end of file diff --git a/article/serializers.py b/article/serializers.py index 256f7cf..1f913a5 100644 --- a/article/serializers.py +++ b/article/serializers.py @@ -1,49 +1,55 @@ +from adrf.serializers import ( + ModelSerializer, CharField, IntegerField, ImageField, DateTimeField, + JSONField, SerializerMethodField, +) from rest_framework import serializers from .models import Article, ArticleComment, ArticleLike, ArticleFavorite, ArticleCommentLike -class ArticleCommentSerializer(serializers.ModelSerializer): - user_id = serializers.IntegerField(source='user.id', read_only=True) - user_name = serializers.CharField(source='user.first_name', read_only=True) - user_avatar = serializers.ImageField(source='user.avatar', read_only=True) - user_location = serializers.CharField(source='user.location', read_only=True) - parent_user_name = serializers.SerializerMethodField() - replies = serializers.SerializerMethodField() - is_liked = serializers.SerializerMethodField() - created_at = serializers.DateTimeField(format='%Y-%m-%d %H:%M') +class ArticleCommentSerializer(ModelSerializer): + user_id = IntegerField(source='user.id', read_only=True) + user_name = CharField(source='user.first_name', read_only=True) + user_avatar = ImageField(source='user.avatar', read_only=True) + user_location = CharField(source='user.location', read_only=True) + parent_user_name = SerializerMethodField() + replies = SerializerMethodField() + is_liked = SerializerMethodField() + created_at = DateTimeField(format='%Y-%m-%d %H:%M') class Meta: model = ArticleComment fields = ['id', 'user_id', 'user_name', 'user_avatar', 'user_location', 'content', 'parent', 'parent_user_name', 'likes', 'created_at', 'replies', 'is_liked'] read_only_fields = ['likes', 'created_at'] - def get_parent_user_name(self, obj): - if obj.parent: - return obj.parent.user.first_name or obj.parent.user.username + async def get_parent_user_name(self, obj): + if obj.parent_id: + # 异步加载父评论作者,避免外键懒加载 + parent = await ArticleComment.objects.select_related('user').aget(pk=obj.parent_id) + return parent.user.first_name or parent.user.username return '' - def get_replies(self, obj): - if obj.parent is None: - replies = obj.replies.all() - return ArticleCommentSerializer(replies, many=True, context=self.context).data + async def get_replies(self, obj): + if obj.parent_id is None: + replies = [r async for r in obj.replies.all()] + return [await ArticleCommentSerializer(r, context=self.context).adata for r in replies] return [] - def get_is_liked(self, obj): + async def get_is_liked(self, obj): request = self.context.get('request') if request and request.user.is_authenticated: - return ArticleCommentLike.objects.filter(user=request.user, comment=obj).exists() + return await ArticleCommentLike.objects.filter(user=request.user, comment=obj).aexists() return False -class ArticleListSerializer(serializers.ModelSerializer): - author_name = serializers.CharField(source='author.first_name', read_only=True) - author_avatar = serializers.ImageField(source='author.avatar', read_only=True) - comments_count = serializers.SerializerMethodField() - read_time = serializers.SerializerMethodField() - is_hot = serializers.SerializerMethodField() - is_new = serializers.SerializerMethodField() - publish_date = serializers.DateTimeField(source='created_at', format='%Y-%m-%d', read_only=True) - image = serializers.SerializerMethodField() +class ArticleListSerializer(ModelSerializer): + author_name = CharField(source='author.first_name', read_only=True) + author_avatar = ImageField(source='author.avatar', read_only=True) + comments_count = SerializerMethodField() + read_time = SerializerMethodField() + is_hot = SerializerMethodField() + is_new = SerializerMethodField() + publish_date = DateTimeField(source='created_at', format='%Y-%m-%d', read_only=True) + image = SerializerMethodField() class Meta: model = Article @@ -53,25 +59,25 @@ class ArticleListSerializer(serializers.ModelSerializer): 'read_time', 'is_featured', 'is_hot', 'is_new', 'is_top', 'status', ] - def get_comments_count(self, obj): - return obj.comments.count() + async def get_comments_count(self, obj): + return await obj.comments.acount() - def get_read_time(self, obj): + async def get_read_time(self, obj): if obj.content: word_count = len(obj.content) minutes = max(1, word_count // 300) return f'{minutes} 分钟' return '1 分钟' - def get_is_hot(self, obj): + async def get_is_hot(self, obj): return obj.views > 1000 or obj.likes > 100 - def get_is_new(self, obj): + async def get_is_new(self, obj): from django.utils import timezone import datetime return (timezone.now() - obj.created_at).days < 7 - def get_image(self, obj): + async def get_image(self, obj): if obj.cover_image: request = self.context.get('request') if request: @@ -80,23 +86,23 @@ class ArticleListSerializer(serializers.ModelSerializer): return '' -class ArticleDetailSerializer(serializers.ModelSerializer): - author_name = serializers.CharField(source='author.first_name', read_only=True) - author_avatar = serializers.ImageField(source='author.avatar', read_only=True) - author_bio = serializers.CharField(source='author.bio', read_only=True, default='') +class ArticleDetailSerializer(ModelSerializer): + author_name = CharField(source='author.first_name', read_only=True) + author_avatar = ImageField(source='author.avatar', read_only=True) + author_bio = CharField(source='author.bio', read_only=True, default='') author_user_id = serializers.PrimaryKeyRelatedField(source='author', read_only=True) - comments_count = serializers.SerializerMethodField() - favorites_count = serializers.SerializerMethodField() - read_time = serializers.SerializerMethodField() - publish_date = serializers.DateTimeField(source='created_at', format='%Y-%m-%d %H:%M', read_only=True) - update_date = serializers.DateTimeField(source='updated_at', format='%Y-%m-%d %H:%M', read_only=True) - image = serializers.SerializerMethodField() - is_liked = serializers.SerializerMethodField() - is_favorited = serializers.SerializerMethodField() - is_following = serializers.SerializerMethodField() - related_articles = serializers.SerializerMethodField() - comments = serializers.SerializerMethodField() - comments_pagination = serializers.SerializerMethodField() + comments_count = SerializerMethodField() + favorites_count = SerializerMethodField() + read_time = SerializerMethodField() + publish_date = DateTimeField(source='created_at', format='%Y-%m-%d %H:%M', read_only=True) + update_date = DateTimeField(source='updated_at', format='%Y-%m-%d %H:%M', read_only=True) + image = SerializerMethodField() + is_liked = SerializerMethodField() + is_favorited = SerializerMethodField() + is_following = SerializerMethodField() + related_articles = SerializerMethodField() + comments = SerializerMethodField() + comments_pagination = SerializerMethodField() class Meta: model = Article @@ -108,20 +114,20 @@ class ArticleDetailSerializer(serializers.ModelSerializer): 'is_liked', 'is_favorited', 'is_following', 'related_articles', 'comments', 'comments_pagination', ] - def get_comments_count(self, obj): - return obj.comments.count() + async def get_comments_count(self, obj): + return await obj.comments.acount() - def get_favorites_count(self, obj): - return obj.article_favorites.count() + async def get_favorites_count(self, obj): + return await obj.article_favorites.acount() - def get_read_time(self, obj): + async def get_read_time(self, obj): if obj.content: word_count = len(obj.content) minutes = max(1, word_count // 300) return f'{minutes} 分钟' return '1 分钟' - def get_image(self, obj): + async def get_image(self, obj): if obj.cover_image: request = self.context.get('request') if request: @@ -129,77 +135,77 @@ class ArticleDetailSerializer(serializers.ModelSerializer): return obj.cover_image.url return '' - def get_is_liked(self, obj): + async def get_is_liked(self, obj): request = self.context.get('request') if request and request.user.is_authenticated: - return ArticleLike.objects.filter(user=request.user, article=obj).exists() + return await ArticleLike.objects.filter(user=request.user, article=obj).aexists() return False - def get_is_favorited(self, obj): + async def get_is_favorited(self, obj): request = self.context.get('request') if request and request.user.is_authenticated: - return ArticleFavorite.objects.filter(user=request.user, article=obj).exists() + return await ArticleFavorite.objects.filter(user=request.user, article=obj).aexists() return False - def get_is_following(self, obj): + async def get_is_following(self, obj): request = self.context.get('request') if request and request.user.is_authenticated: from user.models import Follow - return Follow.objects.filter(follower=request.user, following=obj.author).exists() + return await Follow.objects.filter(follower=request.user, following=obj.author_id).aexists() return False - def get_related_articles(self, obj): + async def get_related_articles(self, obj): queryset = Article.objects.filter(status='published').exclude(id=obj.id) related = queryset.filter(category=obj.category)[:4] - if related.count() < 4: - remaining = 4 - related.count() + related_list = [a async for a in related] + if len(related_list) < 4: + remaining = 4 - len(related_list) extra = queryset.filter(tags__overlap=obj.tags or []).exclude( - id__in=list(related.values_list('id', flat=True)) + id__in=[a.id for a in related_list] )[:remaining] - related = list(related) + list(extra) - return ArticleListSerializer(related, many=True, context=self.context).data + related_list = related_list + [a async for a in extra] + return [await ArticleListSerializer(a, context=self.context).adata for a in related_list] - def get_comments(self, obj): + async def get_comments(self, obj): request = self.context.get('request') page_size = int(request.query_params.get('comment_page_size', 10)) if request else 10 page = int(request.query_params.get('comment_page', 1)) if request else 10 top_comments = obj.comments.filter(parent__isnull=True) start = (page - 1) * page_size end = start + page_size - comments_page = top_comments[start:end] - serializer = ArticleCommentSerializer(comments_page, many=True, context=self.context) - return serializer.data + comments_page = [c async for c in top_comments[start:end]] + return [await ArticleCommentSerializer(c, context=self.context).adata for c in comments_page] - def get_comments_pagination(self, obj): + async def get_comments_pagination(self, obj): request = self.context.get('request') page_size = int(request.query_params.get('comment_page_size', 10)) if request else 10 - top_count = obj.comments.filter(parent__isnull=True).count() + top_count = await obj.comments.filter(parent__isnull=True).acount() return { 'count': top_count, 'next': top_count > page_size, } -class ArticleCreateUpdateSerializer(serializers.ModelSerializer): - tags = serializers.JSONField(required=False, default=list) +class ArticleCreateUpdateSerializer(ModelSerializer): + tags = JSONField(required=False, default=list) class Meta: model = Article fields = ['title', 'excerpt', 'content', 'category', 'tags', 'cover_image', 'status', 'editor_mode'] read_only_fields = ['status'] - def create(self, validated_data): + async def acreate(self, validated_data): validated_data['author'] = self.context['request'].user - return super().create(validated_data) + return await super().acreate(validated_data) -class ArticleManageSerializer(serializers.ModelSerializer): - author_name = serializers.CharField(source='author.first_name', read_only=True) - comments_count = serializers.SerializerMethodField() - favorites_count = serializers.SerializerMethodField() - summary = serializers.CharField(source='excerpt', read_only=True) - publish_date = serializers.DateTimeField(source='created_at', format='%Y-%m-%d', read_only=True) - update_date = serializers.DateTimeField(source='updated_at', format='%Y-%m-%d', read_only=True) +class ArticleManageSerializer(ModelSerializer): + author_name = CharField(source='author.first_name', read_only=True) + comments_count = SerializerMethodField() + favorites_count = SerializerMethodField() + summary = CharField(source='excerpt', read_only=True) + publish_date = DateTimeField(source='created_at', format='%Y-%m-%d', read_only=True) + update_date = DateTimeField(source='updated_at', format='%Y-%m-%d', read_only=True) class Meta: model = Article @@ -209,8 +215,8 @@ class ArticleManageSerializer(serializers.ModelSerializer): 'category', 'tags', 'is_top', 'author_name', ] - def get_comments_count(self, obj): - return obj.comments.count() + async def get_comments_count(self, obj): + return await obj.comments.acount() - def get_favorites_count(self, obj): - return obj.article_favorites.count() + async def get_favorites_count(self, obj): + return await obj.article_favorites.acount() diff --git a/article/views.py b/article/views.py index 8646e53..b40a2ae 100644 --- a/article/views.py +++ b/article/views.py @@ -1,10 +1,13 @@ -from rest_framework import generics, status +from adrf import generics +from adrf.generics import aget_object_or_404 +from adrf.mixins import get_data +from adrf.views import APIView +from asgiref.sync import sync_to_async +from rest_framework import status from rest_framework.permissions import IsAuthenticated, AllowAny, IsAuthenticatedOrReadOnly from rest_framework.parsers import JSONParser, MultiPartParser, FormParser from rest_framework.pagination import PageNumberPagination -from rest_framework.views import APIView from rest_framework.response import Response -from django.shortcuts import get_object_or_404 from django.db.models import F from django.db import transaction from drf_yasg.utils import swagger_auto_schema @@ -77,6 +80,12 @@ class ArticleListCreateView(generics.ListCreateAPIView): ordering = self.request.query_params.get('ordering', '-created_at') return qs.order_by('-is_top', ordering) + async def get(self, request, *args, **kwargs): + return await self.list(request, *args, **kwargs) + + async def post(self, request, *args, **kwargs): + return await self.create(request, *args, **kwargs) + @swagger_auto_schema( tags=['文章'], operation_summary='获取文章列表', @@ -91,17 +100,17 @@ class ArticleListCreateView(generics.ListCreateAPIView): ], responses={200: success_response} ) - def list(self, request, *args, **kwargs): - queryset = self.filter_queryset(self.get_queryset()) - page = self.paginate_queryset(queryset) + async def list(self, request, *args, **kwargs): + queryset = await self.afilter_queryset(self.get_queryset()) + page = await self.apaginate_queryset(queryset) if page is not None: serializer = self.get_serializer(page, many=True) - data = serializer.data + data = await get_data(serializer) if request.user.is_authenticated: from .models import ArticleFavorite - fav_ids = set( - ArticleFavorite.objects.filter(user=request.user).values_list('article_id', flat=True) - ) + fav_ids = set([ + v async for v in ArticleFavorite.objects.filter(user=request.user).values_list('article_id', flat=True) + ]) for item in data: item['is_favorited'] = item['id'] in fav_ids else: @@ -109,12 +118,12 @@ class ArticleListCreateView(generics.ListCreateAPIView): item['is_favorited'] = False return self.get_paginated_response(data) serializer = self.get_serializer(queryset, many=True) - data = serializer.data + data = await get_data(serializer) if request.user.is_authenticated: from .models import ArticleFavorite - fav_ids = set( - ArticleFavorite.objects.filter(user=request.user).values_list('article_id', flat=True) - ) + fav_ids = set([ + v async for v in ArticleFavorite.objects.filter(user=request.user).values_list('article_id', flat=True) + ]) for item in data: item['is_favorited'] = item['id'] in fav_ids else: @@ -133,13 +142,14 @@ class ArticleListCreateView(generics.ListCreateAPIView): 401: unauthorized_response, } ) - def create(self, request, *args, **kwargs): + async def create(self, request, *args, **kwargs): serializer = self.get_serializer(data=request.data) - if serializer.is_valid(): - serializer.save() - track_task(request.user, 'post') + if await sync_to_async(serializer.is_valid)(): + await sync_to_async(serializer.save)() + await sync_to_async(track_task)(request.user, 'post') + data = await get_data(serializer) return create_standardized_response( - data=serializer.data, + data=data, code=ResponseCode.SUCCESS, message='文章创建成功', status_code=status.HTTP_201_CREATED @@ -168,19 +178,19 @@ class ArticleDetailView(APIView): 404: not_found_response, } ) - def get(self, request, pk): - article = get_object_or_404(Article, pk=pk) - Article.objects.filter(pk=pk).update(views=F('views') + 1) - article.refresh_from_db() + async def get(self, request, pk): + article = await aget_object_or_404(Article, pk=pk) + await Article.objects.filter(pk=pk).aupdate(views=F('views') + 1) + await article.arefresh_from_db() serializer = ArticleDetailSerializer(article, context={'request': request}) - data = serializer.data + data = await get_data(serializer) if request.user.is_authenticated: - data['is_favorited'] = ArticleFavorite.objects.filter( + data['is_favorited'] = await ArticleFavorite.objects.filter( user=request.user, article=article - ).exists() + ).aexists() else: data['is_favorited'] = False - data['favorites_count'] = ArticleFavorite.objects.filter(article=article).count() + data['favorites_count'] = await ArticleFavorite.objects.filter(article=article).acount() return create_standardized_response(data=data, code=ResponseCode.SUCCESS) @swagger_auto_schema( @@ -199,8 +209,8 @@ class ArticleDetailView(APIView): 404: not_found_response, } ) - def put(self, request, pk): - article = get_object_or_404(Article, pk=pk) + async def put(self, request, pk): + article = await aget_object_or_404(Article, pk=pk) if article.author != request.user: return create_standardized_error_response( code=ResponseCode.VALIDATION_ERROR, @@ -208,9 +218,10 @@ class ArticleDetailView(APIView): status_code=status.HTTP_403_FORBIDDEN ) serializer = ArticleCreateUpdateSerializer(article, data=request.data, partial=True, context={'request': request}) - if serializer.is_valid(): - serializer.save() - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS, message='文章更新成功') + if await sync_to_async(serializer.is_valid)(): + await sync_to_async(serializer.save)() + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS, message='文章更新成功') return create_standardized_error_response( data=serializer.errors, code=ResponseCode.VALIDATION_ERROR, @@ -231,15 +242,15 @@ class ArticleDetailView(APIView): 404: not_found_response, } ) - def delete(self, request, pk): - article = get_object_or_404(Article, pk=pk) + async def delete(self, request, pk): + article = await aget_object_or_404(Article, pk=pk) if article.author != request.user: return create_standardized_error_response( code=ResponseCode.VALIDATION_ERROR, message='无权操作', status_code=status.HTTP_403_FORBIDDEN ) - article.delete() + await article.adelete() return create_standardized_response(code=ResponseCode.SUCCESS, message='文章删除成功', status_code=status.HTTP_204_NO_CONTENT) @@ -258,6 +269,12 @@ class ArticleCommentListCreateView(generics.ListCreateAPIView): article_id = self.kwargs['article_id'] return ArticleComment.objects.filter(article_id=article_id, parent__isnull=True) + async def get(self, request, *args, **kwargs): + return await self.list(request, *args, **kwargs) + + async def post(self, request, *args, **kwargs): + return await self.create(request, *args, **kwargs) + @swagger_auto_schema( tags=['文章'], operation_summary='获取文章评论列表', @@ -272,14 +289,16 @@ class ArticleCommentListCreateView(generics.ListCreateAPIView): 404: not_found_response, } ) - def list(self, request, *args, **kwargs): + async def list(self, request, *args, **kwargs): queryset = self.get_queryset() - page = self.paginate_queryset(queryset) + page = await self.apaginate_queryset(queryset) if page is not None: serializer = self.get_serializer(page, many=True) - return self.get_paginated_response(serializer.data) + data = await get_data(serializer) + return self.get_paginated_response(data) serializer = self.get_serializer(queryset, many=True) - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS) + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS) @swagger_auto_schema( tags=['文章'], @@ -296,32 +315,36 @@ class ArticleCommentListCreateView(generics.ListCreateAPIView): 404: not_found_response, } ) - def create(self, request, *args, **kwargs): + async def create(self, request, *args, **kwargs): article_id = self.kwargs['article_id'] - article = get_object_or_404(Article, pk=article_id) + article = await aget_object_or_404(Article, pk=article_id) serializer = self.get_serializer(data=request.data) - if serializer.is_valid(): - with transaction.atomic(): - comment = serializer.save(user=request.user, article=article) - comment_id = comment.id - comment_content = serializer.data.get('content', '') - article_title = article.title - article_author_id = article.author_id - parent_comment_id = comment.parent_id - parent_comment_user_id = comment.parent.user_id if comment.parent else None - user_id = request.user.id - user_nickname = request.user.nickname or request.user.username + if await sync_to_async(serializer.is_valid)(): + def _save_and_notify(): + with transaction.atomic(): + comment = serializer.save(user=request.user, article=article) + comment_id = comment.id + comment_content = serializer.data.get('content', '') + article_title = article.title + article_author_id = article.author_id + parent_comment_id = comment.parent_id + parent_comment_user_id = comment.parent.user_id if comment.parent else None + user_id = request.user.id + user_nickname = request.user.nickname or request.user.username - transaction.on_commit(lambda: self._send_notification_messages( - article_id, article_title, article_author_id, comment_id, - comment_content, parent_comment_id, parent_comment_user_id, - user_id, user_nickname - )) + transaction.on_commit(lambda: self._send_notification_messages( + article_id, article_title, article_author_id, comment_id, + comment_content, parent_comment_id, parent_comment_user_id, + user_id, user_nickname + )) + return comment + await sync_to_async(_save_and_notify)() - track_task(request.user, 'post') + await sync_to_async(track_task)(request.user, 'post') + data = await get_data(serializer) return create_standardized_response( - data=serializer.data, + data=data, code=ResponseCode.SUCCESS, message='评论成功', status_code=status.HTTP_201_CREATED @@ -412,23 +435,23 @@ class ArticleLikeToggleView(APIView): 404: not_found_response, } ) - def post(self, request, pk): - article = get_object_or_404(Article, pk=pk) - like, created = ArticleLike.objects.get_or_create(user=request.user, article=article) + async def post(self, request, pk): + article = await aget_object_or_404(Article, pk=pk) + like, created = await ArticleLike.objects.aget_or_create(user=request.user, article=article) if not created: - like.delete() - Article.objects.filter(pk=pk).update(likes=F("likes") - 1) - article.refresh_from_db() + await like.adelete() + await Article.objects.filter(pk=pk).aupdate(likes=F("likes") - 1) + await article.arefresh_from_db() article.likes = max(0, article.likes) return create_standardized_response( data={'liked': False, 'likes_count': article.likes}, code=ResponseCode.SUCCESS ) article.likes = F('likes') + 1 - article.save(update_fields=['likes']) - article.refresh_from_db() + await article.asave(update_fields=['likes']) + await article.arefresh_from_db() if created and article.author != request.user: - create_message( + await sync_to_async(create_message)( recipient=article.author, sender=request.user, msg_type='like', @@ -461,17 +484,17 @@ class ArticleFavoriteToggleView(APIView): 404: not_found_response, } ) - def post(self, request, pk): - article = get_object_or_404(Article, pk=pk) - fav, created = ArticleFavorite.objects.get_or_create(user=request.user, article=article) + async def post(self, request, pk): + article = await aget_object_or_404(Article, pk=pk) + fav, created = await ArticleFavorite.objects.aget_or_create(user=request.user, article=article) if not created: - fav.delete() - count = article.article_favorites.count() + await fav.adelete() + count = await article.article_favorites.acount() return create_standardized_response( data={'favorited': False, 'favorites_count': count}, code=ResponseCode.SUCCESS ) - count = article.article_favorites.count() + count = await article.article_favorites.acount() return create_standardized_response( data={'favorited': True, 'favorites_count': count}, code=ResponseCode.SUCCESS @@ -494,20 +517,20 @@ class ArticleCommentLikeToggleView(APIView): 404: not_found_response, } ) - def post(self, request, pk): - comment = get_object_or_404(ArticleComment, pk=pk) - like, created = ArticleCommentLike.objects.get_or_create(user=request.user, comment=comment) + async def post(self, request, pk): + comment = await aget_object_or_404(ArticleComment, pk=pk) + like, created = await ArticleCommentLike.objects.aget_or_create(user=request.user, comment=comment) if not created: - like.delete() - ArticleComment.objects.filter(pk=pk).update(likes=F("likes") - 1) - comment.refresh_from_db() + await like.adelete() + await ArticleComment.objects.filter(pk=pk).aupdate(likes=F("likes") - 1) + await comment.arefresh_from_db() comment.likes = max(0, comment.likes) return create_standardized_response( data={'liked': False, 'likes_count': comment.likes}, code=ResponseCode.SUCCESS ) - comment.likes = comment.comment_likes.count() - comment.save(update_fields=['likes']) + comment.likes = await comment.comment_likes.acount() + await comment.asave(update_fields=['likes']) return create_standardized_response( data={'liked': True, 'likes_count': comment.likes}, code=ResponseCode.SUCCESS @@ -526,6 +549,9 @@ class MyArticleListView(generics.ListAPIView): qs = qs.filter(status=st) return qs.order_by('-updated_at') + async def get(self, request, *args, **kwargs): + return await self.list(request, *args, **kwargs) + @swagger_auto_schema( tags=['文章'], operation_summary='获取我的文章列表', @@ -538,10 +564,11 @@ class MyArticleListView(generics.ListAPIView): 401: unauthorized_response, } ) - def list(self, request, *args, **kwargs): + async def list(self, request, *args, **kwargs): queryset = self.get_queryset() serializer = self.get_serializer(queryset, many=True) - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS) + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS) class MyArticleBatchView(APIView): @@ -574,7 +601,7 @@ class MyArticleBatchView(APIView): 401: unauthorized_response, } ) - def post(self, request): + async def post(self, request): ids = request.data.get('ids', []) action = request.data.get('action', '') if not ids or action not in ('publish', 'draft', 'delete'): @@ -585,11 +612,11 @@ class MyArticleBatchView(APIView): ) qs = Article.objects.filter(id__in=ids, author=request.user) if action == 'delete': - count = qs.delete()[0] + count = (await qs.adelete())[0] elif action == 'publish': - count = qs.update(status='published') + count = await qs.aupdate(status='published') elif action == 'draft': - count = qs.update(status='draft') + count = await qs.aupdate(status='draft') return create_standardized_response( data={'affected': count}, code=ResponseCode.SUCCESS, @@ -614,8 +641,8 @@ class ArticleToggleTopView(APIView): 404: not_found_response, } ) - def post(self, request, pk): - article = get_object_or_404(Article, pk=pk) + async def post(self, request, pk): + article = await aget_object_or_404(Article, pk=pk) if article.author != request.user: return create_standardized_error_response( code=ResponseCode.VALIDATION_ERROR, @@ -623,9 +650,9 @@ class ArticleToggleTopView(APIView): status_code=status.HTTP_403_FORBIDDEN ) article.is_top = not article.is_top - article.save(update_fields=['is_top']) + await article.asave(update_fields=['is_top']) return create_standardized_response( data={'is_top': article.is_top}, code=ResponseCode.SUCCESS, message='置顶状态已更新' - ) \ No newline at end of file + ) diff --git a/bug/serializers.py b/bug/serializers.py index 44a112a..896de13 100644 --- a/bug/serializers.py +++ b/bug/serializers.py @@ -1,40 +1,44 @@ -from rest_framework import serializers +from adrf.serializers import ( + ModelSerializer, CharField, ImageField, DateTimeField, + ListField, FileField, SerializerMethodField, +) +from rest_framework.exceptions import ValidationError from .models import ( BugReport, BugReportImage, BugReportAttachment, BugReportComment, BugReportCommentImage, BugReportCommentAttachment ) -class BugReportImageSerializer(serializers.ModelSerializer): +class BugReportImageSerializer(ModelSerializer): class Meta: model = BugReportImage fields = ['id', 'image', 'created_at'] -class BugReportAttachmentSerializer(serializers.ModelSerializer): +class BugReportAttachmentSerializer(ModelSerializer): class Meta: model = BugReportAttachment fields = ['id', 'file', 'filename', 'file_size', 'created_at'] -class BugReportCommentImageSerializer(serializers.ModelSerializer): +class BugReportCommentImageSerializer(ModelSerializer): class Meta: model = BugReportCommentImage fields = ['id', 'image', 'created_at'] -class BugReportCommentAttachmentSerializer(serializers.ModelSerializer): +class BugReportCommentAttachmentSerializer(ModelSerializer): class Meta: model = BugReportCommentAttachment fields = ['id', 'file', 'filename', 'file_size', 'created_at'] -class BugReportCommentSerializer(serializers.ModelSerializer): - user_name = serializers.CharField(source='user.nickname', read_only=True) - user_avatar = serializers.ImageField(source='user.avatar', read_only=True) +class BugReportCommentSerializer(ModelSerializer): + user_name = CharField(source='user.nickname', read_only=True) + user_avatar = ImageField(source='user.avatar', read_only=True) images = BugReportCommentImageSerializer(many=True, read_only=True) attachments = BugReportCommentAttachmentSerializer(many=True, read_only=True) - created_at = serializers.DateTimeField(format='%Y-%m-%d %H:%M') + created_at = DateTimeField(format='%Y-%m-%d %H:%M') class Meta: model = BugReportComment @@ -45,9 +49,9 @@ class BugReportCommentSerializer(serializers.ModelSerializer): read_only_fields = ['user', 'is_admin'] -class BugReportListSerializer(serializers.ModelSerializer): - user_name = serializers.CharField(source='user.nickname', read_only=True) - images_count = serializers.SerializerMethodField() +class BugReportListSerializer(ModelSerializer): + user_name = CharField(source='user.nickname', read_only=True) + images_count = SerializerMethodField() class Meta: model = BugReport @@ -56,13 +60,13 @@ class BugReportListSerializer(serializers.ModelSerializer): 'created_at', 'updated_at', 'user_name', 'images_count' ] - def get_images_count(self, obj): - return obj.images.count() + async def get_images_count(self, obj): + return await obj.images.acount() -class BugReportDetailSerializer(serializers.ModelSerializer): - user_name = serializers.CharField(source='user.nickname', read_only=True) - user_avatar = serializers.ImageField(source='user.avatar', read_only=True) +class BugReportDetailSerializer(ModelSerializer): + user_name = CharField(source='user.nickname', read_only=True) + user_avatar = ImageField(source='user.avatar', read_only=True) images = BugReportImageSerializer(many=True, read_only=True) attachments = BugReportAttachmentSerializer(many=True, read_only=True) comments = BugReportCommentSerializer(many=True, read_only=True) @@ -78,14 +82,14 @@ class BugReportDetailSerializer(serializers.ModelSerializer): read_only_fields = ['user', 'status'] -class BugReportCreateSerializer(serializers.ModelSerializer): - images = serializers.ListField( - child=serializers.ImageField(), +class BugReportCreateSerializer(ModelSerializer): + images = ListField( + child=ImageField(), required=False, max_length=5 ) - attachments = serializers.ListField( - child=serializers.FileField(), + attachments = ListField( + child=FileField(), required=False, max_length=3 ) @@ -98,27 +102,27 @@ class BugReportCreateSerializer(serializers.ModelSerializer): max_size = 10 * 1024 * 1024 # 10MB for image in value: if image.size > max_size: - raise serializers.ValidationError(f"图片 {image.name} 大小超过10MB限制") + raise ValidationError(f"图片 {image.name} 大小超过10MB限制") return value def validate_attachments(self, value): max_size = 20 * 1024 * 1024 # 20MB for attachment in value: if attachment.size > max_size: - raise serializers.ValidationError(f"附件 {attachment.name} 大小超过20MB限制") + raise ValidationError(f"附件 {attachment.name} 大小超过20MB限制") return value - def create(self, validated_data): + async def acreate(self, validated_data): images_data = validated_data.pop('images', []) attachments_data = validated_data.pop('attachments', []) - bug_report = BugReport.objects.create(**validated_data) + bug_report = await BugReport.objects.acreate(**validated_data) for image in images_data: - BugReportImage.objects.create(bug_report=bug_report, image=image) + await BugReportImage.objects.acreate(bug_report=bug_report, image=image) for attachment in attachments_data: - BugReportAttachment.objects.create( + await BugReportAttachment.objects.acreate( bug_report=bug_report, file=attachment, filename=attachment.name, @@ -128,14 +132,14 @@ class BugReportCreateSerializer(serializers.ModelSerializer): return bug_report -class BugReportCommentCreateSerializer(serializers.ModelSerializer): - images = serializers.ListField( - child=serializers.ImageField(), +class BugReportCommentCreateSerializer(ModelSerializer): + images = ListField( + child=ImageField(), required=False, max_length=3 ) - attachments = serializers.ListField( - child=serializers.FileField(), + attachments = ListField( + child=FileField(), required=False, max_length=2 ) @@ -148,27 +152,27 @@ class BugReportCommentCreateSerializer(serializers.ModelSerializer): max_size = 10 * 1024 * 1024 # 10MB for image in value: if image.size > max_size: - raise serializers.ValidationError(f"图片 {image.name} 大小超过10MB限制") + raise ValidationError(f"图片 {image.name} 大小超过10MB限制") return value def validate_attachments(self, value): max_size = 20 * 1024 * 1024 # 20MB for attachment in value: if attachment.size > max_size: - raise serializers.ValidationError(f"附件 {attachment.name} 大小超过20MB限制") + raise ValidationError(f"附件 {attachment.name} 大小超过20MB限制") return value - def create(self, validated_data): + async def acreate(self, validated_data): images_data = validated_data.pop('images', []) attachments_data = validated_data.pop('attachments', []) - comment = BugReportComment.objects.create(**validated_data) + comment = await BugReportComment.objects.acreate(**validated_data) for image in images_data: - BugReportCommentImage.objects.create(comment=comment, image=image) + await BugReportCommentImage.objects.acreate(comment=comment, image=image) for attachment in attachments_data: - BugReportCommentAttachment.objects.create( + await BugReportCommentAttachment.objects.acreate( comment=comment, file=attachment, filename=attachment.name, diff --git a/bug/views.py b/bug/views.py index db97a3a..a2bd87c 100644 --- a/bug/views.py +++ b/bug/views.py @@ -1,10 +1,14 @@ -from rest_framework import generics, permissions, status +from rest_framework import status from rest_framework.response import Response from rest_framework.parsers import MultiPartParser, FormParser, JSONParser +from asgiref.sync import sync_to_async from django.utils import timezone 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 adrf import generics + from .models import BugReport, BugReportComment from .serializers import ( BugReportListSerializer, @@ -35,12 +39,15 @@ class BugReportListCreateAPIView(generics.ListCreateAPIView): operation_description='获取当前用户提交的所有Bug反馈列表', responses={200: success_response, 401: unauthorized_response} ) - def list(self, request, *args, **kwargs): + async def list(self, request, *args, **kwargs): queryset = self.get_queryset() - serializer = self.get_serializer(queryset, many=True) + items = [x async for x in queryset] + serializer = self.get_serializer(items, many=True) + # 兜底:user_name 触发 user 外键懒加载、images_count 内部有同步 ORM count() + data = await sync_to_async(lambda: serializer.data)() return create_standardized_response( code=ResponseCode.SUCCESS, - data=serializer.data, + data=data, message='获取Bug反馈列表成功' ) @@ -51,14 +58,18 @@ class BugReportListCreateAPIView(generics.ListCreateAPIView): request_body=BugReportCreateSerializer, responses={201: success_response, 400: error_response, 401: unauthorized_response} ) - def create(self, request, *args, **kwargs): + async def create(self, request, *args, **kwargs): serializer = self.get_serializer(data=request.data) - if serializer.is_valid(): - bug_report = serializer.save(user=request.user) + is_valid = await sync_to_async(serializer.is_valid)() + if is_valid: + # 兜底:serializer.create 内部有同步 ORM(BugReport/BugReportImage/BugReportAttachment 的 create) + bug_report = await sync_to_async(serializer.save)(user=request.user) detail_serializer = BugReportDetailSerializer(bug_report) + # 兜底:detail 序列化会触发 user/images/attachments/comments 的同步懒加载 + data = await sync_to_async(lambda: detail_serializer.data)() return create_standardized_response( code=ResponseCode.SUCCESS, - data=detail_serializer.data, + data=data, message='Bug反馈提交成功', status_code=status.HTTP_201_CREATED ) @@ -88,13 +99,15 @@ class BugReportDetailAPIView(generics.RetrieveAPIView): ], responses={200: success_response, 401: unauthorized_response, 404: not_found_response} ) - def retrieve(self, request, *args, **kwargs): + async def retrieve(self, request, *args, **kwargs): try: - instance = self.get_object() + instance = await self.get_object() serializer = self.get_serializer(instance) + # 兜底:detail 序列化会触发 user/images/attachments/comments 的同步懒加载 + data = await sync_to_async(lambda: serializer.data)() return create_standardized_response( code=ResponseCode.SUCCESS, - data=serializer.data, + data=data, message='获取Bug反馈详情成功' ) except BugReport.DoesNotExist: @@ -119,10 +132,10 @@ class BugReportCommentCreateAPIView(generics.CreateAPIView): request_body=BugReportCommentCreateSerializer, responses={201: success_response, 400: error_response, 401: unauthorized_response, 404: not_found_response} ) - def create(self, request, *args, **kwargs): + async def create(self, request, *args, **kwargs): bug_report_id = kwargs.get('bug_report_id') try: - bug_report = BugReport.objects.get(id=bug_report_id, user=request.user) + bug_report = await BugReport.objects.aget(id=bug_report_id, user=request.user) except BugReport.DoesNotExist: return create_standardized_error_response( code=ResponseCode.NOT_FOUND, @@ -131,16 +144,20 @@ class BugReportCommentCreateAPIView(generics.CreateAPIView): ) serializer = self.get_serializer(data=request.data) - if serializer.is_valid(): - comment = serializer.save( + is_valid = await sync_to_async(serializer.is_valid)() + if is_valid: + # 兜底:serializer.create 内部有同步 ORM(comment 及其图片/附件的 create) + comment = await sync_to_async(serializer.save)( user=request.user, bug_report=bug_report, is_admin=False ) comment_serializer = BugReportCommentSerializer(comment) + # 兜底:comment 序列化会触发 user 外键与 images/attachments 的同步懒加载 + data = await sync_to_async(lambda: comment_serializer.data)() return create_standardized_response( code=ResponseCode.SUCCESS, - data=comment_serializer.data, + data=data, message='回复成功', status_code=status.HTTP_201_CREATED ) diff --git a/chat/serializers.py b/chat/serializers.py index 09f2384..d7e3eba 100644 --- a/chat/serializers.py +++ b/chat/serializers.py @@ -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, + } diff --git a/chat/views.py b/chat/views.py index 3752527..1ddd6f5 100644 --- a/chat/views.py +++ b/chat/views.py @@ -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, ) diff --git a/chunyu_project/settings.py b/chunyu_project/settings.py index 1ff4a72..515462d 100644 --- a/chunyu_project/settings.py +++ b/chunyu_project/settings.py @@ -157,6 +157,7 @@ INSTALLED_APPS = [ 'django.contrib.messages', 'django.contrib.staticfiles', 'rest_framework', + 'adrf', 'rest_framework_simplejwt', 'corsheaders', 'django_filters', @@ -168,26 +169,62 @@ INSTALLED_APPS = [ ] # Database -DATABASES = { - 'default': { - 'ENGINE': 'django.db.backends.mysql', - 'NAME': 'chunyu_project', - 'USER': 'root', - 'PASSWORD': os.environ.get('MYSQL_PASSWORD', ''), - 'HOST': '192.168.5.7', - 'PORT': '3306', - 'OPTIONS': { - 'charset': 'utf8mb4', - 'init_command': "SET sql_mode='STRICT_TRANS_TABLES'", - }, - 'CONN_MAX_AGE': 3600, - 'CONN_HEALTH_CHECKS': True, +# ============================================ +# 异步化架构:默认 PostgreSQL(经 psycopg3 支持全链路异步 ORM:aget/afilter/acreate) +# 通过环境变量切换,兼容 MySQL(仅限本地过渡,MySQL 后端不支持异步 ORM) +# 注意:异步 ORM 不支持持久连接,CONN_MAX_AGE 必须为 0 +# ============================================ +DB_ENGINE = os.environ.get('DB_ENGINE', 'django.db.backends.postgresql') +DB_NAME = os.environ.get('DB_NAME', 'chunyu') +DB_USER = os.environ.get('DB_USER', 'postgres') +DB_PASSWORD = os.environ.get('DB_PASSWORD', '') +DB_HOST = os.environ.get('DB_HOST', 'localhost') +DB_PORT = os.environ.get('DB_PORT', '5432') + +if DB_ENGINE == 'django.db.backends.postgresql': + DATABASES = { + 'default': { + 'ENGINE': DB_ENGINE, + 'NAME': DB_NAME, + 'USER': DB_USER, + 'PASSWORD': DB_PASSWORD, + 'HOST': DB_HOST, + 'PORT': DB_PORT, + 'CONN_MAX_AGE': 0, # 异步 ORM 必须禁用持久连接 + 'CONN_HEALTH_CHECKS': False, + 'OPTIONS': { + 'connect_timeout': 5, + }, + } + } +elif DB_ENGINE == 'django.db.backends.mysql': + DATABASES = { + 'default': { + 'ENGINE': DB_ENGINE, + 'NAME': DB_NAME, + 'USER': DB_USER, + 'PASSWORD': DB_PASSWORD, + 'HOST': DB_HOST, + 'PORT': DB_PORT, + 'OPTIONS': { + 'charset': 'utf8mb4', + 'init_command': "SET sql_mode='STRICT_TRANS_TABLES'", + }, + 'CONN_MAX_AGE': 3600, + 'CONN_HEALTH_CHECKS': True, + } + } +else: + DATABASES = { + 'default': { + 'ENGINE': DB_ENGINE, + 'NAME': DB_NAME, + } } -} # Redis -REDIS_HOST = '192.168.5.7' -REDIS_PORT = 6379 +REDIS_HOST = os.environ.get('REDIS_HOST', '192.168.5.7') +REDIS_PORT = os.environ.get('REDIS_PORT', '6379') REDIS_DB = 0 REDIS_PASSWORD = os.environ.get('REDIS_PASSWORD', '') SESSION_DB = 1 diff --git a/currency/views.py b/currency/views.py index b5d9424..5bde97f 100644 --- a/currency/views.py +++ b/currency/views.py @@ -1,19 +1,44 @@ +import asyncio +import json +import hashlib from datetime import datetime -import requests +import aiohttp +from asgiref.sync import sync_to_async from rest_framework.permissions import AllowAny -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from drf_yasg.utils import swagger_auto_schema from drf_yasg import openapi from chunyu_project.common_schemas import success_response, error_response -from .cache import get_cached_data, set_cached_data +from django.core.cache import caches + +_cache = caches['default'] + +CACHE_TIMEOUT = 900 LATEST_RATES_URL = "https://open.er-api.com/v6/latest/{base}" +def _make_cache_key(prefix, *args, **kwargs): + raw = json.dumps({'args': args, 'kwargs': kwargs}, sort_keys=True, default=str) + suffix = hashlib.md5(raw.encode('utf-8')).hexdigest() + return f'currency_{prefix}_{suffix}' + + +async def get_cached_data(prefix, *args, **kwargs): + cache_key = _make_cache_key(prefix, *args, **kwargs) + # django-redis 为同步客户端,sync_to_async 兜底 + return await sync_to_async(_cache.get)(cache_key) + + +async def set_cached_data(prefix, data, *args, timeout=CACHE_TIMEOUT, **kwargs): + cache_key = _make_cache_key(prefix, *args, **kwargs) + await sync_to_async(_cache.set)(cache_key, data, timeout=timeout) + + class CurrencyRatesView(APIView): """ 汇率查询视图 - 获取指定基准货币对所有支持货币的汇率 @@ -35,11 +60,11 @@ class CurrencyRatesView(APIView): ], responses={200: success_response, 400: error_response, 502: error_response} ) - def get(self, request): + async def get(self, request): base = request.GET.get('base', 'USD').strip().upper() if not base: base = 'USD' - return self._fetch_rates(base) + return await self._fetch_rates(base) @swagger_auto_schema( tags=['Currency'], @@ -53,16 +78,16 @@ class CurrencyRatesView(APIView): ), responses={200: success_response, 400: error_response, 502: error_response} ) - def post(self, request): + async def post(self, request): base = request.data.get('base', 'USD').strip().upper() if isinstance(request.data, dict) else 'USD' if not base: base = 'USD' - return self._fetch_rates(base) + return await self._fetch_rates(base) - def _fetch_rates(self, base): + async def _fetch_rates(self, base): base = base.upper() - cached = get_cached_data('rates', base) + cached = await get_cached_data('rates', base) if cached is not None: return Response( {"code": 200, "message": "success (cached)", "data": cached}, @@ -71,15 +96,16 @@ class CurrencyRatesView(APIView): try: url = LATEST_RATES_URL.format(base=base) - response = requests.get(url, timeout=10) + timeout = aiohttp.ClientTimeout(total=10) + async with aiohttp.ClientSession(timeout=timeout) as session: + async with session.get(url) as response: + if response.status != 200: + return Response( + {"code": 502, "message": "汇率服务异常,请稍后重试", "data": None}, + status=status.HTTP_502_BAD_GATEWAY, + ) + data = await response.json() - if response.status_code != 200: - return Response( - {"code": 502, "message": "汇率服务异常,请稍后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY, - ) - - data = response.json() if data.get('result') != 'success': return Response( {"code": 502, "message": f"汇率服务返回错误:{data.get('error-type', '未知错误')}", "data": None}, @@ -93,23 +119,18 @@ class CurrencyRatesView(APIView): 'next_update': data.get('time_next_update_utc', ''), } - set_cached_data('rates', rates_data, base) + await set_cached_data('rates', rates_data, base) return Response( {"code": 200, "message": "success", "data": rates_data}, status=status.HTTP_200_OK, ) - except requests.exceptions.Timeout: + except (aiohttp.ClientError, asyncio.TimeoutError): return Response( {"code": 502, "message": "请求超时,请稍后重试", "data": None}, status=status.HTTP_502_BAD_GATEWAY, ) - except requests.exceptions.ConnectionError: - return Response( - {"code": 502, "message": "网络连接异常,请检查网络后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY, - ) except Exception as e: return Response( {"code": 500, "message": f"服务器内部错误:{str(e)}", "data": None}, @@ -147,8 +168,8 @@ class CurrenciesListView(APIView): operation_description='返回系统常用的货币代码及名称(带缓存)', responses={200: success_response, 500: error_response} ) - def get(self, request): - cached = get_cached_data('currencies') + async def get(self, request): + cached = await get_cached_data('currencies') if cached is not None: return Response( {"code": 200, "message": "success (cached)", "data": cached}, @@ -162,7 +183,7 @@ class CurrenciesListView(APIView): 'updated_at': datetime.now().strftime('%Y-%m-%d %H:%M:%S'), } - set_cached_data('currencies', data) + await set_cached_data('currencies', data) return Response( {"code": 200, "message": "success", "data": data}, @@ -210,11 +231,11 @@ class CurrencyConvertView(APIView): ], responses={200: success_response, 400: error_response, 502: error_response} ) - def get(self, request): + async def get(self, request): from_code = request.GET.get('from', '').strip().upper() to_code = request.GET.get('to', '').strip().upper() amount = request.GET.get('amount', '').strip() - return self._convert(from_code, to_code, amount) + return await self._convert(from_code, to_code, amount) @swagger_auto_schema( tags=['Currency'], @@ -231,14 +252,14 @@ class CurrencyConvertView(APIView): ), responses={200: success_response, 400: error_response, 502: error_response} ) - def post(self, request): + async def post(self, request): from_code = request.data.get('from', '').strip().upper() if isinstance(request.data, dict) else '' to_code = request.data.get('to', '').strip().upper() if isinstance(request.data, dict) else '' amount = request.data.get('amount', '') if isinstance(request.data, dict) else '' amount = str(amount).strip() if amount else '' - return self._convert(from_code, to_code, amount) + return await self._convert(from_code, to_code, amount) - def _convert(self, from_code, to_code, amount): + async def _convert(self, from_code, to_code, amount): if not from_code: return Response( {"code": 400, "message": "参数错误:from 不能为空", "data": None}, @@ -268,7 +289,7 @@ class CurrencyConvertView(APIView): status=status.HTTP_400_BAD_REQUEST, ) - cached_rates = get_cached_data('rates', from_code) + cached_rates = await get_cached_data('rates', from_code) rates = None if cached_rates is not None: @@ -277,15 +298,16 @@ class CurrencyConvertView(APIView): if rates is None: try: url = LATEST_RATES_URL.format(base=from_code) - response = requests.get(url, timeout=10) + timeout = aiohttp.ClientTimeout(total=10) + async with aiohttp.ClientSession(timeout=timeout) as session: + async with session.get(url) as response: + if response.status != 200: + return Response( + {"code": 502, "message": "汇率服务异常,请稍后重试", "data": None}, + status=status.HTTP_502_BAD_GATEWAY, + ) + data = await response.json() - if response.status_code != 200: - return Response( - {"code": 502, "message": "汇率服务异常,请稍后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY, - ) - - data = response.json() if data.get('result') != 'success': return Response( {"code": 502, "message": f"汇率服务返回错误:{data.get('error-type', '未知错误')}", "data": None}, @@ -298,19 +320,14 @@ class CurrencyConvertView(APIView): 'last_updated': data.get('time_last_update_utc', datetime.now().strftime('%Y-%m-%d %H:%M:%S')), 'next_update': data.get('time_next_update_utc', ''), } - set_cached_data('rates', rates_data, from_code) + await set_cached_data('rates', rates_data, from_code) rates = rates_data['rates'] - except requests.exceptions.Timeout: + except (aiohttp.ClientError, asyncio.TimeoutError): return Response( {"code": 502, "message": "请求超时,请稍后重试", "data": None}, status=status.HTTP_502_BAD_GATEWAY, ) - except requests.exceptions.ConnectionError: - return Response( - {"code": 502, "message": "网络连接异常,请检查网络后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY, - ) except Exception as e: return Response( {"code": 500, "message": f"服务器内部错误:{str(e)}", "data": None}, diff --git a/docs/ADRF_CONVERSION_GUIDE.md b/docs/ADRF_CONVERSION_GUIDE.md new file mode 100644 index 0000000..11caffc --- /dev/null +++ b/docs/ADRF_CONVERSION_GUIDE.md @@ -0,0 +1,74 @@ +# ADRF 异步化转换规范(chunyu_project 视图层) + +## 目标 +将 DRF 同步视图迁移到 ADRF(Asynchronous Django REST Framework)+ Django 5.x 异步 ORM,实现从视图到数据库的全链路非阻塞。 + +## 核心规则 + +### 1. 导入替换 +- `from rest_framework.views import APIView` → `from adrf.views import APIView` +- `from rest_framework import generics` / `rest_framework.generics.XXX` → `from adrf import generics`(或 `from adrf.generics import ListCreateAPIView` 等) +- `from rest_framework.viewsets import ModelViewSet/ReadOnlyModelViewSet` → `from adrf.viewsets import ModelViewSet, ReadOnlyModelViewSet` +- `from rest_framework.mixins import ...` → `from adrf import mixins` +- 保留:`Response`, `status`, `permissions`, `filters`, `pagination`, `serializers`, `django_filters`, `drf_yasg` 的导入不变。 +- `get_object_or_404` → `from adrf.generics import aget_object_or_404`(在 async 上下文中使用) + +### 2. 处理器异步化 +- 视图类中的 HTTP 方法处理器改为 `async def`(get/post/put/patch/delete/list/retrieve/create/update/destroy 及自定义 action)。 +- ADRF 特性:类内任意一个处理器是 async,dispatch 即走事件循环;残留的同步处理器会被自动 `sync_to_async` 兜底(不会崩,但会占线程池——尽量全部转完)。 + +### 3. 同步 ORM → 异步 ORM(Django 5.x,逐个替换,不可遗漏) +| 同步 | 异步 | +|---|---| +| `Model.objects.get(...)` | `await Model.objects.aget(...)` | +| `filter(...).first()` | `await Model.objects.filter(...).afirst()` | +| `filter(...).exists()` | `await Model.objects.filter(...).aexists()` | +| `filter(...).count()` | `await Model.objects.filter(...).acount()` | +| `Model.objects.create(...)` | `await Model.objects.acreate(...)` | +| `obj.save()` | `await obj.asave()` | +| `obj.delete()` | `await obj.adelete()` | +| `queryset.update(...)` | `await queryset.aupdate(...)` | +| `for x in queryset:` | `async for x in queryset:` | +| `list(queryset)` / 切片后遍历 | `results = [x async for x in queryset[start:end]]` | +| `get_object_or_404(...)` | `await aget_object_or_404(...)` | +| `len(queryset)`(避免) | `acount()` | + +注意: +- **切片**:`queryset[0:10]` 本身是惰性的,切片可以保留;但取值/遍历必须 async。 +- **values()/values_list()**:`[x async for x in Model.objects.filter(...).values(...)]` 合法。 +- **select_related/prefetch_related 链**:保留,遍历时 async。 +- **聚合**:`await Model.objects.aggregate(...)`(Django 5.x 支持 aggregate 的异步版本是 `Model.objects.acount` 等;`aggregate()` 无原生异步 —— 用 `await sync_to_async(Model.objects.aggregate)(...)`)。 +- **事务** `transaction.atomic()`:在 async 上下文中改为 `await sync_to_async(fn)()` 整体包裹,或使用 `transaction.atomic` 的异步支持 `async with transaction.acreate_agent`——最稳妥是 `await sync_to_async(sync_business_fn)()`。 +- **Q 对象/复杂查询**:构造部分是同步的,保留。 + +### 4. 外部 HTTP 调用 +- `import requests` + `requests.get/post(...)` → 改用 `aiohttp`(已安装): +```python +import aiohttp +async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=10)) as session: + async with session.get(url, params=..., headers=...) as resp: + data = await resp.json() +``` +- 文件上传(如百度语音/图片识别)用 `aiohttp.FormData()`。 +- 若改写风险过大(复杂 multipart),可用 `await sync_to_async(requests.post)(...)` 兜底,但要加注释 `# TODO: aiohttp 化`。 + +### 5. 必须保持不变 +- 所有 `@swagger_auto_schema` 装饰器原样保留(drf-yasg 只附加元数据,与 async 兼容)。 +- 所有 permission_classes、authentication_classes、pagination/filter 配置不变。 +- URL 路由文件(urls.py)不改 —— as_view() 不变。 +- 返回的 JSON 结构、状态码、错误信息完全不变(前端依赖这些契约)。 +- Celery `.delay()` 调用保留(可在 async 中直接调用;如需保险用 `await sync_to_async(task.delay)(...)`)。 +- `request.data` / `request.query_params` / `request.user` 用法不变(ADRF 的 AsyncRequest 兼容)。 + +### 6. 缓存操作 +- `cache.get/set` 是同步网络 I/O → `await sync_to_async(cache.get)(key)` 或保持 django-redis 同步调用并用 sync_to_async 包裹。 +- 简单做法:`from asgiref.sync import sync_to_async`,然后 `await sync_to_async(cache.set)(key, val, ttl)`。 + +### 7. 文件/验证 +- 每改完一个文件执行:`python -m py_compile ` 确保语法正确。 +- 不要运行服务器/测试(由主会话统一验证)。 +- 不确定某个 ORM 调用是否有异步版本时,用 `sync_to_async` 包裹并加 `# TODO: async ORM` 注释 —— 宁可兜底也不要留下裸同步 ORM 调用(会抛 SynchronousOnlyOperation)。 + +### 8. 输出要求 +- 逐文件报告:改了哪些处理器、哪些 ORM 调用、哪些外部 HTTP、是否有 sync_to_async 兜底点。 +- 列出任何你不敢改的复杂点(如嵌套事务、信号)。 diff --git a/docs/ADRF_SERIALIZER_GUIDE.md b/docs/ADRF_SERIALIZER_GUIDE.md new file mode 100644 index 0000000..e0d1f09 --- /dev/null +++ b/docs/ADRF_SERIALIZER_GUIDE.md @@ -0,0 +1,57 @@ +# ADRF 原生异步序列化器转换规范(第二阶段:100% 异步化) + +## 目标 +将 serializer 层的同步 ORM 与 MethodField 全部原生异步化,消灭视图层 `sync_to_async(serializer...)` 兜底。 + +## 核心 API(adrf 0.1.14) +```python +# 导入替换 +from rest_framework.serializers import ( + ModelSerializer, Serializer, SerializerMethodField, ... +) +# → 全部改为 +from adrf.serializers import ( + ModelSerializer, Serializer, SerializerMethodField, ... +) +# adrf.serializers 同时重导出了全部异步化字段(CharField/IntegerField/... 均已支持) +``` + +### 1. MethodField 异步化 +```python +class XSerializer(ModelSerializer): + foo = SerializerMethodField() + + async def get_foo(self, obj): # async def 即可,adrf.fields.SerializerMethodField 支持 + count = await Related.objects.filter(x=obj).acount() + return count +``` + +### 2. 序列化输出 +```python +# 视图中 +data = await serializer.adata # 异步属性(替代 serializer.data) +# 或在 adrf mixins/泛型内部使用 adrf.mixins.get_data(serializer) +``` +- `adata` 内部逐字段异步调 to_representation,MethodField 异步方法会被正确 await。 +- **many=True** 同样支持:`data = await serializer.adata`(ListSerializer 已被 adrf BaseSerializer.many_init 覆盖)。 +- nested Serializer:嵌套的 serializer 也必须来自 adrf.serializers,否则其 .data 是同步求值。 + +### 3. 写路径 +```python +await serializer.asave() # 替代 sync_to_async(serializer.save) +instance = await serializer.asave() # 返回 instance(同 DRF .save() 语义) +# serializer.create(...) → 改写为 async def acreate(self, validated_data),内部用 await Model.objects.acreate(...) +# serializer.update(...) → async def aupdate(self, instance, validated_data),内部 aget/asave/aupdate +``` +- `is_valid()` 是纯 CPU 校验(无 DB 除非 validator 带查询)——保持同步调用即可;若 validators 内部有 DB 查询(如 UniqueValidator 会查库),视图侧仍需 `await sync_to_async(serializer.is_valid)()` 或改自定义 validator 为 async。**UniqueValidator 场景保留 sync_to_async 包裹 is_valid。** + +### 4. 约束 +- **不得**在同步方法(get_queryset、get_serializer_class、validate 等 hooks)中调用 ORM——保持现状。 +- `validate(self, attrs)` 内的 DB查询(若存在)→ 改 `async def validate`(adrf 支持异步 validate?——不支持!validate 由 is_valid 同步调用。validate 内的 DB查询必须改用 CustomValidator 异步类或保留视图侧包裹)。遇到时在报告中列出。 +- ModelSerializer 字段声明(fields=..., read_only_fields 等)不变。 + +### 5. 每文件验证 +`python -m py_compile ` 必须 0 退出。不运行服务器。 + +### 6. 报告要求 +逐文件列出:改动的 MethodField 数、acreate/aupdate 重写数、残留的 sync_to_async 必要点(含原因)。 diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md new file mode 100644 index 0000000..d0f206c --- /dev/null +++ b/docs/ARCHITECTURE.md @@ -0,0 +1,151 @@ +# 可乐平台 异步化架构(ADRF + Granian + PostgreSQL + Nginx) + +## 架构总览 + +``` +┌──────────┐ :8080 ┌─────────────────────┐ +│ 浏览器 │──────────▶│ Nginx (frontend) │ +└──────────┘ │ ├─ / → React SPA 静态产物 (build/, 动静分离) + │ ├─ /api 等前缀 ─┐ + │ └─ /ws/ ─┤ 反向代理 + └──────────────────┼───┘ + ▼ + ┌───────────────────────────────┐ + │ Granian (Rust) :8000 │ + │ Django 5.2 ASGI + ADRF │ + │ ├─ 异步视图 (async def) │ + │ ├─ 异步 ORM (aget/afilter) │ + │ └─ Channels (WebSocket) │ + └───────┬──────────────┬────────┘ + ▼ ▼ + ┌──────────────┐ ┌──────────────┐ + │ PostgreSQL 16 │ │ Redis 7 │ + │ psycopg3 异步 │ │ cache/session│ + │ ORM 驱动 │ │ celery/ws层 │ + └──────────────┘ └──────────────┘ +``` + +## 技术选型 + +| 层 | 技术 | 说明 | +|---|---|---| +| Web 框架 | Django 5.2 + DRF 3.17 | 原生异步能力 | +| 异步视图 | **ADRF 0.1.14** | 所有 API 视图跑在异步事件循环上 | +| 应用服务器 | **Granian 2.8.2 (Rust)** | 替代 Daphne,ASGI + WebSocket,吞吐量显著提升 | +| 数据库 | **PostgreSQL 16** | Django 5 异步 ORM 必备(MySQL 后端不支持异步) | +| DB 驱动 | **psycopg3 (binary) 3.2.13** | 支持异步连接 | +| 缓存/队列 | Redis 7 | 缓存、Session、Celery broker、Channels layer | +| 前端 | React 19 + Vite 8 | 独立 SPA,与 Django 模板层完全解耦 | +| 反代/静态 | Nginx | SPA 静态资源直出 + API 反向代理 | + +## 关键设计决策 + +1. **全链路异步**:视图 `async def` → ORM `aget/afirst/aexists/acount/acreate/asave/aupdate/adelete` → PostgreSQL(psycopg3)。`CONN_MAX_AGE=0`(异步 ORM 不支持持久连接)。 +2. **ADRF 兼容性**:`view_is_async` 自动检测处理器;纯 CPU 视图(图形验证码、图像压缩、文本 diff)保持同步 `def`,ADRF 自动 `sync_to_async` 线程池兜底。 +3. **外部 HTTP 调用 aiohttp 化**:weather / air_quality / currency / getIpData 全部改为 `aiohttp.ClientSession`,无阻塞 I/O。 +4. **Redis 缓存 / Celery / SMTP / 密码哈希**:同步客户端用 `sync_to_async` 包裹(不阻塞事件循环)。 +5. **事务兜底**:`transaction.atomic` + `select_for_update` 段(钱包、任务领取)抽成同步闭包整体 `sync_to_async` 执行,保住行锁语义。 +6. **数据库可切换**:`DB_ENGINE` 环境变量支持 `django.db.backends.postgresql`(默认)与 `django.db.backends.mysql`(本地过渡)。 + +## 前后端分离 + +- 开发:`pnpm dev`(Vite HMR :5173),`BACKEND_URL` 环境变量指定后端(默认 `http://127.0.0.1:8002`,Docker 后端用 `http://127.0.0.1:8000`)。API 全部走 Vite proxy,无 CORS 负担。 +- 生产:`pnpm build` 产物由 Nginx 容器直出(`try_files → index.html` SPA 回退),带指纹资源 1 年 immutable 缓存,index.html no-cache。 +- 前后端仅通过 RESTful API + WebSocket 交互。 + +## 快速开始 + +### 一键部署(生产形态) +```bash +docker compose up -d --build +# 前端 http://localhost:8080 +# 后端 API http://localhost:8000 +# Swagger http://localhost:8080/swagger/ +``` + +### 本地开发 +```bash +# 1. 启动数据层 +docker compose up -d postgres redis + +# 2. 后端(异步栈) +cd chunyu_project +export DB_ENGINE=django.db.backends.postgresql DB_HOST=localhost DB_PORT=5432 \ + DB_NAME=chunyu DB_USER=chunyu DB_PASSWORD=chunyu_pg_pass \ + REDIS_HOST=localhost +python manage.py migrate +granian --interface asgi --host 0.0.0.0 --port 8000 chunyu_project.asgi:application + +# 3. 前端(Vite HMR) +cd chunyu_project_react +BACKEND_URL=http://127.0.0.1:8000 pnpm dev # http://localhost:5173 +``` + +## 异步化转换统计 + +- 覆盖 16 个应用、45 个视图文件:`api` `user` `article` `learn` `chat` `message` `history` `search` `bug` `tool` `weather` `air_quality` `currency` `shorturl` `logs` `apidirectory` `app` +- 处理器全部 `async def`(约 90+ 个);ORM 调用全部异步化 +- 兜底点(`sync_to_async`):serializer is_valid/save/data、Redis cache、SMTP/邮件任务、密码哈希、文件存储 I/O、`transaction.atomic` 事务段、Celery `.delay` +- 转换规范见 `chunyu_project/docs/ADRF_CONVERSION_GUIDE.md` + +## 验证记录 + +- ✅ `python manage.py check` 0 错误 +- ✅ URLCONF 全量导入成功(所有视图模块可加载) +- ✅ Granian 启动 + 冒烟:`/`、`/article/articles/`、`/learn/courses/`、`/api/apidirectory/categories/`、`/api/weather/?city=Beijing`、`/api/air-quality/?city=Shanghai`、`/api/currency/rates/?base=USD`、`/api/image-captcha/`、`/api/captcha/generate/`、`/swagger.json`(331KB) 全部 200 +- ✅ `docker compose config` 校验通过 +- ✅ Nginx 配置语法 `nginx -t` 通过 + +## 服务器部署(192.168.5.7 生产/联调环境) + +本地 Windows 宿主虚拟化暂时不可用(Hypervisor 未加载,见下),已将整套异步架构部署至本地 Linux 服务器: + +| 项 | 值 | +|---|---| +| 服务器 | Debian,16 核(`lan-19216857` SSH profile) | +| 部署目录 | `/mnt/sda/chunyu` | +| 访问地址 | `http://192.168.5.7:18080`(前端 SPA + API 反代) | +| 后端直连 | `http://192.168.5.7:18000`(Granian) | +| PostgreSQL | `192.168.5.7:15432`(容器内 5432) | +| Redis | `192.168.5.7:16379`(容器内 6379) | + +端口重映射原因:服务器既有服务占用 5432/6379/8080,覆盖文件 `docker-compose.server.yml` 将宿主端口改为 15432/16379/18000/18080,容器间通信仍走内部网络不受影响。 + +部署命令: +```bash +cd /mnt/sda/chunyu +docker compose -f docker-compose.yml -f docker-compose.server.yml up -d --build +``` + +服务器端到端验证记录(全部通过): +- ✅ 四容器运行:chunyu-backend(healthy) / chunyu-frontend / chunyu-postgres(healthy) / chunyu-redis(healthy) +- ✅ SPA 静态资源 Nginx 直出 200;`/api/apidirectory/categories/`、`/article/articles/`、`/swagger.json` 经反代 200 +- ✅ aiohttp 外部 API(天气)经反代 200 +- ✅ 异步 ORM 写路径:短链创建成功落库(PostgreSQL 66 张迁移表) +- ✅ Redis PONG(缓存/Session 层) +- ✅ WebSocket:Granian + Channels,JWT 认证后 accept 成功(匿名拒绝为原有业务设计,ChatConsumer 强制认证) +- ✅ 并发冒烟:80 并发 × 5 轮 = 400 请求,**400/400 成功,187.9 RPS**(穿透 Nginx → Granian → PostgreSQL 全链路) +- ✅ Windows 局域网访问 `http://192.168.5.7:18080/*` 全部 200 + +> 说明:`DJANGO_DEBUG` 插值默认 True(局域网无 TLS;False 会触发 `SECURE_SSL_REDIRECT` 301)。正式上 TLS 后设 `DJANGO_DEBUG=False`。 +> 本地 Windows 宿主:2026-09-05 22:22 重启后 HypervisorPresent=False(BCD hypervisorlaunchtype 被关),Docker Desktop 不可用;管理员执行 `bcdedit /set "{current}" hypervisorlaunchtype auto` 重启后可恢复本地 Docker。 + +## 环境变量 + +见 `docker.env.example`。核心项: + +| 变量 | 默认 | 说明 | +|---|---|---| +| `DB_ENGINE` | postgresql | django.db.backends.postgresql / mysql | +| `DB_HOST/DB_PORT/DB_NAME/DB_USER/DB_PASSWORD` | postgres/5432/chunyu/chunyu | 数据库连接 | +| `REDIS_HOST/REDIS_PORT/REDIS_PASSWORD` | redis/6379 | Redis 连接 | +| `POSTGRES_PASSWORD` | chunyu_pg_pass | PostgreSQL root 密码 | +| `DJANGO_DEBUG` | False | 生产必须 False | +| `CORS_ALLOWED_ORIGINS` | localhost:8080/5173 | 前端来源 | + +## 已知限制 + +1. **serializer 层仍有同步 ORM**(经 sync_to_async 兜底,功能正确但该段走线程池):`chat/serializers.py` ConversationSerializer 的 MethodField、`bug` 列表 get_images_count、若干 `serializer.create()`。后续可对 serializer 原生异步化。 +2. **Celery worker 不在默认编排内**(views 已适配:同步调用失败自动降级 submit_task 异步队列)。如需启用:`docker compose up -d` 后手动 `docker compose run backend celery -A chunyu_project worker -l info`。 +3. Channels ASGI Lifespan 在 Granian 下有 warning(`asginl` 接口可消除),不影响 WebSocket 功能。 +4. 纯 CPU 视图(验证码/压缩/diff)保持同步,ADRF 自动线程池兜底。 diff --git a/history/serializers.py b/history/serializers.py index 44cb364..ecc39c0 100644 --- a/history/serializers.py +++ b/history/serializers.py @@ -1,29 +1,29 @@ -from rest_framework import serializers +from adrf.serializers import ModelSerializer from django.utils import timezone from datetime import timedelta from .models import BrowsingHistory -class BrowsingHistorySerializer(serializers.ModelSerializer): +class BrowsingHistorySerializer(ModelSerializer): class Meta: model = BrowsingHistory fields = ['id', 'type', 'title', 'description', 'image', 'category', 'link', 'viewed_at'] read_only_fields = ['viewed_at'] -class BrowsingHistoryCreateSerializer(serializers.ModelSerializer): +class BrowsingHistoryCreateSerializer(ModelSerializer): class Meta: model = BrowsingHistory fields = ['type', 'title', 'description', 'image', 'category', 'link'] - def create(self, validated_data): + async def acreate(self, validated_data): user = self.context['request'].user link = validated_data.get('link', '') today = timezone.now().date() - existing = BrowsingHistory.objects.filter( + existing = await BrowsingHistory.objects.filter( user=user, link=link, viewed_at__date=today - ).first() + ).afirst() if existing: - existing.save() + await existing.asave() return existing - return BrowsingHistory.objects.create(user=user, **validated_data) + return await BrowsingHistory.objects.acreate(user=user, **validated_data) diff --git a/history/views.py b/history/views.py index c10a2ff..30ef538 100644 --- a/history/views.py +++ b/history/views.py @@ -1,14 +1,17 @@ -from rest_framework import generics, status +from rest_framework import status from rest_framework.permissions import IsAuthenticated from rest_framework.parsers import JSONParser -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.pagination import PageNumberPagination from rest_framework.response import Response -from django.shortcuts import get_object_or_404 +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 adrf import generics +from adrf.generics import aget_object_or_404 + from utils.response_codes import ResponseCode, create_standardized_response, create_standardized_error_response from .models import BrowsingHistory from .serializers import BrowsingHistorySerializer, BrowsingHistoryCreateSerializer @@ -62,13 +65,15 @@ class BrowsingHistoryListCreateView(generics.ListCreateAPIView): ], responses={200: success_response, 401: unauthorized_response} ) - def list(self, request, *args, **kwargs): + async def list(self, request, *args, **kwargs): queryset = self.filter_queryset(self.get_queryset()) - page = self.paginate_queryset(queryset) + # 兜底:DRF 分页器内部同步评估 queryset(count + 切片取值) + page = await sync_to_async(self.paginate_queryset)(queryset) if page is not None: serializer = self.get_serializer(page, many=True) return self.get_paginated_response(serializer.data) - serializer = self.get_serializer(queryset, many=True) + results = [x async for x in queryset] + serializer = self.get_serializer(results, many=True) return Response({'code': 10000, 'message': 'Success', 'data': serializer.data}) @swagger_auto_schema( @@ -78,10 +83,12 @@ class BrowsingHistoryListCreateView(generics.ListCreateAPIView): request_body=BrowsingHistoryCreateSerializer, responses={201: success_response, 400: error_response, 401: unauthorized_response} ) - def create(self, request, *args, **kwargs): + async def create(self, request, *args, **kwargs): serializer = self.get_serializer(data=request.data) - if serializer.is_valid(): - instance = serializer.save() + is_valid = await sync_to_async(serializer.is_valid)() + if is_valid: + # 兜底:serializer.create 内部有同步 ORM(filter/first/save/create) + instance = await sync_to_async(serializer.save)() return create_standardized_response( data=BrowsingHistorySerializer(instance).data, code=ResponseCode.SUCCESS, @@ -107,9 +114,9 @@ class BrowsingHistoryDeleteView(APIView): ], responses={204: '删除成功', 401: unauthorized_response, 404: not_found_response} ) - def delete(self, request, pk): - record = get_object_or_404(BrowsingHistory, pk=pk, user=request.user) - record.delete() + async def delete(self, request, pk): + record = await aget_object_or_404(BrowsingHistory, pk=pk, user=request.user) + await record.adelete() return create_standardized_response( code=ResponseCode.SUCCESS, message='删除成功', @@ -126,8 +133,8 @@ class BrowsingHistoryClearView(APIView): operation_description='清空当前用户的所有浏览历史记录', responses={200: success_response, 401: unauthorized_response} ) - def delete(self, request): - count = BrowsingHistory.objects.filter(user=request.user).delete()[0] + async def delete(self, request): + count = (await BrowsingHistory.objects.filter(user=request.user).adelete())[0] return create_standardized_response( data={'deleted': count}, code=ResponseCode.SUCCESS, diff --git a/learn/serializers.py b/learn/serializers.py index f21e6f5..394506e 100644 --- a/learn/serializers.py +++ b/learn/serializers.py @@ -1,16 +1,18 @@ -from rest_framework import serializers +from adrf.serializers import ( + ModelSerializer, SerializerMethodField, +) from .models import Course, Chapter, ChapterContent, CourseMaterial -class ChapterSerializer(serializers.ModelSerializer): - video_poster_url = serializers.SerializerMethodField() - video_local_url = serializers.SerializerMethodField() +class ChapterSerializer(ModelSerializer): + video_poster_url = SerializerMethodField() + video_local_url = SerializerMethodField() class Meta: model = Chapter fields = ['id', 'title', 'sort_order', 'duration', 'is_free', 'video_url', 'video_source', 'video_local_url', 'video_poster_url', 'created_at', 'updated_at'] - def get_video_poster_url(self, obj): + async def get_video_poster_url(self, obj): if obj.video_poster: request = self.context.get('request') if request: @@ -18,7 +20,7 @@ class ChapterSerializer(serializers.ModelSerializer): return obj.video_poster.url return '' - def get_video_local_url(self, obj): + async def get_video_local_url(self, obj): if obj.video_local: request = self.context.get('request') if request: @@ -27,22 +29,22 @@ class ChapterSerializer(serializers.ModelSerializer): return '' -class ChapterContentSerializer(serializers.ModelSerializer): +class ChapterContentSerializer(ModelSerializer): class Meta: model = ChapterContent fields = ['id', 'chapter', 'content_md', 'content_html', 'md_file_path', 'updated_at'] -class MaterialSerializer(serializers.ModelSerializer): - file_url = serializers.SerializerMethodField() - file_size_display = serializers.SerializerMethodField() - file_type_display = serializers.SerializerMethodField() +class MaterialSerializer(ModelSerializer): + file_url = SerializerMethodField() + file_size_display = SerializerMethodField() + file_type_display = SerializerMethodField() class Meta: model = CourseMaterial fields = ['id', 'title', 'file_url', 'file_type', 'file_type_display', 'file_size', 'file_size_display', 'download_count', 'sort_order', 'created_at'] - def get_file_url(self, obj): + async def get_file_url(self, obj): if obj.file: request = self.context.get('request') if request: @@ -50,7 +52,7 @@ class MaterialSerializer(serializers.ModelSerializer): return obj.file.url return '' - def get_file_size_display(self, obj): + async def get_file_size_display(self, obj): size = obj.file_size if size < 1024: return f'{size} B' @@ -59,7 +61,7 @@ class MaterialSerializer(serializers.ModelSerializer): else: return f'{size / (1024 * 1024):.1f} MB' - def get_file_type_display(self, obj): + async def get_file_type_display(self, obj): type_map = { 'pdf': 'PDF', 'code': '代码', @@ -72,14 +74,14 @@ class MaterialSerializer(serializers.ModelSerializer): return type_map.get(obj.file_type, '其他') -class ChapterSerializer(serializers.ModelSerializer): - video_poster_url = serializers.SerializerMethodField() +class ChapterSerializer(ModelSerializer): + video_poster_url = SerializerMethodField() class Meta: model = Chapter fields = ['id', 'title', 'sort_order', 'duration', 'is_free', 'video_url', 'video_poster_url', 'created_at', 'updated_at'] - def get_video_poster_url(self, obj): + async def get_video_poster_url(self, obj): if obj.video_poster: request = self.context.get('request') if request: @@ -88,17 +90,17 @@ class ChapterSerializer(serializers.ModelSerializer): return '' -class ChapterContentSerializer(serializers.ModelSerializer): +class ChapterContentSerializer(ModelSerializer): class Meta: model = ChapterContent fields = ['id', 'chapter', 'content_md', 'content_html', 'md_file_path', 'updated_at'] -class CourseListSerializer(serializers.ModelSerializer): - author_name = serializers.SerializerMethodField() - chapters_count = serializers.SerializerMethodField() - students_count = serializers.SerializerMethodField() - cover_image_url = serializers.SerializerMethodField() +class CourseListSerializer(ModelSerializer): + author_name = SerializerMethodField() + chapters_count = SerializerMethodField() + students_count = SerializerMethodField() + cover_image_url = SerializerMethodField() class Meta: model = Course @@ -109,13 +111,13 @@ class CourseListSerializer(serializers.ModelSerializer): 'sort_order', 'is_hot', 'is_new', 'created_at', 'updated_at', ] - def get_chapters_count(self, obj): - return obj.chapters.count() + async def get_chapters_count(self, obj): + return await obj.chapters.acount() - def get_students_count(self, obj): + async def get_students_count(self, obj): return 0 - def get_cover_image_url(self, obj): + async def get_cover_image_url(self, obj): if obj.cover_image: request = self.context.get('request') if request: @@ -123,18 +125,18 @@ class CourseListSerializer(serializers.ModelSerializer): return obj.cover_image.url return '' - def get_author_name(self, obj): + async def get_author_name(self, obj): if obj.author: return obj.author.get_full_name() or obj.author.username return '' -class CourseDetailSerializer(serializers.ModelSerializer): - author_name = serializers.SerializerMethodField() +class CourseDetailSerializer(ModelSerializer): + author_name = SerializerMethodField() chapters = ChapterSerializer(many=True, read_only=True) - chapters_count = serializers.SerializerMethodField() - students_count = serializers.SerializerMethodField() - cover_image_url = serializers.SerializerMethodField() + chapters_count = SerializerMethodField() + students_count = SerializerMethodField() + cover_image_url = SerializerMethodField() class Meta: model = Course @@ -145,13 +147,13 @@ class CourseDetailSerializer(serializers.ModelSerializer): 'sort_order', 'is_hot', 'is_new', 'created_at', 'updated_at', ] - def get_chapters_count(self, obj): - return obj.chapters.count() + async def get_chapters_count(self, obj): + return await obj.chapters.acount() - def get_students_count(self, obj): + async def get_students_count(self, obj): return 0 - def get_cover_image_url(self, obj): + async def get_cover_image_url(self, obj): if obj.cover_image: request = self.context.get('request') if request: @@ -159,25 +161,25 @@ class CourseDetailSerializer(serializers.ModelSerializer): return obj.cover_image.url return '' - def get_author_name(self, obj): + async def get_author_name(self, obj): if obj.author: return obj.author.get_full_name() or obj.author.username return '' -class CourseCreateUpdateSerializer(serializers.ModelSerializer): +class CourseCreateUpdateSerializer(ModelSerializer): class Meta: model = Course fields = ['title', 'description', 'category', 'level', 'cover_image', 'color', 'icon_name', 'status', 'sort_order', 'is_hot', 'is_new'] - def create(self, validated_data): + async def acreate(self, validated_data): validated_data['author'] = self.context['request'].user - return super().create(validated_data) + return await super().acreate(validated_data) -class CourseManageSerializer(serializers.ModelSerializer): - author_name = serializers.SerializerMethodField() - chapters_count = serializers.SerializerMethodField() +class CourseManageSerializer(ModelSerializer): + author_name = SerializerMethodField() + chapters_count = SerializerMethodField() class Meta: model = Course @@ -187,10 +189,10 @@ class CourseManageSerializer(serializers.ModelSerializer): 'author_name', 'created_at', 'updated_at', ] - def get_chapters_count(self, obj): - return obj.chapters.count() + async def get_chapters_count(self, obj): + return await obj.chapters.acount() - def get_author_name(self, obj): + async def get_author_name(self, obj): if obj.author: return obj.author.get_full_name() or obj.author.username return '' diff --git a/learn/views.py b/learn/views.py index c788930..c814151 100644 --- a/learn/views.py +++ b/learn/views.py @@ -3,11 +3,14 @@ from django.conf import settings from django.db.models import Count from django.utils import timezone from django.http import FileResponse, Http404 -from django.shortcuts import get_object_or_404 -from rest_framework import generics, status +from adrf import generics +from adrf.generics import aget_object_or_404 +from adrf.mixins import get_data +from adrf.views import APIView +from asgiref.sync import sync_to_async +from rest_framework import status from rest_framework.permissions import IsAuthenticated, AllowAny from rest_framework.parsers import JSONParser, MultiPartParser, FormParser -from rest_framework.views import APIView from drf_yasg import openapi from drf_yasg.utils import swagger_auto_schema @@ -65,6 +68,12 @@ class CourseListCreateView(generics.ListCreateAPIView): return qs + async def get(self, request, *args, **kwargs): + return await self.list(request, *args, **kwargs) + + async def post(self, request, *args, **kwargs): + return await self.create(request, *args, **kwargs) + @swagger_auto_schema( tags=['学习'], operation_summary='获取课程列表', @@ -76,22 +85,22 @@ class CourseListCreateView(generics.ListCreateAPIView): ], responses={200: success_response}, ) - def list(self, request, *args, **kwargs): + async def list(self, request, *args, **kwargs): queryset = self.get_queryset() serializer = self.get_serializer(queryset, many=True) - data = serializer.data + data = await get_data(serializer) if request.user.is_authenticated: - favorite_course_ids = set( - CourseFavorite.objects.filter(user=request.user).values_list('course_id', flat=True) - ) + favorite_course_ids = set([ + v async for v in CourseFavorite.objects.filter(user=request.user).values_list('course_id', flat=True) + ]) for item in data: item['is_favorited'] = item['id'] in favorite_course_ids - item['favorites_count'] = CourseFavorite.objects.filter(course_id=item['id']).count() + item['favorites_count'] = await CourseFavorite.objects.filter(course_id=item['id']).acount() else: for item in data: item['is_favorited'] = False - item['favorites_count'] = CourseFavorite.objects.filter(course_id=item['id']).count() + item['favorites_count'] = await CourseFavorite.objects.filter(course_id=item['id']).acount() return create_standardized_response(data=data, code=ResponseCode.SUCCESS) @@ -106,12 +115,13 @@ class CourseListCreateView(generics.ListCreateAPIView): 401: unauthorized_response, }, ) - def create(self, request, *args, **kwargs): + async def create(self, request, *args, **kwargs): serializer = self.get_serializer(data=request.data) - if serializer.is_valid(): - serializer.save() + if await sync_to_async(serializer.is_valid)(): + await sync_to_async(serializer.save)() + data = await get_data(serializer) return create_standardized_response( - data=serializer.data, + data=data, code=ResponseCode.SUCCESS, message='课程创建成功', status_code=status.HTTP_201_CREATED @@ -140,17 +150,17 @@ class CourseDetailView(APIView): 404: not_found_response, }, ) - def get(self, request, pk): - course = get_object_or_404(Course, pk=pk) + async def get(self, request, pk): + course = await aget_object_or_404(Course, pk=pk) serializer = CourseDetailSerializer(course, context={'request': request}) - data = serializer.data + data = await get_data(serializer) if request.user.is_authenticated: - data['is_favorited'] = CourseFavorite.objects.filter( + data['is_favorited'] = await CourseFavorite.objects.filter( user=request.user, course=course - ).exists() + ).aexists() else: data['is_favorited'] = False - data['favorites_count'] = CourseFavorite.objects.filter(course=course).count() + data['favorites_count'] = await CourseFavorite.objects.filter(course=course).acount() return create_standardized_response(data=data, code=ResponseCode.SUCCESS) @swagger_auto_schema( @@ -169,8 +179,8 @@ class CourseDetailView(APIView): 404: not_found_response, }, ) - def put(self, request, pk): - course = get_object_or_404(Course, pk=pk) + async def put(self, request, pk): + course = await aget_object_or_404(Course, pk=pk) if course.author != request.user: return create_standardized_error_response( code=ResponseCode.VALIDATION_ERROR, @@ -178,9 +188,10 @@ class CourseDetailView(APIView): status_code=status.HTTP_403_FORBIDDEN ) serializer = CourseCreateUpdateSerializer(course, data=request.data, partial=True, context={'request': request}) - if serializer.is_valid(): - serializer.save() - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS, message='课程更新成功') + if await sync_to_async(serializer.is_valid)(): + await sync_to_async(serializer.save)() + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS, message='课程更新成功') return create_standardized_error_response( data=serializer.errors, code=ResponseCode.VALIDATION_ERROR, @@ -201,15 +212,15 @@ class CourseDetailView(APIView): 404: not_found_response, }, ) - def delete(self, request, pk): - course = get_object_or_404(Course, pk=pk) + async def delete(self, request, pk): + course = await aget_object_or_404(Course, pk=pk) if course.author != request.user: return create_standardized_error_response( code=ResponseCode.VALIDATION_ERROR, message='无权操作', status_code=status.HTTP_403_FORBIDDEN ) - course.delete() + await course.adelete() return create_standardized_response(code=ResponseCode.SUCCESS, message='课程删除成功', status_code=status.HTTP_204_NO_CONTENT) @@ -227,6 +238,12 @@ class ChapterListCreateView(generics.ListCreateAPIView): course_id = self.kwargs['course_id'] return Chapter.objects.filter(course_id=course_id) + async def get(self, request, *args, **kwargs): + return await self.list(request, *args, **kwargs) + + async def post(self, request, *args, **kwargs): + return await self.create(request, *args, **kwargs) + @swagger_auto_schema( tags=['学习'], operation_summary='获取章节列表', @@ -236,10 +253,11 @@ class ChapterListCreateView(generics.ListCreateAPIView): ], responses={200: success_response, 404: not_found_response}, ) - def list(self, request, *args, **kwargs): + async def list(self, request, *args, **kwargs): queryset = self.get_queryset() serializer = self.get_serializer(queryset, many=True) - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS) + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS) @swagger_auto_schema( tags=['学习'], @@ -256,14 +274,15 @@ class ChapterListCreateView(generics.ListCreateAPIView): 404: not_found_response, }, ) - def create(self, request, *args, **kwargs): + async def create(self, request, *args, **kwargs): course_id = self.kwargs['course_id'] - course = get_object_or_404(Course, pk=course_id) + course = await aget_object_or_404(Course, pk=course_id) serializer = self.get_serializer(data=request.data) - if serializer.is_valid(): - serializer.save(course=course) + if await sync_to_async(serializer.is_valid)(): + await sync_to_async(serializer.save)(course=course) + data = await get_data(serializer) return create_standardized_response( - data=serializer.data, + data=data, code=ResponseCode.SUCCESS, message='章节创建成功', status_code=status.HTTP_201_CREATED @@ -293,10 +312,11 @@ class ChapterDetailView(APIView): 404: not_found_response, }, ) - def get(self, request, pk): - chapter = get_object_or_404(Chapter, pk=pk) + async def get(self, request, pk): + chapter = await aget_object_or_404(Chapter, pk=pk) serializer = ChapterSerializer(chapter) - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS) + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS) @swagger_auto_schema( tags=['学习'], @@ -313,12 +333,13 @@ class ChapterDetailView(APIView): 404: not_found_response, }, ) - def put(self, request, pk): - chapter = get_object_or_404(Chapter, pk=pk) + async def put(self, request, pk): + chapter = await aget_object_or_404(Chapter, pk=pk) serializer = ChapterSerializer(chapter, data=request.data, partial=True) - if serializer.is_valid(): - serializer.save() - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS, message='章节更新成功') + if await sync_to_async(serializer.is_valid)(): + await sync_to_async(serializer.save)() + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS, message='章节更新成功') return create_standardized_error_response( data=serializer.errors, code=ResponseCode.VALIDATION_ERROR, @@ -338,9 +359,9 @@ class ChapterDetailView(APIView): 404: not_found_response, }, ) - def delete(self, request, pk): - chapter = get_object_or_404(Chapter, pk=pk) - chapter.delete() + async def delete(self, request, pk): + chapter = await aget_object_or_404(Chapter, pk=pk) + await chapter.adelete() return create_standardized_response(code=ResponseCode.SUCCESS, message='章节删除成功', status_code=status.HTTP_204_NO_CONTENT) @@ -361,14 +382,15 @@ class ChapterContentView(APIView): 404: not_found_response, }, ) - def get(self, request, chapter_id): - chapter = get_object_or_404(Chapter, pk=chapter_id) - content, created = ChapterContent.objects.get_or_create( + async def get(self, request, chapter_id): + chapter = await aget_object_or_404(Chapter, pk=chapter_id) + content, created = await ChapterContent.objects.aget_or_create( chapter=chapter, defaults={'content_md': '', 'content_html': ''} ) serializer = ChapterContentSerializer(content) - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS) + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS) @swagger_auto_schema( tags=['学习'], @@ -385,16 +407,17 @@ class ChapterContentView(APIView): 404: not_found_response, }, ) - def put(self, request, chapter_id): - chapter = get_object_or_404(Chapter, pk=chapter_id) - content, created = ChapterContent.objects.get_or_create( + async def put(self, request, chapter_id): + chapter = await aget_object_or_404(Chapter, pk=chapter_id) + content, created = await ChapterContent.objects.aget_or_create( chapter=chapter, defaults={'content_md': '', 'content_html': ''} ) serializer = ChapterContentSerializer(content, data=request.data, partial=True) - if serializer.is_valid(): - serializer.save() - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS, message='内容保存成功') + if await sync_to_async(serializer.is_valid)(): + await sync_to_async(serializer.save)() + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS, message='内容保存成功') return create_standardized_error_response( data=serializer.errors, code=ResponseCode.VALIDATION_ERROR, @@ -414,6 +437,9 @@ class MyCourseListView(generics.ListAPIView): qs = qs.filter(status=st) return qs.order_by('-updated_at') + async def get(self, request, *args, **kwargs): + return await self.list(request, *args, **kwargs) + @swagger_auto_schema( tags=['学习'], operation_summary='获取我的课程列表', @@ -426,10 +452,11 @@ class MyCourseListView(generics.ListAPIView): 401: unauthorized_response, }, ) - def list(self, request, *args, **kwargs): + async def list(self, request, *args, **kwargs): queryset = self.get_queryset() serializer = self.get_serializer(queryset, many=True) - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS) + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS) class MyCourseBatchView(APIView): @@ -454,7 +481,7 @@ class MyCourseBatchView(APIView): 401: unauthorized_response, }, ) - def post(self, request): + async def post(self, request): ids = request.data.get('ids', []) action = request.data.get('action', '') if not ids or action not in ('publish', 'draft', 'delete'): @@ -465,11 +492,11 @@ class MyCourseBatchView(APIView): ) qs = Course.objects.filter(id__in=ids, author=request.user) if action == 'delete': - count = qs.delete()[0] + count = (await qs.adelete())[0] elif action == 'publish': - count = qs.update(status='published') + count = await qs.aupdate(status='published') elif action == 'draft': - count = qs.update(status='draft') + count = await qs.aupdate(status='draft') return create_standardized_response( data={'affected': count}, code=ResponseCode.SUCCESS, @@ -493,21 +520,26 @@ class CDNStaticFileView(APIView): 404: not_found_response, }, ) - def get(self, request, course_id, chapter_id): + async def get(self, request, course_id, chapter_id): file_path = os.path.join( settings.MEDIA_ROOT, 'learn', 'courses', str(course_id), 'chapters', f'{chapter_id}.md' ) - if not os.path.exists(file_path): - raise Http404 - fh = open(file_path, "rb") - try: - response = FileResponse(fh, content_type="text/markdown; charset=utf-8") - response["Cache-Control"] = "max-age=86400" - return response - except Exception: - fh.close() - raise + + def _open_file(): + if not os.path.exists(file_path): + raise Http404 + fh = open(file_path, "rb") + try: + response = FileResponse(fh, content_type="text/markdown; charset=utf-8") + response["Cache-Control"] = "max-age=86400" + return response + except Exception: + fh.close() + raise + # TODO: aiohttp 化(当前用 sync_to_async 兜底避免阻塞事件循环) + return await sync_to_async(_open_file)() + class CourseFavoriteToggleView(APIView): permission_classes = [IsAuthenticated] @@ -525,22 +557,22 @@ class CourseFavoriteToggleView(APIView): 404: not_found_response, }, ) - def post(self, request, pk): - course = Course.objects.filter(pk=pk, status='published').first() + async def post(self, request, pk): + course = await Course.objects.filter(pk=pk, status='published').afirst() if not course: return create_standardized_error_response( code=ResponseCode.VALIDATION_ERROR, message='课程不存在', status_code=status.HTTP_404_NOT_FOUND ) - fav, created = CourseFavorite.objects.get_or_create(user=request.user, course=course) + fav, created = await CourseFavorite.objects.aget_or_create(user=request.user, course=course) if not created: - fav.delete() - count = CourseFavorite.objects.filter(course=course).count() + await fav.adelete() + count = await CourseFavorite.objects.filter(course=course).acount() return create_standardized_response( data={'favorited': False, 'favorites_count': count} ) - count = CourseFavorite.objects.filter(course=course).count() + count = await CourseFavorite.objects.filter(course=course).acount() return create_standardized_response( data={'favorited': True, 'favorites_count': count} ) @@ -568,8 +600,8 @@ class ChapterMarkCompletedView(APIView): 404: not_found_response, }, ) - def post(self, request, pk): - chapter = Chapter.objects.filter(pk=pk).select_related('course').first() + async def post(self, request, pk): + chapter = await Chapter.objects.filter(pk=pk).select_related('course').afirst() if not chapter: return create_standardized_error_response( code=ResponseCode.VALIDATION_ERROR, @@ -580,7 +612,7 @@ class ChapterMarkCompletedView(APIView): if not isinstance(completed, bool): completed = str(completed).lower() in ('true', '1', 'yes') - read_record, created = ChapterRead.objects.get_or_create( + read_record, created = await ChapterRead.objects.aget_or_create( user=request.user, chapter=chapter, defaults={ @@ -592,12 +624,12 @@ class ChapterMarkCompletedView(APIView): if not created: read_record.completed = completed read_record.completed_at = timezone.now() if completed else None - read_record.save(update_fields=['completed', 'completed_at', 'updated_at']) + await read_record.asave(update_fields=['completed', 'completed_at', 'updated_at']) - total_chapters = chapter.course.chapters.count() - completed_chapters = ChapterRead.objects.filter( + total_chapters = await chapter.course.chapters.acount() + completed_chapters = await ChapterRead.objects.filter( user=request.user, course=chapter.course, completed=True - ).count() + ).acount() progress = round((completed_chapters / total_chapters) * 100) if total_chapters > 0 else 0 return create_standardized_response( @@ -627,20 +659,20 @@ class CourseProgressView(APIView): 404: not_found_response, }, ) - def get(self, request, course_id): - course = Course.objects.filter(pk=course_id, status='published').first() + async def get(self, request, course_id): + course = await Course.objects.filter(pk=course_id, status='published').afirst() if not course: return create_standardized_error_response( code=ResponseCode.VALIDATION_ERROR, message='课程不存在', status_code=status.HTTP_404_NOT_FOUND ) - completed_ids = set( - ChapterRead.objects.filter( + completed_ids = set([ + v async for v in ChapterRead.objects.filter( user=request.user, course=course, completed=True ).values_list('chapter_id', flat=True) - ) - total_chapters = course.chapters.count() + ]) + total_chapters = await course.chapters.acount() completed_count = len(completed_ids) progress = round((completed_count / total_chapters) * 100) if total_chapters > 0 else 0 @@ -667,7 +699,7 @@ class MaterialListView(APIView): ], responses={200: success_response}, ) - def get(self, request): + async def get(self, request): chapter_id = request.query_params.get('chapter_id') if not chapter_id: return create_standardized_error_response( @@ -677,7 +709,8 @@ class MaterialListView(APIView): ) materials = CourseMaterial.objects.filter(chapter_id=chapter_id) serializer = MaterialSerializer(materials, many=True, context={'request': request}) - return create_standardized_response(data=serializer.data, code=ResponseCode.SUCCESS) + data = await get_data(serializer) + return create_standardized_response(data=data, code=ResponseCode.SUCCESS) class MaterialDownloadView(APIView): @@ -692,12 +725,12 @@ class MaterialDownloadView(APIView): ], responses={200: openapi.Response(description='文件流'), 404: not_found_response}, ) - def get(self, request, pk): - material = CourseMaterial.objects.filter(pk=pk).first() + async def get(self, request, pk): + material = await CourseMaterial.objects.filter(pk=pk).afirst() if not material or not material.file: raise Http404 material.download_count += 1 - material.save(update_fields=['download_count']) + await material.asave(update_fields=['download_count']) response = FileResponse(material.file.open('rb'), content_type='application/octet-stream') response['Content-Disposition'] = f'attachment; filename="{material.title}"' return response @@ -715,8 +748,8 @@ class ChapterVideoDownloadView(APIView): ], responses={200: openapi.Response(description='文件流'), 404: not_found_response}, ) - def get(self, request, pk): - chapter = Chapter.objects.filter(pk=pk).first() + async def get(self, request, pk): + chapter = await Chapter.objects.filter(pk=pk).afirst() if not chapter or not chapter.video_local: raise Http404 response = FileResponse(chapter.video_local.open('rb'), content_type='application/octet-stream') @@ -733,11 +766,11 @@ class MyProgressView(APIView): operation_description='获取当前用户所有课程的学习进度列表,需要登录', responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user reads = ChapterRead.objects.filter(user=user, completed=True).select_related('course', 'chapter') course_map = {} - for read in reads: + async for read in reads: course_id = read.course.id if course_id not in course_map: course_map[course_id] = { @@ -745,7 +778,7 @@ class MyProgressView(APIView): 'title': read.course.title, 'cover_image': read.course.cover_image.url if read.course.cover_image else '', 'completed_chapters': 0, - 'total_chapters': read.course.chapters.count(), + 'total_chapters': await read.course.chapters.acount(), 'last_studied_at': read.completed_at, } course_map[course_id]['completed_chapters'] += 1 @@ -772,4 +805,4 @@ class MyProgressView(APIView): 'total_courses': len(courses), 'total_completed': total_completed, } - ) \ No newline at end of file + ) diff --git a/logs/views.py b/logs/views.py index b71aeab..706ca76 100644 --- a/logs/views.py +++ b/logs/views.py @@ -1,12 +1,13 @@ from django.db.models import Count, Avg, Max, Q, F from django.db.models.functions import TruncHour, TruncDate from django.utils import timezone -from rest_framework import viewsets, status +from rest_framework import status from rest_framework.decorators import action from rest_framework.permissions import IsAuthenticated, IsAdminUser from rest_framework.response import Response from django_filters.rest_framework import DjangoFilterBackend from rest_framework.filters import SearchFilter, OrderingFilter +from adrf.viewsets import ModelViewSet, ReadOnlyModelViewSet, ViewSet from .models import ApiRequestLog, ErrorLog, SystemEventLog from .serializers import ( @@ -17,7 +18,7 @@ from .serializers import ( from .filters import ApiRequestLogFilter, ErrorLogFilter, SystemEventLogFilter -class ApiRequestLogViewSet(viewsets.ReadOnlyModelViewSet): +class ApiRequestLogViewSet(ReadOnlyModelViewSet): queryset = ApiRequestLog.objects.all() permission_classes = [IsAuthenticated, IsAdminUser] filter_backends = [DjangoFilterBackend, SearchFilter, OrderingFilter] @@ -32,7 +33,7 @@ class ApiRequestLogViewSet(viewsets.ReadOnlyModelViewSet): return ApiRequestLogSerializer -class ErrorLogViewSet(viewsets.ModelViewSet): +class ErrorLogViewSet(ModelViewSet): queryset = ErrorLog.objects.all() permission_classes = [IsAuthenticated, IsAdminUser] filter_backends = [DjangoFilterBackend, SearchFilter, OrderingFilter] @@ -54,18 +55,18 @@ class ErrorLogViewSet(viewsets.ModelViewSet): return ErrorLog.objects.all() @action(detail=True, methods=['post']) - def resolve(self, request, pk=None): - error_log = self.get_object() + async def resolve(self, request, pk=None): + error_log = await self.aget_object() error_log.is_resolved = True error_log.resolved_at = timezone.now() - error_log.save(update_fields=['is_resolved', 'resolved_at']) + await error_log.asave(update_fields=['is_resolved', 'resolved_at']) return Response({'status': 'resolved'}) @action(detail=False, methods=['post']) - def resolve_all(self, request): + async def resolve_all(self, request): ids = request.data.get('ids', []) if ids: - ErrorLog.objects.filter(id__in=ids, is_resolved=False).update( + await ErrorLog.objects.filter(id__in=ids, is_resolved=False).aupdate( is_resolved=True, resolved_at=timezone.now() ) @@ -73,7 +74,7 @@ class ErrorLogViewSet(viewsets.ModelViewSet): return Response({'error': '请提供 ids 列表'}, status=status.HTTP_400_BAD_REQUEST) -class SystemEventLogViewSet(viewsets.ReadOnlyModelViewSet): +class SystemEventLogViewSet(ReadOnlyModelViewSet): queryset = SystemEventLog.objects.select_related('user', 'target_user').all() permission_classes = [IsAuthenticated, IsAdminUser] filter_backends = [DjangoFilterBackend, SearchFilter, OrderingFilter] @@ -88,28 +89,28 @@ class SystemEventLogViewSet(viewsets.ReadOnlyModelViewSet): return SystemEventLogSerializer -class LogStatsViewSet(viewsets.ViewSet): +class LogStatsViewSet(ViewSet): permission_classes = [IsAuthenticated, IsAdminUser] @action(detail=False, methods=['get']) - def summary(self, request): + async def summary(self, request): days = int(request.query_params.get('days', 7)) start_date = timezone.now() - timezone.timedelta(days=days) request_qs = ApiRequestLog.objects.filter(timestamp__gte=start_date) error_qs = ErrorLog.objects.filter(timestamp__gte=start_date) - total_requests = request_qs.count() - total_errors = request_qs.filter(is_error=True).count() + total_requests = await request_qs.acount() + total_errors = await request_qs.filter(is_error=True).acount() error_rate = (total_errors / total_requests * 100) if total_requests > 0 else 0 - duration_stats = request_qs.aggregate( + duration_stats = await request_qs.aaggregate( avg_duration=Avg('duration_ms'), max_duration=Max('duration_ms') ) - unique_users = request_qs.exclude(user__isnull=True).values('user').distinct().count() - unique_ips = request_qs.exclude(ip_address__isnull=True).values('ip_address').distinct().count() + unique_users = await request_qs.exclude(user__isnull=True).values('user').distinct().acount() + unique_ips = await request_qs.exclude(ip_address__isnull=True).values('ip_address').distinct().acount() top_status_codes = ( request_qs.values('status_code') @@ -118,8 +119,8 @@ class LogStatsViewSet(viewsets.ViewSet): ) today_start = timezone.now().replace(hour=0, minute=0, second=0, microsecond=0) - error_count_today = error_qs.filter(timestamp__gte=today_start).count() - request_count_today = request_qs.filter(timestamp__gte=today_start).count() + error_count_today = await error_qs.filter(timestamp__gte=today_start).acount() + request_count_today = await request_qs.filter(timestamp__gte=today_start).acount() return Response({ 'total_requests': total_requests, @@ -129,13 +130,13 @@ class LogStatsViewSet(viewsets.ViewSet): 'max_duration_ms': round(duration_stats['max_duration'] or 0, 2), 'unique_users': unique_users, 'unique_ips': unique_ips, - 'top_status_codes': list(top_status_codes), + 'top_status_codes': [item async for item in top_status_codes], 'error_count_today': error_count_today, 'request_count_today': request_count_today, }) @action(detail=False, methods=['get']) - def by_path(self, request): + async def by_path(self, request): days = int(request.query_params.get('days', 7)) limit = int(request.query_params.get('limit', 20)) start_date = timezone.now() - timezone.timedelta(days=days) @@ -153,7 +154,7 @@ class LogStatsViewSet(viewsets.ViewSet): ) result = [] - for item in stats: + async for item in stats: item['error_rate'] = round( (item['error_count'] / item['request_count'] * 100) if item['request_count'] > 0 else 0, 2 ) @@ -163,7 +164,7 @@ class LogStatsViewSet(viewsets.ViewSet): return Response(result) @action(detail=False, methods=['get']) - def by_hour(self, request): + async def by_hour(self, request): days = int(request.query_params.get('days', 1)) start_date = timezone.now() - timezone.timedelta(days=days) @@ -181,14 +182,14 @@ class LogStatsViewSet(viewsets.ViewSet): ) result = [] - for item in stats: + async for item in stats: item['avg_duration_ms'] = round(item['avg_duration_ms'] or 0, 2) result.append(item) return Response(result) @action(detail=False, methods=['get']) - def slow_requests(self, request): + async def slow_requests(self, request): days = int(request.query_params.get('days', 7)) threshold = float(request.query_params.get('threshold', 1000)) limit = int(request.query_params.get('limit', 50)) @@ -207,7 +208,7 @@ class LogStatsViewSet(viewsets.ViewSet): ) result = [] - for item in slow_requests: + async for item in slow_requests: item['avg_duration_ms'] = round(item['avg_duration_ms'] or 0, 2) item['max_duration_ms'] = round(item['max_duration_ms'] or 0, 2) result.append(item) @@ -215,7 +216,7 @@ class LogStatsViewSet(viewsets.ViewSet): return Response(result) @action(detail=False, methods=['get']) - def error_types(self, request): + async def error_types(self, request): days = int(request.query_params.get('days', 7)) limit = int(request.query_params.get('limit', 20)) start_date = timezone.now() - timezone.timedelta(days=days) @@ -231,4 +232,4 @@ class LogStatsViewSet(viewsets.ViewSet): .order_by('-count')[:limit] ) - return Response(list(error_types)) + return Response([item async for item in error_types]) diff --git a/message/serializers.py b/message/serializers.py index 168cd35..a59c350 100644 --- a/message/serializers.py +++ b/message/serializers.py @@ -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 diff --git a/message/views.py b/message/views.py index 444cf1a..88318f4 100644 --- a/message/views.py +++ b/message/views.py @@ -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) diff --git a/requirements.txt b/requirements.txt index 8d91d04..2322c0b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -102,3 +102,9 @@ xmltodict==1.0.4 zope.interface==8.2 gunicorn==21.2.0 pillow-avif-plugin>=1.3.0 +# ============================================ +# 异步化架构升级:ADRF + Granian + PostgreSQL(psycopg3) +# ============================================ +adrf==0.1.14 +granian==2.8.2 +psycopg[binary]==3.2.13 diff --git a/search/views.py b/search/views.py index 22c6324..3900431 100644 --- a/search/views.py +++ b/search/views.py @@ -1,4 +1,4 @@ -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from rest_framework.permissions import AllowAny @@ -43,7 +43,7 @@ class GlobalSearchView(APIView): ], responses={200: success_response} ) - def get(self, request): + async def get(self, request): q = request.query_params.get('q', '').strip() content_type = request.query_params.get('type', 'all') page = int(request.query_params.get('page', 1)) @@ -65,7 +65,7 @@ class GlobalSearchView(APIView): Q(title__icontains=q) | Q(excerpt__icontains=q) | Q(content__icontains=q) ).filter(status='published').order_by('-created_at')[:100] - for article in articles: + async for article in articles: cover_url = '' if article.cover_image and hasattr(article.cover_image, 'url'): cover_url = request.build_absolute_uri(article.cover_image.url) @@ -90,7 +90,7 @@ class GlobalSearchView(APIView): Q(name__icontains=q) | Q(description__icontains=q) ).order_by('-created_at')[:100] - for tool in tools: + async for tool in tools: results.append({ 'id': tool.id, 'type': '工具', @@ -106,7 +106,7 @@ class GlobalSearchView(APIView): Q(title__icontains=q) | Q(description__icontains=q) ).order_by('-created_at')[:100] - for course in courses: + async for course in courses: cover_url = '' if course.cover_image and hasattr(course.cover_image, 'url'): cover_url = request.build_absolute_uri(course.cover_image.url) @@ -126,7 +126,7 @@ class GlobalSearchView(APIView): Q(name__icontains=q) | Q(description__icontains=q) | Q(url_path__icontains=q) ).filter(is_enabled=True).order_by('-created_at')[:100] - for api in apis: + async for api in apis: desc = api.description or f'{api.method} {api.url_path}' results.append({ 'id': api.id, @@ -175,7 +175,7 @@ class SearchSuggestionsView(APIView): ], responses={200: success_response} ) - def get(self, request): + async def get(self, request): q = request.query_params.get('q', '').strip() limit = int(request.query_params.get('limit', 8)) @@ -187,22 +187,22 @@ class SearchSuggestionsView(APIView): article_titles = Article.objects.filter( title__icontains=q ).filter(status='published').values_list('title', flat=True)[:5] - suggestions.update(article_titles) + suggestions.update([t async for t in article_titles]) tool_names = Tool.objects.filter( name__icontains=q ).values_list('name', flat=True)[:5] - suggestions.update(tool_names) + suggestions.update([t async for t in tool_names]) course_titles = Course.objects.filter( title__icontains=q ).values_list('title', flat=True)[:5] - suggestions.update(course_titles) + suggestions.update([t async for t in course_titles]) api_names = ApiItem.objects.filter( name__icontains=q ).filter(is_enabled=True).values_list('name', flat=True)[:5] - suggestions.update(api_names) + suggestions.update([t async for t in api_names]) result = list(suggestions)[:limit] return create_standardized_response(data=result, code=ResponseCode.SUCCESS) @@ -217,7 +217,7 @@ class HotKeywordsView(APIView): operation_description='返回热门搜索关键词列表', responses={200: success_response} ) - def get(self, request): + async def get(self, request): hot_keywords = [ 'React Hooks', 'RESTful API', @@ -230,4 +230,4 @@ class HotKeywordsView(APIView): 'Vue3', 'Python', ] - return create_standardized_response(data=hot_keywords, code=ResponseCode.SUCCESS) \ No newline at end of file + return create_standardized_response(data=hot_keywords, code=ResponseCode.SUCCESS) diff --git a/shorturl/views.py b/shorturl/views.py index b74eff2..006b079 100644 --- a/shorturl/views.py +++ b/shorturl/views.py @@ -3,7 +3,7 @@ import re from django.http import HttpResponseRedirect, HttpResponseNotFound, HttpResponseGone from django.utils import timezone from rest_framework.permissions import AllowAny, IsAuthenticated -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from drf_yasg.utils import swagger_auto_schema @@ -54,7 +54,7 @@ class ShortUrlShortenView(APIView): ), responses={200: success_response, 400: error_response, 409: error_response} ) - def post(self, request): + async def post(self, request): url = request.data.get('url', '').strip() custom_code = request.data.get('custom_code', '').strip() or None expire_days = request.data.get('expire_days') @@ -71,17 +71,17 @@ class ShortUrlShortenView(APIView): {"code": 400, "message": "自定义短码仅允许字母、数字、连字符、下划线,长度3-20字符"}, status=status.HTTP_400_BAD_REQUEST ) - if ShortUrl.objects.filter(code=custom_code).exists(): + if await ShortUrl.objects.filter(code=custom_code).aexists(): return Response( {"code": 409, "message": "该短码已被使用,请更换"}, status=status.HTTP_409_CONFLICT ) code = custom_code else: - last = ShortUrl.objects.order_by('-id').first() + last = await ShortUrl.objects.order_by('-id').afirst() next_id = (last.id + 1) if last else 1 code = encode_base62(next_id) - while ShortUrl.objects.filter(code=code).exists(): + while await ShortUrl.objects.filter(code=code).aexists(): next_id += 1 code = encode_base62(next_id) @@ -96,7 +96,7 @@ class ShortUrlShortenView(APIView): user = request.user if request.user.is_authenticated else None - short_url = ShortUrl.objects.create( + short_url = await ShortUrl.objects.acreate( code=code, original_url=url, custom_code=custom_code, @@ -129,9 +129,9 @@ class ShortUrlInfoView(APIView): operation_description='通过短码查询短链接的详细信息', responses={200: success_response, 404: error_response} ) - def get(self, request, code): + async def get(self, request, code): try: - short_url = ShortUrl.objects.get(code=code) + short_url = await ShortUrl.objects.aget(code=code) except ShortUrl.DoesNotExist: return Response( {"code": 404, "message": "短链接不存在"}, @@ -157,16 +157,16 @@ class ShortUrlInfoView(APIView): class ShortUrlRedirectView(APIView): permission_classes = [AllowAny] - def get(self, request, code): + async def get(self, request, code): try: - short_url = ShortUrl.objects.get(code=code) + short_url = await ShortUrl.objects.aget(code=code) except ShortUrl.DoesNotExist: return HttpResponseNotFound('

404 - 短链接不存在

') if short_url.expire_at and short_url.expire_at < timezone.now(): return HttpResponseGone('

410 - 短链接已过期

') - ShortUrl.objects.filter(pk=short_url.pk).update(click_count=short_url.click_count + 1) + await ShortUrl.objects.filter(pk=short_url.pk).aupdate(click_count=short_url.click_count + 1) return HttpResponseRedirect(short_url.original_url) @@ -180,7 +180,7 @@ class ShortUrlListView(APIView): operation_description='分页返回当前登录用户创建的所有短链接', responses={200: success_response, 401: unauthorized_response} ) - def get(self, request): + async def get(self, request): queryset = ShortUrl.objects.filter(creator=request.user).order_by('-created_at') page = int(request.GET.get('page', 1)) @@ -189,8 +189,8 @@ class ShortUrlListView(APIView): start = (page - 1) * page_size end = start + page_size - total = queryset.count() - results = queryset[start:end] + total = await queryset.acount() + results = [item async for item in queryset[start:end]] data = [{ "code": item.code, diff --git a/tool/serializers.py b/tool/serializers.py index 3b60a0b..0ae8217 100644 --- a/tool/serializers.py +++ b/tool/serializers.py @@ -1,20 +1,23 @@ -from rest_framework import serializers +from adrf.serializers import ( + ModelSerializer, Serializer, CharField, ChoiceField, FileField, + IntegerField, BooleanField, ListField, SerializerMethodField, +) from .models import CompressionHistory, ToolCategory, Tool, ColorHistory -class ToolCategorySerializer(serializers.ModelSerializer): - tool_count = serializers.SerializerMethodField() +class ToolCategorySerializer(ModelSerializer): + tool_count = SerializerMethodField() class Meta: model = ToolCategory fields = ['id', 'name', 'icon', 'sort_order', 'tool_count', 'created_at'] - def get_tool_count(self, obj): - return obj.tools.filter(is_enabled=True).count() + async def get_tool_count(self, obj): + return await obj.tools.filter(is_enabled=True).acount() -class ToolSerializer(serializers.ModelSerializer): - category_name = serializers.CharField(source='category.name', read_only=True, default='') +class ToolSerializer(ModelSerializer): + category_name = CharField(source='category.name', read_only=True, default='') class Meta: model = Tool @@ -26,10 +29,10 @@ class ToolSerializer(serializers.ModelSerializer): read_only_fields = ['id', 'created_at', 'updated_at'] -class CompressionHistorySerializer(serializers.ModelSerializer): - original_size_display = serializers.SerializerMethodField() - compressed_size_display = serializers.SerializerMethodField() - created_at_display = serializers.SerializerMethodField() +class CompressionHistorySerializer(ModelSerializer): + original_size_display = SerializerMethodField() + compressed_size_display = SerializerMethodField() + created_at_display = SerializerMethodField() class Meta: model = CompressionHistory @@ -41,13 +44,13 @@ class CompressionHistorySerializer(serializers.ModelSerializer): ] read_only_fields = ['id', 'created_at'] - def get_original_size_display(self, obj): + async def get_original_size_display(self, obj): return self._format_size(obj.original_size) - def get_compressed_size_display(self, obj): + async def get_compressed_size_display(self, obj): return self._format_size(obj.compressed_size) - def get_created_at_display(self, obj): + async def get_created_at_display(self, obj): return obj.created_at.strftime('%Y-%m-%d %H:%M:%S') @staticmethod @@ -60,29 +63,29 @@ class CompressionHistorySerializer(serializers.ModelSerializer): return f"{size_bytes / (1024 * 1024):.1f} MB" -class ImageCompressRequestSerializer(serializers.Serializer): - file = serializers.FileField() - mode = serializers.ChoiceField(choices=['lossy', 'lossless'], default='lossy') - quality = serializers.IntegerField(min_value=1, max_value=100, default=80) - format = serializers.ChoiceField( +class ImageCompressRequestSerializer(Serializer): + file = FileField() + mode = ChoiceField(choices=['lossy', 'lossless'], default='lossy') + quality = IntegerField(min_value=1, max_value=100, default=80) + format = ChoiceField( choices=['jpeg', 'png', 'webp', 'avif', 'original'], default='original' ) - keep_exif = serializers.BooleanField(default=True) - width = serializers.IntegerField(required=False, min_value=1, max_value=10000) - height = serializers.IntegerField(required=False, min_value=1, max_value=10000) - maintain_aspect_ratio = serializers.BooleanField(default=True) + keep_exif = BooleanField(default=True) + width = IntegerField(required=False, min_value=1, max_value=10000) + height = IntegerField(required=False, min_value=1, max_value=10000) + maintain_aspect_ratio = BooleanField(default=True) -class ColorHistorySerializer(serializers.ModelSerializer): +class ColorHistorySerializer(ModelSerializer): class Meta: model = ColorHistory fields = ['id', 'color', 'created_at'] read_only_fields = ['id', 'created_at'] -class ColorHistorySyncSerializer(serializers.Serializer): - colors = serializers.ListField( - child=serializers.CharField(max_length=7), +class ColorHistorySyncSerializer(Serializer): + colors = ListField( + child=CharField(max_length=7), help_text='颜色HEX值列表,如 ["#FF5733", "#33FF57"]' ) diff --git a/tool/views/color_history_view.py b/tool/views/color_history_view.py index 8bf85d1..b1b14e7 100644 --- a/tool/views/color_history_view.py +++ b/tool/views/color_history_view.py @@ -1,5 +1,5 @@ import re -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from rest_framework.permissions import IsAuthenticated @@ -17,17 +17,17 @@ def validate_hex_color(color): class ColorHistoryListView(APIView): permission_classes = [IsAuthenticated] - def get(self, request): + async def get(self, request): days = int(request.query_params.get('days', 30)) cutoff = timezone.now() - timedelta(days=days) - histories = ColorHistory.objects.filter( + histories = [h async for h in ColorHistory.objects.filter( user=request.user, created_at__gte=cutoff - ).order_by('-created_at')[:50] + ).order_by('-created_at')[:50]] serializer = ColorHistorySerializer(histories, many=True) return Response(serializer.data) - def post(self, request): + async def post(self, request): color = request.data.get('color', '') if not validate_hex_color(color): return Response( @@ -35,16 +35,16 @@ class ColorHistoryListView(APIView): status=status.HTTP_400_BAD_REQUEST ) color = color.upper() - existing = ColorHistory.objects.filter( + existing = await ColorHistory.objects.filter( user=request.user, color=color - ).first() + ).afirst() if existing: existing.created_at = timezone.now() - existing.save() + await existing.asave() serializer = ColorHistorySerializer(existing) return Response(serializer.data, status=status.HTTP_200_OK) - history = ColorHistory.objects.create( + history = await ColorHistory.objects.acreate( user=request.user, color=color ) @@ -55,10 +55,10 @@ class ColorHistoryListView(APIView): class ColorHistoryDetailView(APIView): permission_classes = [IsAuthenticated] - def delete(self, request, pk): + async def delete(self, request, pk): try: - history = ColorHistory.objects.get(pk=pk, user=request.user) - history.delete() + history = await ColorHistory.objects.aget(pk=pk, user=request.user) + await history.adelete() return Response(status=status.HTTP_204_NO_CONTENT) except ColorHistory.DoesNotExist: return Response( @@ -70,34 +70,32 @@ class ColorHistoryDetailView(APIView): class ColorHistoryClearView(APIView): permission_classes = [IsAuthenticated] - def post(self, request): - deleted_count, _ = ColorHistory.objects.filter(user=request.user).delete() + async def post(self, request): + deleted_count, _ = await ColorHistory.objects.filter(user=request.user).adelete() return Response({'deleted': deleted_count}) class ColorHistorySyncView(APIView): permission_classes = [IsAuthenticated] - def post(self, request): + async def post(self, request): serializer = ColorHistorySyncSerializer(data=request.data) if not serializer.is_valid(): return Response(serializer.errors, status=status.HTTP_400_BAD_REQUEST) colors = serializer.validated_data['colors'] valid_colors = [c.upper() for c in colors if validate_hex_color(c)] - existing_colors = set( - ColorHistory.objects.filter( + existing_colors = set([ + c async for c in ColorHistory.objects.filter( user=request.user, color__in=valid_colors ).values_list('color', flat=True) - ) + ]) new_colors = [c for c in valid_colors if c not in existing_colors] now = timezone.now() - for color in new_colors: - ColorHistory.objects.create( - user=request.user, - color=color, - created_at=now - ) + await ColorHistory.objects.abulk_create([ + ColorHistory(user=request.user, color=color, created_at=now) + for color in new_colors + ]) return Response({ 'synced': len(new_colors), 'duplicates': len(valid_colors) - len(new_colors), diff --git a/tool/views/compression_history_view.py b/tool/views/compression_history_view.py index 7f24288..30b3500 100644 --- a/tool/views/compression_history_view.py +++ b/tool/views/compression_history_view.py @@ -1,15 +1,16 @@ from django.http import JsonResponse -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework import status from rest_framework.decorators import permission_classes from rest_framework.permissions import AllowAny +from asgiref.sync import sync_to_async from ..models import CompressionHistory from ..serializers import CompressionHistorySerializer @permission_classes([AllowAny]) class CompressionHistoryListView(APIView): - def get(self, request): + async def get(self, request): user = request.user if request.user.is_authenticated else None session_id = request.query_params.get('session_id') @@ -29,8 +30,8 @@ class CompressionHistoryListView(APIView): start = (page - 1) * page_size end = start + page_size - total = histories.count() - serializer = CompressionHistorySerializer(histories[start:end], many=True) + total = await histories.acount() + serializer = CompressionHistorySerializer([h async for h in histories[start:end]], many=True) return JsonResponse({ 'success': True, @@ -40,7 +41,7 @@ class CompressionHistoryListView(APIView): 'page_size': page_size }) - def post(self, request): + async def post(self, request): data = request.data.copy() if request.user.is_authenticated: @@ -48,7 +49,8 @@ class CompressionHistoryListView(APIView): serializer = CompressionHistorySerializer(data=data) if serializer.is_valid(): - serializer.save() + # DRF 序列化器 save() 为同步 ORM 操作,sync_to_async 兜底 + await sync_to_async(serializer.save)() return JsonResponse({ 'success': True, 'data': serializer.data @@ -62,22 +64,22 @@ class CompressionHistoryListView(APIView): @permission_classes([AllowAny]) class CompressionHistoryDetailView(APIView): - def delete(self, request, pk): + async def delete(self, request, pk): user = request.user if request.user.is_authenticated else None session_id = request.query_params.get('session_id') try: if user: - history = CompressionHistory.objects.get(pk=pk, user=user) + history = await CompressionHistory.objects.aget(pk=pk, user=user) elif session_id: - history = CompressionHistory.objects.get(pk=pk, session_id=session_id) + history = await CompressionHistory.objects.aget(pk=pk, session_id=session_id) else: return JsonResponse( {'error': '无权限删除'}, status=status.HTTP_403_FORBIDDEN ) - history.delete() + await history.adelete() return JsonResponse({'success': True}) except CompressionHistory.DoesNotExist: return JsonResponse( @@ -88,14 +90,14 @@ class CompressionHistoryDetailView(APIView): @permission_classes([AllowAny]) class CompressionHistoryClearView(APIView): - def delete(self, request): + async def delete(self, request): user = request.user if request.user.is_authenticated else None session_id = request.query_params.get('session_id') if user: - count, _ = CompressionHistory.objects.filter(user=user).delete() + count, _ = await CompressionHistory.objects.filter(user=user).adelete() elif session_id: - count, _ = CompressionHistory.objects.filter(session_id=session_id).delete() + count, _ = await CompressionHistory.objects.filter(session_id=session_id).adelete() else: return JsonResponse( {'error': '无权限清空'}, diff --git a/tool/views/image_compress_view.py b/tool/views/image_compress_view.py index 71d164b..26b080d 100644 --- a/tool/views/image_compress_view.py +++ b/tool/views/image_compress_view.py @@ -3,7 +3,7 @@ import uuid from PIL import Image from PIL.ExifTags import TAGS from django.http import JsonResponse, FileResponse -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.parsers import MultiPartParser, FormParser from rest_framework import status from rest_framework.decorators import permission_classes diff --git a/tool/views/text_diff_view.py b/tool/views/text_diff_view.py index c3b2443..104c344 100644 --- a/tool/views/text_diff_view.py +++ b/tool/views/text_diff_view.py @@ -1,6 +1,6 @@ import difflib from rest_framework.decorators import permission_classes -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from rest_framework.permissions import AllowAny diff --git a/tool/views/tool_favorite_view.py b/tool/views/tool_favorite_view.py index 2382459..6b76a08 100644 --- a/tool/views/tool_favorite_view.py +++ b/tool/views/tool_favorite_view.py @@ -1,8 +1,8 @@ from django.http import JsonResponse -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.permissions import IsAuthenticated from rest_framework.parsers import JSONParser -from django.db import transaction +from asgiref.sync import sync_to_async from ..models import Tool, ToolFavorite @@ -10,60 +10,75 @@ class ToolFavoriteToggleView(APIView): permission_classes = [IsAuthenticated] parser_classes = [JSONParser] - @transaction.atomic - def post(self, request, tool_id): - try: - tool = Tool.objects.get(pk=tool_id, is_enabled=True) - except Tool.DoesNotExist: + async def post(self, request, tool_id): + # 事务段(get_or_create + delete + count)整体放入同步函数,sync_to_async 执行保证原子性 + def _toggle_in_txn(): + from django.db import transaction + with transaction.atomic(): + try: + tool = Tool.objects.get(pk=tool_id, is_enabled=True) + except Tool.DoesNotExist: + return None + favorite, created = ToolFavorite.objects.get_or_create( + user=request.user, + tool=tool + ) + if not created: + favorite.delete() + is_favorite = False + message = '已取消收藏' + else: + is_favorite = True + message = '已添加收藏' + favorite_count = ToolFavorite.objects.filter(tool=tool).count() + return { + 'tool_id': tool.id, + 'is_favorite': is_favorite, + 'favorite_count': favorite_count, + 'message': message, + } + + result = await sync_to_async(_toggle_in_txn)() + if result is None: return JsonResponse({'success': False, 'error': '工具不存在'}, status=404) - - favorite, created = ToolFavorite.objects.get_or_create( - user=request.user, - tool=tool - ) - if not created: - favorite.delete() - is_favorite = False - message = '已取消收藏' - else: - is_favorite = True - message = '已添加收藏' - - favorite_count = ToolFavorite.objects.filter(tool=tool).count() - return JsonResponse({ 'success': True, 'data': { - 'id': tool.id, - 'is_favorite': is_favorite, - 'favorite_count': favorite_count + 'id': result['tool_id'], + 'is_favorite': result['is_favorite'], + 'favorite_count': result['favorite_count'] }, - 'message': message + 'message': result['message'] }) class ToolFavoriteListView(APIView): permission_classes = [IsAuthenticated] - def get(self, request): - favorites = ToolFavorite.objects.select_related('tool', 'tool__category').filter( + async def get(self, request): + favorites = [fav async for fav in ToolFavorite.objects.select_related( + 'tool', 'tool__category' + ).filter( user=request.user, tool__is_enabled=True - ).order_by('-created_at') + ).order_by('-created_at')] - data = [{ - 'id': fav.tool.id, - 'name': fav.tool.name, - 'description': fav.tool.description, - 'icon': fav.tool.icon, - 'url_path': fav.tool.url_path, - 'color': fav.tool.color, - 'category': fav.tool.category.id if fav.tool.category else None, - 'category_name': fav.tool.category.name if fav.tool.category else '', - 'is_favorite': True, - 'favorite_count': ToolFavorite.objects.filter(tool=fav.tool).count(), - 'created_at': fav.created_at.isoformat() - } for fav in favorites] + data = [] + for fav in favorites: + favorite_count = await ToolFavorite.objects.filter(tool=fav.tool).acount() + data.append({ + 'id': fav.tool.id, + 'name': fav.tool.name, + 'description': fav.tool.description, + 'icon': fav.tool.icon, + 'url_path': fav.tool.url_path, + 'color': fav.tool.color, + 'category': fav.tool.category.id if fav.tool.category else None, + 'category_name': fav.tool.category.name if fav.tool.category else '', + 'is_favorite': True, + 'favorite_count': favorite_count, + 'created_at': fav.created_at.isoformat() + }) return JsonResponse({ 'success': True, @@ -77,13 +92,13 @@ class ToolFavoriteListView(APIView): class ToolFavoriteStatusView(APIView): permission_classes = [IsAuthenticated] - def get(self, request, tool_id): - is_favorite = ToolFavorite.objects.filter( + async def get(self, request, tool_id): + is_favorite = await ToolFavorite.objects.filter( user=request.user, tool_id=tool_id - ).exists() + ).aexists() - favorite_count = ToolFavorite.objects.filter(tool_id=tool_id).count() + favorite_count = await ToolFavorite.objects.filter(tool_id=tool_id).acount() return JsonResponse({ 'success': True, diff --git a/tool/views/tool_manage_view.py b/tool/views/tool_manage_view.py index b1dbd1d..33458ce 100644 --- a/tool/views/tool_manage_view.py +++ b/tool/views/tool_manage_view.py @@ -1,5 +1,5 @@ from django.http import JsonResponse -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework import status from rest_framework.permissions import AllowAny, IsAuthenticated from django.db.models import Count, Q @@ -19,8 +19,8 @@ class ToolCategoryListView(APIView): tags=['工具'], responses={200: success_response} ) - def get(self, request): - categories = ToolCategory.objects.prefetch_related('tools').all() + async def get(self, request): + categories = [c async for c in ToolCategory.objects.prefetch_related('tools').all()] serializer = ToolCategorySerializer(categories, many=True) return JsonResponse({ 'success': True, @@ -42,7 +42,7 @@ class ToolListView(APIView): ], responses={200: success_response} ) - def get(self, request): + async def get(self, request): category_id = request.query_params.get('category_id') enabled_only = request.query_params.get('enabled_only', 'true') ordering = request.query_params.get('ordering', '') @@ -68,18 +68,19 @@ class ToolListView(APIView): else: tools = tools.order_by('sort_order', '-created_at') - serializer = ToolSerializer(tools, many=True) + tool_list = [t async for t in tools] + serializer = ToolSerializer(tool_list, many=True) data = serializer.data # 构建 favorites_count 映射 favorites_count_map = { - t.id: t.annotated_favorites_count for t in tools + t.id: t.annotated_favorites_count for t in tool_list } if request.user.is_authenticated: - favorite_tool_ids = set( - ToolFavorite.objects.filter(user=request.user).values_list('tool_id', flat=True) - ) + favorite_tool_ids = set([ + tid async for tid in ToolFavorite.objects.filter(user=request.user).values_list('tool_id', flat=True) + ]) for item in data: item['is_favorited'] = item['id'] in favorite_tool_ids item['favorites_count'] = favorites_count_map.get(item['id'], 0) @@ -106,18 +107,18 @@ class ToolDetailView(APIView): ], responses={200: success_response, 404: not_found_response} ) - def get(self, request, pk): + async def get(self, request, pk): try: - tool = Tool.objects.select_related('category').get(pk=pk, is_enabled=True) + tool = await Tool.objects.select_related('category').aget(pk=pk, is_enabled=True) serializer = ToolSerializer(tool) data = serializer.data if request.user.is_authenticated: - data['is_favorited'] = ToolFavorite.objects.filter( + data['is_favorited'] = await ToolFavorite.objects.filter( user=request.user, tool=tool - ).exists() + ).aexists() else: data['is_favorited'] = False - data['favorites_count'] = ToolFavorite.objects.filter(tool=tool).count() + data['favorites_count'] = await ToolFavorite.objects.filter(tool=tool).acount() return JsonResponse({ 'success': True, 'data': data @@ -141,20 +142,20 @@ class ToolFavoriteToggleView(APIView): ], responses={200: success_response, 401: unauthorized_response, 404: not_found_response} ) - def post(self, request, pk): - tool = Tool.objects.filter(pk=pk, is_enabled=True).first() + async def post(self, request, pk): + tool = await Tool.objects.filter(pk=pk, is_enabled=True).afirst() if not tool: return JsonResponse( {'success': False, 'error': '工具不存在'}, status=status.HTTP_404_NOT_FOUND ) - fav, created = ToolFavorite.objects.get_or_create(user=request.user, tool=tool) + fav, created = await ToolFavorite.objects.aget_or_create(user=request.user, tool=tool) if not created: - fav.delete() + await fav.adelete() is_favorited = False else: is_favorited = True - favorites_count = ToolFavorite.objects.filter(tool=tool).count() + favorites_count = await ToolFavorite.objects.filter(tool=tool).acount() return JsonResponse({ 'success': True, 'data': { @@ -174,17 +175,19 @@ class ToolFavoriteListView(APIView): tags=['工具'], responses={200: success_response, 401: unauthorized_response} ) - def get(self, request): - favorites = ToolFavorite.objects.select_related('tool', 'tool__category').filter( + async def get(self, request): + favorites = [fav async for fav in ToolFavorite.objects.select_related( + 'tool', 'tool__category' + ).filter( user=request.user, tool__is_enabled=True - ).order_by('-created_at') + ).order_by('-created_at')] result = [] for fav in favorites: tool_data = ToolSerializer(fav.tool).data tool_data['is_favorited'] = True - tool_data['favorites_count'] = ToolFavorite.objects.filter(tool=fav.tool).count() + tool_data['favorites_count'] = await ToolFavorite.objects.filter(tool=fav.tool).acount() result.append(tool_data) return JsonResponse({ @@ -206,13 +209,13 @@ class ToolUsageIncrementView(APIView): ], responses={200: success_response, 404: not_found_response} ) - def post(self, request, pk): + async def post(self, request, pk): try: - tool = Tool.objects.get(pk=pk, is_enabled=True) + tool = await Tool.objects.aget(pk=pk, is_enabled=True) # 使用F()表达式避免竞态条件 from django.db.models import F - Tool.objects.filter(pk=pk).update(usage_count=F('usage_count') + 1) - tool.refresh_from_db() + await Tool.objects.filter(pk=pk).aupdate(usage_count=F('usage_count') + 1) + await tool.arefresh_from_db() return JsonResponse({ 'success': True, 'data': { diff --git a/user/serializers/region_serializers.py b/user/serializers/region_serializers.py index 4735e73..28ba819 100644 --- a/user/serializers/region_serializers.py +++ b/user/serializers/region_serializers.py @@ -1,22 +1,23 @@ -from rest_framework import serializers +from adrf.serializers import ModelSerializer, SerializerMethodField from ..models import Region -class RegionSerializer(serializers.ModelSerializer): - children = serializers.SerializerMethodField() +class RegionSerializer(ModelSerializer): + children = SerializerMethodField() class Meta: model = Region fields = ['id', 'name', 'code', 'level', 'parent_id', 'pinyin', 'children'] - def get_children(self, obj): - children = obj.children.all() - if children.exists(): - return RegionSerializer(children, many=True).data + async def get_children(self, obj): + children_qs = obj.children.all() + if await children_qs.aexists(): + # 递归异步序列化 + return [await RegionSerializer(c, context=self.context).adata async for c in children_qs] return None -class RegionSimpleSerializer(serializers.ModelSerializer): +class RegionSimpleSerializer(ModelSerializer): class Meta: model = Region fields = ['id', 'name', 'code', 'level', 'parent_id', 'pinyin'] diff --git a/user/serializers/user_serializers.py b/user/serializers/user_serializers.py index f41c619..104e130 100644 --- a/user/serializers/user_serializers.py +++ b/user/serializers/user_serializers.py @@ -1,14 +1,17 @@ -from rest_framework import serializers +from adrf.serializers import ( + ModelSerializer, Serializer, CharField, IntegerField, EmailField, + ChoiceField, IPAddressField, DateTimeField, ImageField, +) +from rest_framework.exceptions import ValidationError from django.contrib.auth import get_user_model from django.utils import timezone -from django.core.cache import caches +from utils.async_cache import aget_cache, aset_cache, adelete_cache import os import re from utils import RandCode FUser = get_user_model() -default_cache = caches['default'] def get_fields_to_extract(validated_data, fields_to_extract): @@ -18,7 +21,7 @@ def get_fields_to_extract(validated_data, fields_to_extract): extracted[field] = validated_data[field] return extracted -class UserSerializer(serializers.ModelSerializer): +class UserSerializer(ModelSerializer): PROTECTED_FIELDS = ['is_active', 'password', 'email', 'is_staff', 'is_superuser', 'last_login'] class Meta: model = FUser @@ -56,29 +59,30 @@ class UserSerializer(serializers.ModelSerializer): } - def create_by_email(self, validated_data): + async def acreate_by_email(self, validated_data): fields_to_extract = ['email'] extracted = get_fields_to_extract(validated_data, fields_to_extract) extracted['username'] = extracted['email'] extracted['password'] = RandCode.get_alphanumeric_characters_code_8() - # 2. 创建用户 - user = FUser.objects.create_user(**extracted) + # 2. 创建用户(create_user 内含密码哈希,为同步调用,sync_to_async 兜底) + from asgiref.sync import sync_to_async + user = await sync_to_async(FUser.objects.create_user)(**extracted) user.is_active = True user.is_staff = False user.is_superuser = False - user.save() + await user.asave() return user - def update(self, instance, validated_data): + async def aupdate(self, instance, validated_data): for attr, value in validated_data.items(): if attr in self.PROTECTED_FIELDS : continue setattr(instance, attr, value) - instance.save() + await instance.asave() return instance def __init__(self, *args, **kwargs): @@ -89,42 +93,45 @@ class UserSerializer(serializers.ModelSerializer): field.required = False -class ChangePasswordSerializer(serializers.Serializer): - old_password = serializers.CharField(required=False, allow_blank=True) - new_password = serializers.CharField(required=True, min_length=8, max_length=128) - confirm_password = serializers.CharField(required=True, min_length=8, max_length=128) +class ChangePasswordSerializer(Serializer): + old_password = CharField(required=False, allow_blank=True) + new_password = CharField(required=True, min_length=8, max_length=128) + confirm_password = CharField(required=True, min_length=8, max_length=128) def validate_old_password(self, value): user = self.context['request'].user + # check_password 为 CPU 密集(argon2),同步调用(is_valid 由视图 sync_to_async 包裹) if user.has_usable_password() and not user.check_password(value): - raise serializers.ValidationError('当前密码错误') + raise ValidationError('当前密码错误') return value def validate_new_password(self, value): if len(value) < 8: - raise serializers.ValidationError('密码长度不能少于8位') + raise ValidationError('密码长度不能少于8位') has_letter = any(c.isalpha() for c in value) has_digit = any(c.isdigit() for c in value) if not (has_letter and has_digit): - raise serializers.ValidationError('密码必须包含字母和数字') + raise ValidationError('密码必须包含字母和数字') return value def validate(self, attrs): if attrs['new_password'] != attrs['confirm_password']: - raise serializers.ValidationError({'confirm_password': '两次输入的新密码不一致'}) + raise ValidationError({'confirm_password': '两次输入的新密码不一致'}) return attrs - def save(self): + async def asave(self): + from asgiref.sync import sync_to_async user = self.context['request'].user - user.set_password(self.validated_data['new_password']) + # 密码哈希 CPU 密集,sync_to_async 兜底 + await sync_to_async(user.set_password)(self.validated_data['new_password']) if not user.isSetPassword: user.isSetPassword = True - user.save() + await user.asave() return user -class UserUpdateSerializer(serializers.ModelSerializer): - username = serializers.CharField( +class UserUpdateSerializer(ModelSerializer): + username = CharField( required=False, max_length=150, min_length=2, @@ -133,14 +140,14 @@ class UserUpdateSerializer(serializers.ModelSerializer): 'min_length': '用户名至少2个字符', } ) - gender = serializers.ChoiceField( + gender = ChoiceField( required=False, choices=[(0, '保密'), (1, '男'), (2, '女')], error_messages={ 'invalid_choice': '性别值无效,可选值为:0(保密)、1(男)、2(女)', } ) - bio = serializers.CharField( + bio = CharField( required=False, max_length=500, allow_blank=True, @@ -148,7 +155,7 @@ class UserUpdateSerializer(serializers.ModelSerializer): 'max_length': '个人简介不能超过500个字符', } ) - location = serializers.CharField( + location = CharField( required=False, max_length=100, allow_blank=True, @@ -156,7 +163,7 @@ class UserUpdateSerializer(serializers.ModelSerializer): 'max_length': '所在地区不能超过100个字符', } ) - phone_number = serializers.CharField( + phone_number = CharField( required=False, max_length=15, allow_blank=True, @@ -180,32 +187,33 @@ class UserUpdateSerializer(serializers.ModelSerializer): def validate_username(self, value): user = self.context['request'].user + # is_valid 由视图侧 sync_to_async 包裹(DB 查询) if FUser.objects.filter(username=value).exclude(id=user.id).exists(): - raise serializers.ValidationError('该用户名已被使用') + raise ValidationError('该用户名已被使用') return value def validate_phone_number(self, value): if value and not value.isdigit(): - raise serializers.ValidationError('手机号只能包含数字') + raise ValidationError('手机号只能包含数字') return value - def update(self, instance, validated_data): + async def aupdate(self, instance, validated_data): for attr, value in validated_data.items(): setattr(instance, attr, value) - instance.save() + await instance.asave() return instance -class PointTransactionSerializer(serializers.Serializer): - id = serializers.IntegerField(read_only=True) - transaction_type = serializers.CharField(read_only=True) - currency_type = serializers.CharField(read_only=True) - amount = serializers.IntegerField(read_only=True) - balance_after = serializers.IntegerField(read_only=True) - description = serializers.CharField(read_only=True) - created_at = serializers.DateTimeField(read_only=True) +class PointTransactionSerializer(Serializer): + id = IntegerField(read_only=True) + transaction_type = CharField(read_only=True) + currency_type = CharField(read_only=True) + amount = IntegerField(read_only=True) + balance_after = IntegerField(read_only=True) + description = CharField(read_only=True) + created_at = DateTimeField(read_only=True) - def to_representation(self, instance): + async def ato_representation(self, instance): return { 'id': instance.id, 'transaction_type': instance.transaction_type, @@ -219,12 +227,12 @@ class PointTransactionSerializer(serializers.Serializer): } -class SendEmailCodeSerializer(serializers.Serializer): - email = serializers.EmailField(required=True, error_messages={ +class SendEmailCodeSerializer(Serializer): + email = EmailField(required=True, error_messages={ 'required': '邮箱地址不能为空', 'invalid': '邮箱格式不正确', }) - target = serializers.ChoiceField( + target = ChoiceField( required=True, choices=['old', 'new'], error_messages={ @@ -239,36 +247,38 @@ class SendEmailCodeSerializer(serializers.Serializer): if target == 'old': if user.email != value: - raise serializers.ValidationError('请输入当前绑定的邮箱地址') + raise ValidationError('请输入当前绑定的邮箱地址') elif target == 'new': if user.email == value: - raise serializers.ValidationError('新邮箱与当前邮箱相同') + raise ValidationError('新邮箱与当前邮箱相同') + # is_valid 由视图侧 sync_to_async 包裹(DB 查询) if FUser.objects.filter(email=value).exists(): - raise serializers.ValidationError('该邮箱已被其他账号使用') + raise ValidationError('该邮箱已被其他账号使用') return value - def save(self): + async def asave(self): + from utils.safe_task import submit_task + from asgiref.sync import sync_to_async email = self.validated_data['email'] target = self.validated_data['target'] user = self.context['request'].user code = RandCode.get_digit_characters_code_6() cache_key = f"email_change_{target}_{user.id}" - default_cache.set(cache_key, {'code': code, 'email': email}, timeout=600) + await aset_cache(cache_key, {'code': code, 'email': email}, timeout=600) from ..tasks import send_change_email_task - from utils.safe_task import submit_task - submit_task(send_change_email_task, email, code, target) + await sync_to_async(submit_task)(send_change_email_task, email, code, target) return email -class BlacklistSerializer(serializers.ModelSerializer): - blocked_user_id = serializers.IntegerField(write_only=True) - blocked_user_username = serializers.CharField(source='blocked_user.username', read_only=True) - blocked_user_email = serializers.CharField(source='blocked_user.email', read_only=True) - blocked_user_avatar = serializers.ImageField(source='blocked_user.avatar', read_only=True) +class BlacklistSerializer(ModelSerializer): + blocked_user_id = IntegerField(write_only=True) + blocked_user_username = CharField(source='blocked_user.username', read_only=True) + blocked_user_email = CharField(source='blocked_user.email', read_only=True) + blocked_user_avatar = ImageField(source='blocked_user.avatar', read_only=True) class Meta: model = None @@ -282,27 +292,28 @@ class BlacklistSerializer(serializers.ModelSerializer): def validate_blocked_user_id(self, value): from ..models import FUser + # is_valid 由视图侧 sync_to_async 包裹(DB 查询) try: FUser.objects.get(id=value) except FUser.DoesNotExist: - raise serializers.ValidationError('用户不存在') + raise ValidationError('用户不存在') return value def validate(self, attrs): user = self.context['request'].user blocked_user_id = attrs.get('blocked_user_id') if user.id == blocked_user_id: - raise serializers.ValidationError({'blocked_user_id': '不能将自己加入黑名单'}) + raise ValidationError({'blocked_user_id': '不能将自己加入黑名单'}) from ..models import Blacklist if Blacklist.objects.filter(user=user, blocked_user_id=blocked_user_id).exists(): - raise serializers.ValidationError({'blocked_user_id': '该用户已在黑名单中'}) + raise ValidationError({'blocked_user_id': '该用户已在黑名单中'}) return attrs - def create(self, validated_data): + async def acreate(self, validated_data): from ..models import Blacklist, FUser user = self.context['request'].user - blocked_user = FUser.objects.get(id=validated_data['blocked_user_id']) - blacklist = Blacklist.objects.create( + blocked_user = await FUser.objects.aget(id=validated_data['blocked_user_id']) + blacklist = await Blacklist.objects.acreate( user=user, blocked_user=blocked_user, reason=validated_data.get('reason', '') @@ -310,12 +321,12 @@ class BlacklistSerializer(serializers.ModelSerializer): return blacklist -class ChangeEmailSerializer(serializers.Serializer): - new_email = serializers.EmailField(required=True, error_messages={ +class ChangeEmailSerializer(Serializer): + new_email = EmailField(required=True, error_messages={ 'required': '新邮箱地址不能为空', 'invalid': '邮箱格式不正确', }) - code = serializers.CharField(required=True, max_length=6, min_length=6, error_messages={ + code = CharField(required=True, max_length=6, min_length=6, error_messages={ 'required': '验证码不能为空', 'min_length': '验证码必须为6位', 'max_length': '验证码必须为6位', @@ -324,50 +335,52 @@ class ChangeEmailSerializer(serializers.Serializer): def validate_new_email(self, value): user = self.context['request'].user if user.email == value: - raise serializers.ValidationError('新邮箱与当前邮箱相同') + raise ValidationError('新邮箱与当前邮箱相同') + # is_valid 由视图侧 sync_to_async 包裹(DB 查询) if FUser.objects.filter(email=value).exists(): - raise serializers.ValidationError('该邮箱已被其他账号使用') + raise ValidationError('该邮箱已被其他账号使用') return value - def validate(self, attrs): + async def avalidate(self, attrs): + """异步版 validate(验证码在 Redis),由视图在 is_valid 后调用 await serializer.avalidate(attrs) 复核""" user = self.context['request'].user code = attrs['code'] new_email = attrs['new_email'] cache_key = f"email_change_new_{user.id}" - cached = default_cache.get(cache_key) + cached = await aget_cache(cache_key) if cached is None: - raise serializers.ValidationError({'code': '验证码已过期,请重新获取'}) + raise ValidationError({'code': '验证码已过期,请重新获取'}) if cached['code'] != code: - raise serializers.ValidationError({'code': '验证码错误'}) + raise ValidationError({'code': '验证码错误'}) if cached['email'] != new_email: - raise serializers.ValidationError({'new_email': '邮箱与发送验证码时的邮箱不一致'}) + raise ValidationError({'new_email': '邮箱与发送验证码时的邮箱不一致'}) return attrs - def save(self): + async def asave(self): user = self.context['request'].user user.email = self.validated_data['new_email'] - user.save() + await user.asave() cache_key_old = f"email_change_old_{user.id}" cache_key_new = f"email_change_new_{user.id}" - default_cache.delete(cache_key_old) - default_cache.delete(cache_key_new) + await adelete_cache(cache_key_old) + await adelete_cache(cache_key_new) return user -class LoginRecordSerializer(serializers.Serializer): - id = serializers.IntegerField(read_only=True) - device = serializers.CharField(read_only=True) - ip_address = serializers.IPAddressField(read_only=True) - location = serializers.CharField(read_only=True) - login_time = serializers.DateTimeField(read_only=True) - status = serializers.CharField(read_only=True) +class LoginRecordSerializer(Serializer): + id = IntegerField(read_only=True) + device = CharField(read_only=True) + ip_address = IPAddressField(read_only=True) + location = CharField(read_only=True) + login_time = DateTimeField(read_only=True) + status = CharField(read_only=True) - def to_representation(self, instance): + async def ato_representation(self, instance): return { 'id': instance.id, 'device': instance.device, @@ -379,12 +392,12 @@ class LoginRecordSerializer(serializers.Serializer): } -class SendPhoneCodeSerializer(serializers.Serializer): - phone = serializers.CharField(required=True, max_length=15, error_messages={ +class SendPhoneCodeSerializer(Serializer): + phone = CharField(required=True, max_length=15, error_messages={ 'required': '手机号不能为空', 'max_length': '手机号不能超过15个字符', }) - target = serializers.ChoiceField( + target = ChoiceField( required=True, choices=['old', 'new'], error_messages={ @@ -395,29 +408,29 @@ class SendPhoneCodeSerializer(serializers.Serializer): def validate_phone(self, value): if not value.isdigit(): - raise serializers.ValidationError('手机号只能包含数字') + raise ValidationError('手机号只能包含数字') if len(value) != 11: - raise serializers.ValidationError('手机号必须为11位') + raise ValidationError('手机号必须为11位') return value - def save(self): + async def asave(self): phone = self.validated_data['phone'] target = self.validated_data['target'] user = self.context['request'].user code = RandCode.get_digit_characters_code_6() cache_key = f"phone_change_{target}_{user.id}" - default_cache.set(cache_key, {'code': code, 'phone': phone}, timeout=300) + await aset_cache(cache_key, {'code': code, 'phone': phone}, timeout=300) return phone -class ChangePhoneSerializer(serializers.Serializer): - new_phone = serializers.CharField(required=True, max_length=15, error_messages={ +class ChangePhoneSerializer(Serializer): + new_phone = CharField(required=True, max_length=15, error_messages={ 'required': '新手机号不能为空', 'max_length': '手机号不能超过15个字符', }) - code = serializers.CharField(required=True, max_length=6, min_length=6, error_messages={ + code = CharField(required=True, max_length=6, min_length=6, error_messages={ 'required': '验证码不能为空', 'min_length': '验证码必须为6位', 'max_length': '验证码必须为6位', @@ -425,39 +438,40 @@ class ChangePhoneSerializer(serializers.Serializer): def validate_new_phone(self, value): if not value.isdigit(): - raise serializers.ValidationError('手机号只能包含数字') + raise ValidationError('手机号只能包含数字') if len(value) != 11: - raise serializers.ValidationError('手机号必须为11位') + raise ValidationError('手机号必须为11位') user = self.context['request'].user if user.phone_number == value: - raise serializers.ValidationError('新手机号与当前手机号相同') + raise ValidationError('新手机号与当前手机号相同') return value - def validate(self, attrs): + async def avalidate(self, attrs): + """异步版 validate(验证码在 Redis),由视图在 is_valid 后调用""" user = self.context['request'].user code = attrs['code'] new_phone = attrs['new_phone'] cache_key = f"phone_change_new_{user.id}" - cached = default_cache.get(cache_key) + cached = await aget_cache(cache_key) if cached is None: - raise serializers.ValidationError({'code': '验证码已过期,请重新获取'}) + raise ValidationError({'code': '验证码已过期,请重新获取'}) if cached['code'] != code: - raise serializers.ValidationError({'code': '验证码错误'}) + raise ValidationError({'code': '验证码错误'}) if cached['phone'] != new_phone: - raise serializers.ValidationError({'new_phone': '手机号与发送验证码时的手机号不一致'}) + raise ValidationError({'new_phone': '手机号与发送验证码时的手机号不一致'}) return attrs - def save(self): + async def asave(self): user = self.context['request'].user user.phone_number = self.validated_data['new_phone'] - user.save() + await user.asave() cache_key_old = f"phone_change_old_{user.id}" cache_key_new = f"phone_change_new_{user.id}" - default_cache.delete(cache_key_old) - default_cache.delete(cache_key_new) + await adelete_cache(cache_key_old) + await adelete_cache(cache_key_new) return user diff --git a/user/views/activities.py b/user/views/activities.py index 009ca77..af581ee 100644 --- a/user/views/activities.py +++ b/user/views/activities.py @@ -1,5 +1,5 @@ from django.db.models import Q -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.permissions import IsAuthenticated from rest_framework import status from drf_yasg.utils import swagger_auto_schema @@ -27,7 +27,7 @@ class UserActivityView(APIView): ], responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user page = int(request.query_params.get('page', 1)) page_size = int(request.query_params.get('page_size', 20)) @@ -36,7 +36,7 @@ class UserActivityView(APIView): # 文章活动 articles = Article.objects.filter(author=user, status='published').order_by('-created_at')[:50] - for article in articles: + async for article in articles: activities.append({ 'id': f'article_{article.id}', 'type': 'article', @@ -50,7 +50,7 @@ class UserActivityView(APIView): # 收藏活动 favorites = CourseFavorite.objects.filter(user=user).select_related('course').order_by('-created_at')[:50] - for fav in favorites: + async for fav in favorites: activities.append({ 'id': f'favorite_{fav.id}', 'type': 'favorite', @@ -64,7 +64,7 @@ class UserActivityView(APIView): # 学习进度 reads = ChapterRead.objects.filter(user=user, completed=True).select_related('chapter', 'course').order_by('-completed_at')[:50] - for read in reads: + async for read in reads: activities.append({ 'id': f'study_{read.id}', 'type': 'study', diff --git a/user/views/blacklist.py b/user/views/blacklist.py index fe21e71..9d62ab9 100644 --- a/user/views/blacklist.py +++ b/user/views/blacklist.py @@ -1,7 +1,8 @@ -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from django.db.models import Q +from asgiref.sync import sync_to_async from ..models import Blacklist, FUser from ..serializers.user_serializers import BlacklistSerializer from utils.response_codes import ( @@ -26,7 +27,7 @@ class BlacklistListAPIView(APIView): ], responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): queryset = Blacklist.objects.filter(user=request.user) search = request.query_params.get('search', '') @@ -39,11 +40,12 @@ class BlacklistListAPIView(APIView): page = int(request.query_params.get('page', 1)) page_size = int(request.query_params.get('page_size', 10)) - total = queryset.count() + total = await queryset.acount() start = (page - 1) * page_size end = start + page_size - serializer = BlacklistSerializer(queryset[start:end], many=True) + items = [item async for item in queryset[start:end]] + serializer = BlacklistSerializer(items, many=True) return create_standardized_response( data={ @@ -66,13 +68,14 @@ class BlacklistAddAPIView(APIView): request_body=BlacklistSerializer, responses={201: success_response, 400: error_response, 401: unauthorized_response}, ) - def post(self, request): + async def post(self, request): serializer = BlacklistSerializer( data=request.data, context={'request': request} ) - if not serializer.is_valid(): + # 校验器内部含同步 ORM 查询(validate_blocked_user_id / validate),线程池兜底 + if not await sync_to_async(serializer.is_valid)(): errors = serializer.errors first_error = '' for field, msgs in errors.items(): @@ -87,7 +90,7 @@ class BlacklistAddAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - blacklist = serializer.save() + blacklist = await sync_to_async(serializer.save)() result_serializer = BlacklistSerializer(blacklist) return create_standardized_response( @@ -107,11 +110,11 @@ class BlacklistCheckAPIView(APIView): ], responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request, user_id): - is_blocked = Blacklist.objects.filter( + async def get(self, request, user_id): + is_blocked = await Blacklist.objects.filter( user=request.user, blocked_user_id=user_id - ).exists() + ).aexists() return create_standardized_response( data={'is_blocked': is_blocked}, code=ResponseCode.SUCCESS, @@ -129,9 +132,9 @@ class BlacklistRemoveAPIView(APIView): ], responses={200: success_response, 401: unauthorized_response, 404: not_found_response}, ) - def delete(self, request, pk): + async def delete(self, request, pk): try: - blacklist = Blacklist.objects.get(id=pk, user=request.user) + blacklist = await Blacklist.objects.aget(id=pk, user=request.user) except Blacklist.DoesNotExist: return create_standardized_error_response( code=ResponseCode.NOT_FOUND, @@ -139,7 +142,7 @@ class BlacklistRemoveAPIView(APIView): status_code=status.HTTP_404_NOT_FOUND ) - blacklist.delete() + await blacklist.adelete() return create_standardized_response( data={'deleted': True}, diff --git a/user/views/captcha.py b/user/views/captcha.py index 763c3bb..db99c6e 100644 --- a/user/views/captcha.py +++ b/user/views/captcha.py @@ -1,5 +1,5 @@ from rest_framework.permissions import AllowAny -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework import status from drf_yasg.utils import swagger_auto_schema from drf_yasg import openapi diff --git a/user/views/email.py b/user/views/email.py index 5660f05..1fc74cb 100644 --- a/user/views/email.py +++ b/user/views/email.py @@ -1,6 +1,7 @@ from rest_framework import status -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.permissions import IsAuthenticated +from asgiref.sync import sync_to_async from drf_yasg.utils import swagger_auto_schema from drf_yasg import openapi import logging @@ -38,12 +39,12 @@ class SendChangeEmailCodeAPIView(APIView): ), responses={200: success_response, 400: error_response, 401: unauthorized_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): from utils.captcha import check_captcha_required, verify_captcha, record_failure, reset_failures identifier = str(request.user.id) operation = 'change_email' - captcha_required = check_captcha_required(operation, identifier) + captcha_required = await sync_to_async(check_captcha_required)(operation, identifier) if captcha_required: captcha_key = request.data.get('captcha_key', None) @@ -55,14 +56,14 @@ class SendChangeEmailCodeAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - captcha_result = verify_captcha(captcha_key, captcha_code) + captcha_result = await sync_to_async(verify_captcha)(captcha_key, captcha_code) if captcha_result == 'expired': return create_standardized_error_response( code=ResponseCode.CAPTCHA_EXPIRED, status_code=status.HTTP_400_BAD_REQUEST ) elif captcha_result == 'wrong': - record_failure(operation, identifier) + await sync_to_async(record_failure)(operation, identifier) return create_standardized_error_response( code=ResponseCode.CAPTCHA_ERROR, status_code=status.HTTP_400_BAD_REQUEST @@ -73,7 +74,8 @@ class SendChangeEmailCodeAPIView(APIView): context={'request': request} ) - if not serializer.is_valid(): + # 校验器内含同步 ORM(validate_email 查重)与 cache 写入,线程池兜底 + if not await sync_to_async(serializer.is_valid)(): errors = serializer.errors first_error = '' for field, msgs in errors.items(): @@ -81,7 +83,7 @@ class SendChangeEmailCodeAPIView(APIView): first_error = str(msgs[0]) break - record_failure(operation, identifier) + await sync_to_async(record_failure)(operation, identifier) return create_standardized_error_response( data=errors, code=ResponseCode.PARAMETER_ERROR, @@ -90,24 +92,25 @@ class SendChangeEmailCodeAPIView(APIView): ) email = serializer.validated_data.get('email') - if email and not validate_email_mx(email): + if email and not await sync_to_async(validate_email_mx)(email): logger.warning(f'[ChangeEmail] Domain MX check failed: email={email}') - record_failure(operation, identifier) + await sync_to_async(record_failure)(operation, identifier) return create_standardized_error_response( code=ResponseCode.EMAIL_DOMAIN_INVALID, status_code=status.HTTP_400_BAD_REQUEST ) try: - email = serializer.save() - reset_failures(operation, identifier) + # save() 内部触发验证码邮件发送(Celery/cache/SMTP 链路),线程池兜底 + email = await sync_to_async(serializer.save)() + await sync_to_async(reset_failures)(operation, identifier) return create_standardized_response( data={'email_sent': True}, code=ResponseCode.EMAIL_CHANGE_CODE_SENT, status_code=status.HTTP_200_OK ) except Exception as e: - record_failure(operation, identifier) + await sync_to_async(record_failure)(operation, identifier) return create_standardized_error_response( message=f'邮件发送失败: {str(e)}', code=ResponseCode.EMAIL_SEND_FAILED, @@ -131,13 +134,14 @@ class ChangeEmailAPIView(APIView): ), responses={200: success_response, 400: error_response, 401: unauthorized_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): serializer = ChangeEmailSerializer( data=request.data, context={'request': request} ) - if not serializer.is_valid(): + # 校验器内含同步 ORM 查询,线程池兜底 + if not await sync_to_async(serializer.is_valid)(): errors = serializer.errors first_error = '' for field, msgs in errors.items(): @@ -153,7 +157,7 @@ class ChangeEmailAPIView(APIView): ) try: - updated_user = serializer.save() + updated_user = await sync_to_async(serializer.save)() user_serializer = UserSerializer(updated_user) return create_standardized_response( diff --git a/user/views/favorites.py b/user/views/favorites.py index d5ba5bb..ca93344 100644 --- a/user/views/favorites.py +++ b/user/views/favorites.py @@ -1,4 +1,4 @@ -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.permissions import IsAuthenticated from rest_framework import status @@ -24,7 +24,7 @@ class MyFavoritesView(APIView): ], responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): content_type = request.query_params.get('type', 'all') results = [] @@ -32,7 +32,7 @@ class MyFavoritesView(APIView): tool_favs = ToolFavorite.objects.filter( user=request.user ).select_related('tool', 'tool__category').order_by('-created_at') - for tf in tool_favs: + async for tf in tool_favs: results.append({ 'id': tf.tool.id, 'type': 'tool', @@ -49,7 +49,7 @@ class MyFavoritesView(APIView): article_favs = ArticleFavorite.objects.filter( user=request.user ).select_related('article', 'article__author').order_by('-created_at') - for af in article_favs: + async for af in article_favs: results.append({ 'id': af.article.id, 'type': 'article', @@ -69,7 +69,7 @@ class MyFavoritesView(APIView): course_favs = CourseFavorite.objects.filter( user=request.user ).select_related('course', 'course__author').order_by('-created_at') - for cf in course_favs: + async for cf in course_favs: results.append({ 'id': cf.course.id, 'type': 'course', @@ -79,7 +79,7 @@ class MyFavoritesView(APIView): 'level': cf.course.level, 'icon_name': getattr(cf.course, 'icon_name', ''), 'color': getattr(cf.course, 'color', ''), - 'chapters_count': cf.course.chapters.count() if hasattr(cf.course, 'chapters') else 0, + 'chapters_count': await cf.course.chapters.acount() if hasattr(cf.course, 'chapters') else 0, 'author': getattr(cf.course.author, 'nickname', '') if cf.course.author else '', 'author_id': cf.course.author_id, 'created_at': cf.created_at.strftime('%Y-%m-%d %H:%M:%S'), @@ -89,7 +89,7 @@ class MyFavoritesView(APIView): api_favs = ApiFavorite.objects.filter( user=request.user ).select_related('api_item').order_by('-created_at') - for af in api_favs: + async for af in api_favs: results.append({ 'id': af.api_item.id, 'type': 'api', @@ -104,10 +104,10 @@ class MyFavoritesView(APIView): counts = { 'total': len(results), - 'tool': ToolFavorite.objects.filter(user=request.user).count(), - 'article': ArticleFavorite.objects.filter(user=request.user).count(), - 'course': CourseFavorite.objects.filter(user=request.user).count(), - 'api': ApiFavorite.objects.filter(user=request.user).count(), + 'tool': await ToolFavorite.objects.filter(user=request.user).acount(), + 'article': await ArticleFavorite.objects.filter(user=request.user).acount(), + 'course': await CourseFavorite.objects.filter(user=request.user).acount(), + 'api': await ApiFavorite.objects.filter(user=request.user).acount(), } return create_standardized_response( diff --git a/user/views/index.py b/user/views/index.py index 77efa3c..b147f93 100644 --- a/user/views/index.py +++ b/user/views/index.py @@ -1,7 +1,7 @@ from django.http import JsonResponse -def user_index(request): +async def user_index(request): return JsonResponse({ 'code': 200, 'message': 'User API index', diff --git a/user/views/invitation.py b/user/views/invitation.py index ea97ce2..f5370ca 100644 --- a/user/views/invitation.py +++ b/user/views/invitation.py @@ -1,4 +1,4 @@ -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework.permissions import IsAuthenticated from user.models import Invitation, FUser @@ -16,12 +16,12 @@ class InviteCodeAPIView(APIView): operation_description='获取当前用户的邀请码和邀请链接', responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user # Generate or get existing invite code invite_code = f"CY{str(user.id).zfill(6)}" # Ensure the invitation record exists - Invitation.objects.get_or_create( + await Invitation.objects.aget_or_create( inviter=user, invite_code=invite_code, defaults={'is_used': False} @@ -41,11 +41,11 @@ class InviteStatsAPIView(APIView): operation_description='获取当前用户的邀请统计信息,包括已邀请人数、获得积分和待使用邀请数', responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user - used_count = Invitation.objects.filter(inviter=user, is_used=True).count() + used_count = await Invitation.objects.filter(inviter=user, is_used=True).acount() total_points = used_count * 500 # 500 points per invite - pending_count = Invitation.objects.filter(inviter=user, is_used=False).count() + pending_count = await Invitation.objects.filter(inviter=user, is_used=False).acount() return Response({ 'invited_count': used_count, 'total_points': total_points, @@ -62,11 +62,11 @@ class InviteRecordsAPIView(APIView): operation_description='获取当前用户的所有邀请记录列表', responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user records = Invitation.objects.filter(inviter=user).order_by('-created_at') data = [] - for r in records: + async for r in records: data.append({ 'id': r.id, 'invite_code': r.invite_code, diff --git a/user/views/login_record.py b/user/views/login_record.py index c6789be..8ce40a2 100644 --- a/user/views/login_record.py +++ b/user/views/login_record.py @@ -1,5 +1,5 @@ from rest_framework import status -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.permissions import IsAuthenticated from django.utils import timezone from django.utils.decorators import method_decorator @@ -36,7 +36,7 @@ class LoginRecordListAPIView(APIView): ], responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user page = int(request.query_params.get('page', 1)) page_size = int(request.query_params.get('page_size', 10)) @@ -67,11 +67,12 @@ class LoginRecordListAPIView(APIView): elif start_date and end_date: records = records.filter(login_time__date__gte=start_date, login_time__date__lte=end_date) - total = records.count() + total = await records.acount() start = (page - 1) * page_size end = start + page_size - serializer = LoginRecordSerializer(records[start:end], many=True) + items = [item async for item in records[start:end]] + serializer = LoginRecordSerializer(items, many=True) return create_standardized_response( data={ @@ -99,11 +100,11 @@ class LoginRecordDeleteAPIView(APIView): ], responses={200: success_response, 401: unauthorized_response, 404: not_found_response}, ) - def delete(self, request, pk): + async def delete(self, request, pk): user = request.user try: - record = LoginRecord.objects.get(pk=pk, user=user) - record.delete() + record = await LoginRecord.objects.aget(pk=pk, user=user) + await record.adelete() return create_standardized_response( code=ResponseCode.SUCCESS, status_code=status.HTTP_200_OK @@ -126,9 +127,9 @@ class LoginRecordClearAPIView(APIView): operation_description='清空当前用户的所有登录记录', responses={200: success_response, 401: unauthorized_response}, ) - def delete(self, request): + async def delete(self, request): user = request.user - deleted_count, _ = LoginRecord.objects.filter(user=user).delete() + deleted_count, _ = await LoginRecord.objects.filter(user=user).adelete() return create_standardized_response( data={'deleted_count': deleted_count}, code=ResponseCode.SUCCESS, diff --git a/user/views/phone.py b/user/views/phone.py index 496e8dd..e59558b 100644 --- a/user/views/phone.py +++ b/user/views/phone.py @@ -1,6 +1,7 @@ from rest_framework import status -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.permissions import IsAuthenticated +from asgiref.sync import sync_to_async from utils.response_codes import ( ResponseCode, @@ -17,13 +18,14 @@ from ..serializers.user_serializers import ( class SendPhoneCodeAPIView(APIView): permission_classes = [IsAuthenticated] - def post(self, request): + async def post(self, request): serializer = SendPhoneCodeSerializer( data=request.data, context={'request': request} ) - if not serializer.is_valid(): + # 校验/save 内含 cache 写入,线程池兜底 + if not await sync_to_async(serializer.is_valid)(): errors = serializer.errors first_error = '' for field, msgs in errors.items(): @@ -39,7 +41,7 @@ class SendPhoneCodeAPIView(APIView): ) try: - phone = serializer.save() + phone = await sync_to_async(serializer.save)() return create_standardized_response( data={'phone_sent': True}, code=ResponseCode.SUCCESS, @@ -56,13 +58,14 @@ class SendPhoneCodeAPIView(APIView): class ChangePhoneAPIView(APIView): permission_classes = [IsAuthenticated] - def post(self, request): + async def post(self, request): serializer = ChangePhoneSerializer( data=request.data, context={'request': request} ) - if not serializer.is_valid(): + # 校验/save 内含 cache 读写与 user.save() ORM 调用,线程池兜底 + if not await sync_to_async(serializer.is_valid)(): errors = serializer.errors first_error = '' for field, msgs in errors.items(): @@ -78,7 +81,7 @@ class ChangePhoneAPIView(APIView): ) try: - updated_user = serializer.save() + updated_user = await sync_to_async(serializer.save)() user_serializer = UserSerializer(updated_user) return create_standardized_response( diff --git a/user/views/qr_login.py b/user/views/qr_login.py index 2396f53..25efb40 100644 --- a/user/views/qr_login.py +++ b/user/views/qr_login.py @@ -1,6 +1,7 @@ -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.permissions import AllowAny, IsAuthenticated from rest_framework import status +from asgiref.sync import sync_to_async from drf_yasg.utils import swagger_auto_schema from drf_yasg import openapi @@ -22,8 +23,8 @@ class QRTokenView(APIView): operation_description='生成唯一的 QR Token,返回 token 供前端生成二维码', responses={200: success_response}, ) - def post(self, request): - token = generate_qr_token() + async def post(self, request): + token = await sync_to_async(generate_qr_token)() return create_standardized_response(data={ "token": token, "expires_in": 300, @@ -42,8 +43,9 @@ class QRStatusView(APIView): ], responses={200: success_response}, ) - def get(self, request, token): - data = get_qr_status(token) + async def get(self, request, token): + # get_qr_status 走同步 cache,线程池兜底 + data = await sync_to_async(get_qr_status)(token) response_data = { "status": data.get("status", "expired"), "username": data.get("scan_username"), @@ -51,7 +53,7 @@ class QRStatusView(APIView): if data.get("status") == "confirmed": scan_user_id = data.get("scan_user_id") - user = FUser.objects.filter(id=scan_user_id).first() + user = await FUser.objects.filter(id=scan_user_id).afirst() if user: refresh = RefreshToken.for_user(user) response_data["auth"] = { @@ -81,7 +83,7 @@ class QRScanView(APIView): ), responses={200: success_response, 400: error_response, 401: error_response}, ) - def post(self, request): + async def post(self, request): token = request.data.get("token") if not token: return create_standardized_error_response( @@ -90,7 +92,8 @@ class QRScanView(APIView): status_code=status.HTTP_400_BAD_REQUEST, ) user = request.user - success = scan_qr(token, user.id, user.username) + # scan_qr 走同步 cache,线程池兜底 + success = await sync_to_async(scan_qr)(token, user.id, user.username) if not success: return create_standardized_error_response( code=ResponseCode.VALIDATION_ERROR, @@ -116,7 +119,7 @@ class QRConfirmView(APIView): ), responses={200: success_response, 400: error_response, 401: error_response}, ) - def post(self, request): + async def post(self, request): token = request.data.get("token") if not token: return create_standardized_error_response( @@ -124,14 +127,16 @@ class QRConfirmView(APIView): message="缺少 token", status_code=status.HTTP_400_BAD_REQUEST, ) - data = confirm_qr(token) + # confirm_qr 走同步 cache,线程池兜底 + data = await sync_to_async(confirm_qr)(token) if not data: return create_standardized_error_response( code=ResponseCode.VALIDATION_ERROR, message="确认失败,请重新扫码", status_code=status.HTTP_400_BAD_REQUEST, ) - create_login_record(request, request.user, 'success') + # create_login_record 内含 ORM 写入,线程池兜底 + await sync_to_async(create_login_record)(request, request.user, 'success') return create_standardized_response(message="登录确认成功") @@ -150,8 +155,9 @@ class QRCancelView(APIView): ), responses={200: success_response, 401: error_response}, ) - def post(self, request): + async def post(self, request): token = request.data.get("token") if token: - cancel_qr(token) + # cancel_qr 走同步 cache,线程池兜底 + await sync_to_async(cancel_qr)(token) return create_standardized_response(message="已取消") diff --git a/user/views/region.py b/user/views/region.py index 8bac34e..3d35010 100644 --- a/user/views/region.py +++ b/user/views/region.py @@ -1,5 +1,5 @@ from rest_framework.permissions import AllowAny -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from drf_yasg.utils import swagger_auto_schema @@ -59,7 +59,7 @@ class RegionListView(APIView): ), } ) - def get(self, request): + async def get(self, request): level = request.GET.get('level') parent_code = request.GET.get('parent_code') @@ -78,7 +78,8 @@ class RegionListView(APIView): else: queryset = queryset.filter(level=1) - serializer = RegionSimpleSerializer(queryset, many=True) + regions = [r async for r in queryset] + serializer = RegionSimpleSerializer(regions, many=True) return Response( {"code": 200, "message": "success", "data": serializer.data}, status=status.HTTP_200_OK diff --git a/user/views/settings.py b/user/views/settings.py index d769be9..75c106e 100644 --- a/user/views/settings.py +++ b/user/views/settings.py @@ -1,7 +1,8 @@ import json -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.permissions import IsAuthenticated from rest_framework.parsers import MultiPartParser, FormParser +from asgiref.sync import sync_to_async from drf_yasg.utils import swagger_auto_schema from drf_yasg import openapi from user.models import UserDevice @@ -25,7 +26,7 @@ class UploadCoverAPIView(APIView): ], responses={200: success_response, 400: '参数错误', 401: unauthorized_response}, ) - def post(self, request): + async def post(self, request): cover_file = request.FILES.get('cover') if not cover_file: return create_standardized_error_response( @@ -49,12 +50,13 @@ class UploadCoverAPIView(APIView): user = request.user if user.cover_image: try: - user.cover_image.delete(save=False) + # FieldFile.delete 是同步存储 I/O,线程池兜底 + await sync_to_async(user.cover_image.delete)(save=False) except Exception: pass user.cover_image = cover_file - user.save(update_fields=['cover_image']) + await user.asave(update_fields=['cover_image']) return create_standardized_response( data={ @@ -73,7 +75,7 @@ class PrivacySettingsAPIView(APIView): operation_description='获取当前用户的隐私设置', responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user return create_standardized_response( data={ @@ -97,7 +99,7 @@ class PrivacySettingsAPIView(APIView): ), responses={200: success_response, 400: '参数错误', 401: unauthorized_response}, ) - def put(self, request): + async def put(self, request): user = request.user valid_choices = ['public', 'friends_only', 'private'] @@ -112,7 +114,7 @@ class PrivacySettingsAPIView(APIView): ) setattr(user, field, value) - user.save(update_fields=fields) + await user.asave(update_fields=fields) return create_standardized_response( data={ 'privacy_profile': user.privacy_profile, @@ -132,7 +134,7 @@ class NotificationSettingsAPIView(APIView): operation_description='获取当前用户的消息通知偏好设置', responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user return create_standardized_response( data={ @@ -162,7 +164,7 @@ class NotificationSettingsAPIView(APIView): ), responses={200: success_response, 401: unauthorized_response}, ) - def put(self, request): + async def put(self, request): user = request.user fields = ['notify_email', 'notify_browser', 'notify_reply', 'notify_like', 'notify_follow', 'notify_system'] @@ -173,7 +175,7 @@ class NotificationSettingsAPIView(APIView): value = str(value).lower() in ('true', '1', 'yes') setattr(user, field, value) - user.save(update_fields=fields) + await user.asave(update_fields=fields) return create_standardized_response( data={ 'notify_email': user.notify_email, @@ -196,11 +198,11 @@ class UserDeviceListView(APIView): operation_description='获取当前用户的所有登录设备列表', responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user devices = UserDevice.objects.filter(user=user).order_by('-last_active') data = [] - for d in devices: + async for d in devices: data.append({ 'id': d.id, 'device_name': d.device_name, @@ -228,9 +230,9 @@ class UserDeviceRemoveView(APIView): ], responses={200: success_response, 404: '设备不存在', 401: unauthorized_response}, ) - def delete(self, request, pk): + async def delete(self, request, pk): user = request.user - device = UserDevice.objects.filter(user=user, pk=pk).first() + device = await UserDevice.objects.filter(user=user, pk=pk).afirst() if not device: return create_standardized_error_response( code=40401, @@ -241,7 +243,7 @@ class UserDeviceRemoveView(APIView): code=40001, message='不能移除当前设备' ) - device.delete() + await device.adelete() return create_standardized_response(data={'message': '设备已移除'}) @@ -254,7 +256,7 @@ class UserDeviceClearOthersView(APIView): operation_description='移除除当前设备外的所有其他登录设备', responses={200: success_response, 401: unauthorized_response}, ) - def post(self, request): + async def post(self, request): user = request.user - count, _ = UserDevice.objects.filter(user=user).exclude(is_current=True).delete() + count, _ = await UserDevice.objects.filter(user=user).exclude(is_current=True).adelete() return create_standardized_response(data={'message': f'已移除 {count} 个设备', 'removed_count': count}) diff --git a/user/views/slider_captcha.py b/user/views/slider_captcha.py index b2fe58f..24b729f 100644 --- a/user/views/slider_captcha.py +++ b/user/views/slider_captcha.py @@ -1,5 +1,5 @@ from rest_framework.permissions import AllowAny -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework import status from drf_yasg.utils import swagger_auto_schema from drf_yasg import openapi diff --git a/user/views/tasks.py b/user/views/tasks.py index e61a29a..62d84f6 100644 --- a/user/views/tasks.py +++ b/user/views/tasks.py @@ -1,6 +1,7 @@ from rest_framework import status -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.permissions import IsAuthenticated +from asgiref.sync import sync_to_async from django.db import transaction from django.db.models import F from django.utils import timezone @@ -121,7 +122,7 @@ class TaskListAPIView(APIView): operation_description='获取所有活跃任务列表及当前用户的进度和等级信息', responses={200: success_response, 401: unauthorized_response, 500: error_response}, ) - def get(self, request): + async def get(self, request): try: user = request.user now = timezone.localtime(timezone.now()) @@ -131,7 +132,7 @@ class TaskListAPIView(APIView): tasks = TaskDefinition.objects.filter(is_active=True) task_list = [] - for task in tasks: + async for task in tasks: if task.task_type == 'daily': period_key = daily_period elif task.task_type == 'weekly': @@ -139,9 +140,9 @@ class TaskListAPIView(APIView): else: period_key = 'permanent' - progress = UserTaskProgress.objects.filter( + progress = await UserTaskProgress.objects.filter( user=user, task=task, period_key=period_key - ).first() + ).afirst() task_list.append({ 'id': task.id, @@ -160,15 +161,19 @@ class TaskListAPIView(APIView): 'is_claimed': progress.is_claimed if progress else False, }) - user_level = UserLevel.objects.filter(user=user).first() - next_threshold = LevelThreshold.objects.filter( + user_level = await UserLevel.objects.filter(user=user).afirst() + next_threshold = await LevelThreshold.objects.filter( level__gt=user_level.level if user_level else 1 - ).order_by('level').first() if user_level else None + ).order_by('level').afirst() if user_level else None + + cur_threshold = await LevelThreshold.objects.filter( + level=user_level.level + ).afirst() if user_level else None level_data = { 'level': user_level.level if user_level else 1, 'xp': user_level.xp if user_level else 0, - 'title': (LevelThreshold.objects.filter(level=user_level.level).first().title if user_level and LevelThreshold.objects.filter(level=user_level.level).exists() else '新手'), + 'title': cur_threshold.title if cur_threshold else '新手', 'next_level_xp': next_threshold.xp_required if next_threshold else None, 'next_level_title': next_threshold.title if next_threshold else None, } @@ -208,7 +213,7 @@ class TaskTrackAPIView(APIView): ), responses={200: success_response, 400: error_response, 401: unauthorized_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): action_type = request.data.get('action_type') count = int(request.data.get('count', 1)) @@ -228,7 +233,8 @@ class TaskTrackAPIView(APIView): ) try: - updated_tasks = track_user_action(request.user, action_type, count) + # track_user_action 为同步 helper(user.py 也以 sync_to_async 调用),线程池兜底 + updated_tasks = await sync_to_async(track_user_action)(request.user, action_type, count) return create_standardized_response( data={'updated_tasks': updated_tasks}, code=ResponseCode.SUCCESS, @@ -256,10 +262,10 @@ class TaskClaimAPIView(APIView): ], responses={200: success_response, 400: error_response, 401: unauthorized_response, 500: error_response}, ) - def post(self, request, task_id): + async def post(self, request, task_id): try: user = request.user - task = TaskDefinition.objects.filter(id=task_id, is_active=True).first() + task = await TaskDefinition.objects.filter(id=task_id, is_active=True).afirst() if not task: return create_standardized_error_response( @@ -276,9 +282,9 @@ class TaskClaimAPIView(APIView): else: period_key = 'permanent' - progress = UserTaskProgress.objects.filter( + progress = await UserTaskProgress.objects.filter( user=user, task=task, period_key=period_key - ).first() + ).afirst() if not progress or not progress.is_completed: return create_standardized_error_response( @@ -294,53 +300,59 @@ class TaskClaimAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - with transaction.atomic(): - locked_user = FUser.objects.select_for_update().get(pk=user.pk) - user_level, _ = UserLevel.objects.select_for_update().get_or_create( - user=locked_user, - defaults={'xp': 0, 'level': 1} - ) - - if task.reward_points > 0: - locked_user.points = F('points') + task.reward_points - locked_user.save(update_fields=['points']) - locked_user.refresh_from_db() - PointTransaction.objects.create( + # select_for_update + 事务 + refresh_from_db 整体在同步函数内执行,线程池兜底 + def _claim_reward(): + with transaction.atomic(): + locked_user = FUser.objects.select_for_update().get(pk=user.pk) + user_level, _ = UserLevel.objects.select_for_update().get_or_create( user=locked_user, - transaction_type='earn', - currency_type='points', - amount=task.reward_points, - balance_after=locked_user.points, - description=f'完成任务: {task.name}', + defaults={'xp': 0, 'level': 1} ) - if task.reward_coins > 0: - locked_user.coins = F('coins') + task.reward_coins - locked_user.save(update_fields=['coins']) - locked_user.refresh_from_db() - PointTransaction.objects.create( - user=locked_user, - transaction_type='earn', - currency_type='coins', - amount=task.reward_coins, - balance_after=locked_user.coins, - description=f'完成任务: {task.name}', - ) + if task.reward_points > 0: + locked_user.points = F('points') + task.reward_points + locked_user.save(update_fields=['points']) + locked_user.refresh_from_db() + PointTransaction.objects.create( + user=locked_user, + transaction_type='earn', + currency_type='points', + amount=task.reward_points, + balance_after=locked_user.points, + description=f'完成任务: {task.name}', + ) + + if task.reward_coins > 0: + locked_user.coins = F('coins') + task.reward_coins + locked_user.save(update_fields=['coins']) + locked_user.refresh_from_db() + PointTransaction.objects.create( + user=locked_user, + transaction_type='earn', + currency_type='coins', + amount=task.reward_coins, + balance_after=locked_user.coins, + description=f'完成任务: {task.name}', + ) + + if task.reward_xp > 0: + user_level.xp = F('xp') + task.reward_xp + user_level.save(update_fields=['xp', 'updated_at']) + user_level.refresh_from_db() + check_level_up(user_level) + + progress.is_claimed = True + progress.claimed_at = now + progress.save(update_fields=['is_claimed', 'claimed_at']) - if task.reward_xp > 0: - user_level.xp = F('xp') + task.reward_xp - user_level.save(update_fields=['xp', 'updated_at']) user_level.refresh_from_db() - check_level_up(user_level) + next_threshold = LevelThreshold.objects.filter( + level__gt=user_level.level + ).order_by('level').first() - progress.is_claimed = True - progress.claimed_at = now - progress.save(update_fields=['is_claimed', 'claimed_at']) + return locked_user, user_level, next_threshold - user_level.refresh_from_db() - next_threshold = LevelThreshold.objects.filter( - level__gt=user_level.level - ).order_by('level').first() + locked_user, user_level, next_threshold = await sync_to_async(_claim_reward)() return create_standardized_response( data={ @@ -375,19 +387,19 @@ class UserLevelAPIView(APIView): operation_description='获取当前用户的等级、经验值、等级称号和升级进度', responses={200: success_response, 401: unauthorized_response, 500: error_response}, ) - def get(self, request): + async def get(self, request): try: user = request.user - user_level, _ = UserLevel.objects.get_or_create( + user_level, _ = await UserLevel.objects.aget_or_create( user=user, defaults={'xp': 0, 'level': 1} ) - current_threshold = LevelThreshold.objects.filter( + current_threshold = await LevelThreshold.objects.filter( level=user_level.level - ).first() - next_threshold = LevelThreshold.objects.filter( + ).afirst() + next_threshold = await LevelThreshold.objects.filter( level__gt=user_level.level - ).order_by('level').first() + ).order_by('level').afirst() current_xp = current_threshold.xp_required if current_threshold else 0 next_xp = next_threshold.xp_required if next_threshold else user_level.xp diff --git a/user/views/user.py b/user/views/user.py index 5a830cc..e1b8c0b 100644 --- a/user/views/user.py +++ b/user/views/user.py @@ -2,11 +2,12 @@ import os import uuid import logging +from asgiref.sync import sync_to_async from utils.email_utils import validate_email_mx from rest_framework.decorators import permission_classes from rest_framework.parsers import MultiPartParser, FormParser, JSONParser from rest_framework.permissions import AllowAny, IsAuthenticated -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from django.contrib.auth import authenticate @@ -50,6 +51,11 @@ def _create_login_record(request, user, record_status): create_login_record(request, user, record_status) +def _serialize_user(user): + """同步序列化助手:在 async 视图中通过 sync_to_async 调用,避免 SynchronousOnlyOperation。""" + return UserSerializer(user).data + + class SendUserEmailAPIView(APIView): permission_classes = [AllowAny] @@ -66,7 +72,7 @@ class SendUserEmailAPIView(APIView): ), responses={201: success_response, 200: success_response, 400: error_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): to_email = request.data.get('to_email', None) if to_email is None or to_email == "": @@ -75,7 +81,7 @@ class SendUserEmailAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - if not validate_email_mx(to_email): + if not await sync_to_async(validate_email_mx)(to_email): logger.warning(f'[Email] Domain MX check failed: email={to_email}') return create_standardized_error_response( code=ResponseCode.EMAIL_DOMAIN_INVALID, @@ -83,13 +89,13 @@ class SendUserEmailAPIView(APIView): ) try: - user_exists_result = FUser.objects.filter(email=to_email).exists() + user_exists_result = await FUser.objects.filter(email=to_email).aexists() if not user_exists_result: code = RandCode.get_digit_characters_code_8() - default_cache.set(f"register_{to_email}", code, timeout=600) + await sync_to_async(default_cache.set)(f"register_{to_email}", code, timeout=600) - result = submit_task(send_verification_email_task, to_email, code, 'register') + result = await sync_to_async(submit_task)(send_verification_email_task, to_email, code, 'register') if result is None: logger.warning(f'[Register] Email send failed: email={to_email}') else: @@ -102,9 +108,9 @@ class SendUserEmailAPIView(APIView): ) else: code = RandCode.get_digit_characters_code_8() - default_cache.set(f"login_{to_email}", code, timeout=600) + await sync_to_async(default_cache.set)(f"login_{to_email}", code, timeout=600) - result = submit_task(send_verification_email_task, to_email, code, 'login') + result = await sync_to_async(submit_task)(send_verification_email_task, to_email, code, 'login') if result is None: logger.warning(f'[Login] Email send failed: email={to_email}') else: @@ -141,7 +147,7 @@ class UserLoginOrRegisterAPIView(APIView): ), responses={200: success_response, 400: error_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): code = request.data.get('code', None) to_email = request.data.get('email', None) @@ -152,11 +158,11 @@ class UserLoginOrRegisterAPIView(APIView): ) try: - user = FUser.objects.filter(email=to_email).first() + user = await FUser.objects.filter(email=to_email).afirst() if user is None: # Registration flow - vcode = default_cache.get(f"register_{to_email}") + vcode = await sync_to_async(default_cache.get)(f"register_{to_email}") if vcode is None: return create_standardized_error_response( @@ -166,20 +172,21 @@ class UserLoginOrRegisterAPIView(APIView): if code == vcode: user_serializer = UserSerializer(data=request.data) - if user_serializer.is_valid(): - user = user_serializer.create_by_email(request.data) + if await sync_to_async(user_serializer.is_valid)(): + user = await sync_to_async(user_serializer.create_by_email)(request.data) refresh = RefreshToken.for_user(user) + user_data = await sync_to_async(_serialize_user)(user) # Prepare response data response_data = { - 'user': UserSerializer(user).data, + 'user': user_data, 'refresh': str(refresh), 'access': str(refresh.access_token), 'token_type': 'bearer', 'expires_at_timestamp': refresh.access_token.payload['exp'] } - _create_login_record(request, user, 'success') + await sync_to_async(_create_login_record)(request, user, 'success') return create_standardized_response( data=response_data, @@ -198,7 +205,7 @@ class UserLoginOrRegisterAPIView(APIView): ) else: # Login flow - vcode = default_cache.get(f"login_{to_email}") + vcode = await sync_to_async(default_cache.get)(f"login_{to_email}") if vcode is None: return create_standardized_error_response( @@ -208,18 +215,18 @@ class UserLoginOrRegisterAPIView(APIView): if code == vcode: refresh = RefreshToken.for_user(user) - user_serializer = UserSerializer(user) + user_data = await sync_to_async(_serialize_user)(user) # Prepare response data response_data = { - 'user': user_serializer.data, + 'user': user_data, 'refresh': str(refresh), 'access': str(refresh.access_token), 'token_type': 'bearer', 'expires_at_timestamp': refresh.access_token.payload['exp'] } - _create_login_record(request, user, 'success') + await sync_to_async(_create_login_record)(request, user, 'success') return create_standardized_response( data=response_data, @@ -227,7 +234,7 @@ class UserLoginOrRegisterAPIView(APIView): status_code=status.HTTP_200_OK ) else: - _create_login_record(request, user, 'failed') + await sync_to_async(_create_login_record)(request, user, 'failed') return create_standardized_error_response( code=ResponseCode.LOGIN_VERIFICATION_ERROR, @@ -260,7 +267,7 @@ class ForgotPasswordSendCodeAPIView(APIView): ), responses={200: success_response, 400: error_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): from utils.captcha import check_captcha_required, verify_captcha, record_failure, reset_failures to_email = request.data.get('email', None) @@ -271,7 +278,7 @@ class ForgotPasswordSendCodeAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - if not validate_email_mx(to_email): + if not await sync_to_async(validate_email_mx)(to_email): return create_standardized_error_response( code=ResponseCode.EMAIL_DOMAIN_INVALID, status_code=status.HTTP_400_BAD_REQUEST @@ -279,7 +286,7 @@ class ForgotPasswordSendCodeAPIView(APIView): identifier = to_email operation = 'forgot_send' - captcha_required = check_captcha_required(operation, identifier) + captcha_required = await sync_to_async(check_captcha_required)(operation, identifier) if captcha_required: captcha_key = request.data.get('captcha_key', None) @@ -291,39 +298,39 @@ class ForgotPasswordSendCodeAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - captcha_result = verify_captcha(captcha_key, captcha_code) + captcha_result = await sync_to_async(verify_captcha)(captcha_key, captcha_code) if captcha_result == 'expired': return create_standardized_error_response( code=ResponseCode.CAPTCHA_EXPIRED, status_code=status.HTTP_400_BAD_REQUEST ) elif captcha_result == 'wrong': - record_failure(operation, identifier) + await sync_to_async(record_failure)(operation, identifier) return create_standardized_error_response( code=ResponseCode.CAPTCHA_ERROR, status_code=status.HTTP_400_BAD_REQUEST ) try: - user = FUser.objects.filter(email=to_email).first() + user = await FUser.objects.filter(email=to_email).afirst() if user is None: - record_failure(operation, to_email) + await sync_to_async(record_failure)(operation, to_email) return create_standardized_error_response( code=ResponseCode.USER_NOT_FOUND, status_code=status.HTTP_400_BAD_REQUEST ) code = RandCode.get_digit_characters_code_8() - default_cache.set(f"reset_password_{to_email}", code, timeout=600) + await sync_to_async(default_cache.set)(f"reset_password_{to_email}", code, timeout=600) - result = submit_task(send_reset_password_email_task, to_email, code) + result = await sync_to_async(submit_task)(send_reset_password_email_task, to_email, code) if result is None: logger.warning(f'[ResetPassword] Email send failed (both sync and async): email={to_email}') else: logger.info(f'[ResetPassword] Email sent successfully: email={to_email}') - reset_failures(operation, to_email) + await sync_to_async(reset_failures)(operation, to_email) return create_standardized_response( data={"email_sent": True}, code=ResponseCode.RESET_CODE_SENT, @@ -331,7 +338,7 @@ class ForgotPasswordSendCodeAPIView(APIView): ) except Exception as e: - record_failure(operation, to_email) + await sync_to_async(record_failure)(operation, to_email) return create_standardized_error_response( message=f"邮件发送失败: {str(e)}", code=ResponseCode.EMAIL_SEND_FAILED, @@ -360,7 +367,7 @@ class ForgotPasswordResetAPIView(APIView): ), responses={200: success_response, 400: error_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): from utils.captcha import check_captcha_required, verify_captcha, record_failure, reset_failures to_email = request.data.get('email', None) @@ -377,7 +384,7 @@ class ForgotPasswordResetAPIView(APIView): identifier = to_email operation = 'forgot_reset' - captcha_required = check_captcha_required(operation, identifier) + captcha_required = await sync_to_async(check_captcha_required)(operation, identifier) if captcha_required: captcha_key = request.data.get('captcha_key', None) @@ -389,28 +396,28 @@ class ForgotPasswordResetAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - captcha_result = verify_captcha(captcha_key, captcha_code) + captcha_result = await sync_to_async(verify_captcha)(captcha_key, captcha_code) if captcha_result == 'expired': return create_standardized_error_response( code=ResponseCode.CAPTCHA_EXPIRED, status_code=status.HTTP_400_BAD_REQUEST ) elif captcha_result == 'wrong': - record_failure(operation, identifier) + await sync_to_async(record_failure)(operation, identifier) return create_standardized_error_response( code=ResponseCode.CAPTCHA_ERROR, status_code=status.HTTP_400_BAD_REQUEST ) if new_password != confirm_password: - record_failure(operation, to_email) + await sync_to_async(record_failure)(operation, to_email) return create_standardized_error_response( code=ResponseCode.PASSWORD_MISMATCH, status_code=status.HTTP_400_BAD_REQUEST ) if len(new_password) < 8: - record_failure(operation, to_email) + await sync_to_async(record_failure)(operation, to_email) return create_standardized_error_response( code=ResponseCode.PASSWORD_TOO_SHORT, status_code=status.HTTP_400_BAD_REQUEST @@ -419,51 +426,51 @@ class ForgotPasswordResetAPIView(APIView): has_letter = any(c.isalpha() for c in new_password) has_digit = any(c.isdigit() for c in new_password) if not (has_letter and has_digit): - record_failure(operation, to_email) + await sync_to_async(record_failure)(operation, to_email) return create_standardized_error_response( code=ResponseCode.PASSWORD_TOO_WEAK, status_code=status.HTTP_400_BAD_REQUEST ) try: - vcode = default_cache.get(f"reset_password_{to_email}") + vcode = await sync_to_async(default_cache.get)(f"reset_password_{to_email}") if vcode is None: - record_failure(operation, to_email) + await sync_to_async(record_failure)(operation, to_email) return create_standardized_error_response( code=ResponseCode.VERIFICATION_CODE_EXPIRED, status_code=status.HTTP_400_BAD_REQUEST ) if code != vcode: - record_failure(operation, to_email) + await sync_to_async(record_failure)(operation, to_email) return create_standardized_error_response( code=ResponseCode.VERIFICATION_CODE_ERROR, status_code=status.HTTP_400_BAD_REQUEST ) - user = FUser.objects.filter(email=to_email).first() + user = await FUser.objects.filter(email=to_email).afirst() if user is None: - record_failure(operation, to_email) + await sync_to_async(record_failure)(operation, to_email) return create_standardized_error_response( code=ResponseCode.USER_NOT_FOUND, status_code=status.HTTP_400_BAD_REQUEST ) - user.set_password(new_password) - user.save() + await sync_to_async(user.set_password)(new_password) + await user.asave() - default_cache.delete(f"reset_password_{to_email}") + await sync_to_async(default_cache.delete)(f"reset_password_{to_email}") - reset_failures(operation, to_email) + await sync_to_async(reset_failures)(operation, to_email) return create_standardized_response( code=ResponseCode.PASSWORD_RESET_SUCCESS, status_code=status.HTTP_200_OK ) except Exception as e: - record_failure(operation, to_email) + await sync_to_async(record_failure)(operation, to_email) return create_standardized_error_response( message=f"服务器内部错误: {str(e)}", code=ResponseCode.SERVER_INTERNAL_ERROR, @@ -478,12 +485,12 @@ class UserUpdateAPIView(APIView): operation_description='获取当前登录用户的个人信息', responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user - user_serializer = UserSerializer(user) + user_data = await sync_to_async(_serialize_user)(user) return create_standardized_response( - data={'user': user_serializer.data}, + data={'user': user_data}, code=ResponseCode.SUCCESS, status_code=status.HTTP_200_OK ) @@ -495,7 +502,7 @@ class UserUpdateAPIView(APIView): request_body=UserUpdateSerializer, responses={200: success_response, 400: error_response, 401: unauthorized_response}, ) - def put(self, request): + async def put(self, request): user = request.user serializer = UserUpdateSerializer( user, @@ -504,7 +511,7 @@ class UserUpdateAPIView(APIView): context={'request': request} ) - if not serializer.is_valid(): + if not await sync_to_async(serializer.is_valid)(): errors = serializer.errors first_error = '' for field, msgs in errors.items(): @@ -519,23 +526,23 @@ class UserUpdateAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - updated_user = serializer.save() - user_serializer = UserSerializer(updated_user) + updated_user = await sync_to_async(serializer.save)() + user_data = await sync_to_async(_serialize_user)(updated_user) try: - log_event( + await sync_to_async(log_event)( event_type='profile_update', user=updated_user, description=f'更新了{len(serializer.validated_data)}项个人资料', metadata={'updated_fields': list(serializer.validated_data.keys())}, request=request ) - track_user_action(updated_user, 'profile', count=1) + await sync_to_async(track_user_action)(updated_user, 'profile', count=1) except Exception as e: logging.getLogger(__name__).warning(f'Profile track failed: {e}') return create_standardized_response( - data={'user': user_serializer.data}, + data={'user': user_data}, code=ResponseCode.SUCCESS, status_code=status.HTTP_200_OK ) @@ -547,8 +554,8 @@ class UserUpdateAPIView(APIView): request_body=UserUpdateSerializer, responses={200: success_response, 400: error_response, 401: unauthorized_response}, ) - def patch(self, request): - return self.put(request) + async def patch(self, request): + return await self.put(request) class ChangePasswordAPIView(APIView): @@ -568,12 +575,12 @@ class ChangePasswordAPIView(APIView): ), responses={200: success_response, 400: error_response, 401: unauthorized_response}, ) - def post(self, request): + async def post(self, request): from utils.captcha import check_captcha_required, verify_captcha, record_failure, reset_failures identifier = str(request.user.id) operation = 'change_pwd' - captcha_required = check_captcha_required(operation, identifier) + captcha_required = await sync_to_async(check_captcha_required)(operation, identifier) if captcha_required: captcha_key = request.data.get('captcha_key', None) @@ -585,14 +592,14 @@ class ChangePasswordAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - captcha_result = verify_captcha(captcha_key, captcha_code) + captcha_result = await sync_to_async(verify_captcha)(captcha_key, captcha_code) if captcha_result == 'expired': return create_standardized_error_response( code=ResponseCode.CAPTCHA_EXPIRED, status_code=status.HTTP_400_BAD_REQUEST ) elif captcha_result == 'wrong': - record_failure(operation, identifier) + await sync_to_async(record_failure)(operation, identifier) return create_standardized_error_response( code=ResponseCode.CAPTCHA_ERROR, status_code=status.HTTP_400_BAD_REQUEST @@ -603,7 +610,7 @@ class ChangePasswordAPIView(APIView): context={'request': request} ) - if not serializer.is_valid(): + if not await sync_to_async(serializer.is_valid)(): errors = serializer.errors first_error = '' for field, msgs in errors.items(): @@ -617,7 +624,7 @@ class ChangePasswordAPIView(APIView): break break - record_failure(operation, identifier) + await sync_to_async(record_failure)(operation, identifier) return create_standardized_error_response( data=errors, code=ResponseCode.PARAMETER_ERROR, @@ -625,15 +632,15 @@ class ChangePasswordAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - user = serializer.save() - user_serializer = UserSerializer(user) + user = await sync_to_async(serializer.save)() + user_data = await sync_to_async(_serialize_user)(user) is_new_set = not request.user.has_usable_password() or request.data.get('old_password', '') == '' - reset_failures(operation, identifier) + await sync_to_async(reset_failures)(operation, identifier) return create_standardized_response( data={ - 'user': user_serializer.data, + 'user': user_data, 'is_new_set': is_new_set, }, code=ResponseCode.PASSWORD_SET if is_new_set else ResponseCode.PASSWORD_CHANGED, @@ -658,7 +665,7 @@ class UserLoginAPIView(APIView): ), responses={200: success_response, 400: error_response, 401: unauthorized_response, 403: error_response}, ) - def post(self, request): + async def post(self, request): account = request.data.get('account', None) password = request.data.get('password', None) @@ -670,22 +677,22 @@ class UserLoginAPIView(APIView): ) try: - user = authenticate(username=account, password=password) + user = await sync_to_async(authenticate)(username=account, password=password) if user is not None: if user.is_active: refresh = RefreshToken.for_user(user) - user_serializer = UserSerializer(user) + user_data = await sync_to_async(_serialize_user)(user) response_data = { - 'user': user_serializer.data, + 'user': user_data, 'refresh': str(refresh), 'access': str(refresh.access_token), 'token_type': 'bearer', 'expires_at_timestamp': refresh.access_token.payload['exp'] } - _create_login_record(request, user, 'success') + await sync_to_async(_create_login_record)(request, user, 'success') return create_standardized_response( data=response_data, @@ -693,7 +700,7 @@ class UserLoginAPIView(APIView): status_code=status.HTTP_200_OK ) else: - _create_login_record(request, user, 'failed') + await sync_to_async(_create_login_record)(request, user, 'failed') return create_standardized_error_response( code=ResponseCode.PARAMETER_ERROR, @@ -701,9 +708,9 @@ class UserLoginAPIView(APIView): status_code=status.HTTP_403_FORBIDDEN ) else: - login_user = FUser.objects.filter(username=account).first() or FUser.objects.filter(email=account).first() + login_user = await FUser.objects.filter(username=account).afirst() or await FUser.objects.filter(email=account).afirst() if login_user: - _create_login_record(request, login_user, 'failed') + await sync_to_async(_create_login_record)(request, login_user, 'failed') return create_standardized_error_response( code=ResponseCode.PARAMETER_ERROR, @@ -744,7 +751,7 @@ class LoginView(APIView): ), responses={200: success_response, 400: error_response, 401: unauthorized_response, 403: error_response}, ) - def post(self, request): + async def post(self, request): account = request.data.get('account', None) password = request.data.get('password', None) @@ -757,7 +764,7 @@ class LoginView(APIView): identifier = self._get_identifier(request) operation = 'login' - captcha_required = check_captcha_required(operation, identifier) + captcha_required = await sync_to_async(check_captcha_required)(operation, identifier) if captcha_required: slider_captcha_key = request.data.get('slider_captcha_key', None) @@ -770,9 +777,9 @@ class LoginView(APIView): ) try: - captcha_valid = verify_slider_captcha(slider_captcha_key, int(slider_captcha_x)) + captcha_valid = await sync_to_async(verify_slider_captcha)(slider_captcha_key, int(slider_captcha_x)) if not captcha_valid: - record_failure(operation, identifier) + await sync_to_async(record_failure)(operation, identifier) return create_standardized_error_response( code=ResponseCode.CAPTCHA_ERROR, status_code=status.HTTP_400_BAD_REQUEST @@ -784,23 +791,23 @@ class LoginView(APIView): ) try: - user = authenticate(username=account, password=password) + user = await sync_to_async(authenticate)(username=account, password=password) if user is not None: if user.is_active: - reset_failures(operation, identifier) + await sync_to_async(reset_failures)(operation, identifier) refresh = RefreshToken.for_user(user) - user_serializer = UserSerializer(user) + user_data = await sync_to_async(_serialize_user)(user) response_data = { - 'user': user_serializer.data, + 'user': user_data, 'refresh': str(refresh), 'access': str(refresh.access_token), 'token_type': 'bearer', 'expires_at_timestamp': refresh.access_token.payload['exp'] } - _create_login_record(request, user, 'success') + await sync_to_async(_create_login_record)(request, user, 'success') return create_standardized_response( data=response_data, @@ -808,8 +815,8 @@ class LoginView(APIView): status_code=status.HTTP_200_OK ) else: - record_failure(operation, identifier) - _create_login_record(request, user, 'failed') + await sync_to_async(record_failure)(operation, identifier) + await sync_to_async(_create_login_record)(request, user, 'failed') return create_standardized_error_response( code=ResponseCode.PARAMETER_ERROR, @@ -817,10 +824,10 @@ class LoginView(APIView): status_code=status.HTTP_403_FORBIDDEN ) else: - record_failure(operation, identifier) - login_user = FUser.objects.filter(username=account).first() or FUser.objects.filter(email=account).first() + await sync_to_async(record_failure)(operation, identifier) + login_user = await FUser.objects.filter(username=account).afirst() or await FUser.objects.filter(email=account).afirst() if login_user: - _create_login_record(request, login_user, 'failed') + await sync_to_async(_create_login_record)(request, login_user, 'failed') return create_standardized_error_response( code=ResponseCode.PARAMETER_ERROR, @@ -848,16 +855,17 @@ class PublicProfileAPIView(APIView): ], responses={200: success_response, 404: not_found_response}, ) - def get(self, request, user_id): + async def get(self, request, user_id): try: - user = FUser.objects.get(id=user_id, is_active=True) + user = await FUser.objects.aget(id=user_id, is_active=True) except FUser.DoesNotExist: return Response({'error': '用户不存在'}, status=status.HTTP_404_NOT_FOUND) try: from article.models import Article - article_count = Article.objects.filter(author=user, status='published').count() + article_count = await Article.objects.filter(author=user, status='published').acount() recent_articles = Article.objects.filter(author=user, status='published').order_by('-created_at')[:5] + recent_articles = [a async for a in recent_articles] articles_data = [] for article in recent_articles: articles_data.append({ @@ -945,7 +953,7 @@ class UploadAvatarAPIView(APIView): ], responses={200: success_response, 400: error_response, 401: unauthorized_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): file = request.FILES.get('avatar') if not file: @@ -979,17 +987,17 @@ class UploadAvatarAPIView(APIView): try: filename = f'avatars/{uuid.uuid4().hex}{ext}' - saved_path = default_storage.save(filename, file) + saved_path = await sync_to_async(default_storage.save)(filename, file) user = request.user if user.avatar and user.avatar.name: try: - default_storage.delete(user.avatar.name) + await sync_to_async(default_storage.delete)(user.avatar.name) except Exception: pass user.avatar = saved_path - user.save(update_fields=['avatar']) + await user.asave(update_fields=['avatar']) avatar_url = request.build_absolute_uri(settings.MEDIA_URL + saved_path) @@ -1024,9 +1032,9 @@ class FollowToggleView(APIView): 404: not_found_response, }, ) - def post(self, request, user_id): + async def post(self, request, user_id): try: - target_user = FUser.objects.get(id=user_id, is_active=True) + target_user = await FUser.objects.aget(id=user_id, is_active=True) except FUser.DoesNotExist: return create_standardized_error_response( code=ResponseCode.NOT_FOUND, @@ -1041,19 +1049,19 @@ class FollowToggleView(APIView): status_code=status.HTTP_400_BAD_REQUEST, ) - follow, created = Follow.objects.get_or_create( + follow, created = await Follow.objects.aget_or_create( follower=request.user, following=target_user, ) if not created: - follow.delete() + await follow.adelete() is_following = False else: is_following = True - follower_count = target_user.followers.count() - following_count = request.user.following.count() + follower_count = await target_user.followers.acount() + following_count = await request.user.following.acount() return create_standardized_response( data={ diff --git a/user/views/wallet.py b/user/views/wallet.py index 7b74eee..986ecfa 100644 --- a/user/views/wallet.py +++ b/user/views/wallet.py @@ -1,5 +1,5 @@ from rest_framework import status -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.permissions import IsAuthenticated, IsAdminUser from django.db import transaction from django.db.models import F @@ -9,6 +9,8 @@ from django.views.decorators.cache import never_cache from datetime import timedelta import logging +from asgiref.sync import sync_to_async + from utils.response_codes import ( ResponseCode, create_standardized_response, @@ -33,7 +35,7 @@ class WalletBalanceAPIView(APIView): operation_description='获取当前用户的积分和y币余额', responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user return create_standardized_response( data={ @@ -59,7 +61,7 @@ class WalletTransactionsAPIView(APIView): ], responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user currency_type = request.query_params.get('currency_type', None) page = int(request.query_params.get('page', 1)) @@ -70,11 +72,12 @@ class WalletTransactionsAPIView(APIView): if currency_type: transactions = transactions.filter(currency_type=currency_type) - total = transactions.count() + total = await transactions.acount() start = (page - 1) * page_size end = start + page_size - serializer = PointTransactionSerializer(transactions[start:end], many=True) + rows = [t async for t in transactions[start:end]] + serializer = PointTransactionSerializer(rows, many=True) return create_standardized_response( data={ @@ -107,7 +110,7 @@ class EarnPointsAPIView(APIView): ), responses={200: success_response, 400: error_response, 401: unauthorized_response}, ) - def post(self, request): + async def post(self, request): user_id = request.data.get('user_id') amount = request.data.get('amount', 0) description = request.data.get('description', '') @@ -129,7 +132,7 @@ class EarnPointsAPIView(APIView): ) try: - target_user = FUser.objects.get(pk=user_id) if user_id else request.user + target_user = await FUser.objects.aget(pk=user_id) if user_id else request.user except FUser.DoesNotExist: return create_standardized_error_response( message='目标用户不存在', @@ -137,7 +140,7 @@ class EarnPointsAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - try: + def _earn_points_txn(): with transaction.atomic(): target_user.points = F('points') + amount target_user.save(update_fields=['points']) @@ -152,6 +155,9 @@ class EarnPointsAPIView(APIView): description=description or '获得积分', ) + try: + await sync_to_async(_earn_points_txn)() + return create_standardized_response( data={ 'points': target_user.points, @@ -173,9 +179,9 @@ CHECKIN_POINTS = 30 FULL_WEEK_BONUS = 100 -def track_checkin_task(user): +async def track_checkin_task(user): """ - 签到成功后,自动更新签到任务的进度 + 签到成功后,自动更新签到任务的进度(异步版) """ try: now = timezone.localtime(timezone.now()) @@ -187,12 +193,12 @@ def track_checkin_task(user): is_active=True, ) - for task in checkin_tasks: + async for task in checkin_tasks: period_key = daily_period if task.task_type == 'daily' else ( now.strftime('%Y-W%W') if task.task_type == 'weekly' else 'permanent' ) - progress, created = UserTaskProgress.objects.get_or_create( + progress, created = await UserTaskProgress.objects.aget_or_create( user=user, task=task, period_key=period_key, @@ -211,7 +217,7 @@ def track_checkin_task(user): progress.is_completed = True progress.completed_at = now - progress.save() + await progress.asave() except Exception as e: logger.error(f'签到任务进度更新失败: {e}') @@ -226,7 +232,7 @@ class CheckinStatusAPIView(APIView): operation_description='获取当前用户的签到状态,包括本周签到记录和积分信息', responses={200: success_response, 401: unauthorized_response}, ) - def get(self, request): + async def get(self, request): user = request.user today = timezone.localdate() @@ -241,7 +247,7 @@ class CheckinStatusAPIView(APIView): checkin_date__lte=week_dates[6], ).values_list('checkin_date', flat=True) - checked_dates = set(week_checkins) + checked_dates = set([d async for d in week_checkins]) signed_today = today in checked_dates week_days = [] @@ -281,18 +287,18 @@ class CheckinAPIView(APIView): operation_description='执行每日签到,获得积分奖励,连续签到满一周可获得额外奖励', responses={200: success_response, 400: error_response, 401: unauthorized_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): user = request.user today = timezone.localdate() - if DailyCheckin.objects.filter(user=user, checkin_date=today).exists(): + if await DailyCheckin.objects.filter(user=user, checkin_date=today).aexists(): return create_standardized_error_response( message='今日已签到,请勿重复签到', code=ResponseCode.PARAMETER_ERROR, status_code=status.HTTP_400_BAD_REQUEST ) - try: + def _checkin_txn(): with transaction.atomic(): DailyCheckin.objects.create(user=user, checkin_date=today) @@ -320,8 +326,13 @@ class CheckinAPIView(APIView): description=f'每日签到' + (f'(含满周奖励{FULL_WEEK_BONUS}积分)' if bonus else ''), ) + return week_checkin_count, bonus, total_points + + try: + week_checkin_count, bonus, total_points = await sync_to_async(_checkin_txn)() + # 更新签到任务进度 - track_checkin_task(user) + await track_checkin_task(user) return create_standardized_response( data={ @@ -361,7 +372,7 @@ class EarnCoinsAPIView(APIView): ), responses={200: success_response, 400: error_response, 401: unauthorized_response}, ) - def post(self, request): + async def post(self, request): user_id = request.data.get('user_id') amount = request.data.get('amount', 0) description = request.data.get('description', '') @@ -383,7 +394,7 @@ class EarnCoinsAPIView(APIView): ) try: - target_user = FUser.objects.get(pk=user_id) if user_id else request.user + target_user = await FUser.objects.aget(pk=user_id) if user_id else request.user except FUser.DoesNotExist: return create_standardized_error_response( message='目标用户不存在', @@ -391,7 +402,7 @@ class EarnCoinsAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - try: + def _earn_coins_txn(): with transaction.atomic(): target_user.coins = F('coins') + amount target_user.save(update_fields=['coins']) @@ -406,6 +417,9 @@ class EarnCoinsAPIView(APIView): description=description or '获得y币', ) + try: + await sync_to_async(_earn_coins_txn)() + return create_standardized_response( data={ 'points': target_user.points, @@ -440,7 +454,7 @@ class SpendPointsAPIView(APIView): ), responses={200: success_response, 400: error_response, 401: unauthorized_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): amount = request.data.get('amount', 0) description = request.data.get('description', '') @@ -460,16 +474,12 @@ class SpendPointsAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - try: + def _spend_points_txn(): with transaction.atomic(): user = FUser.objects.select_for_update().get(pk=request.user.pk) if user.points < amount: - return create_standardized_error_response( - message='积分余额不足', - code=ResponseCode.PARAMETER_ERROR, - status_code=status.HTTP_400_BAD_REQUEST - ) + return None user.points = F('points') - amount user.save(update_fields=['points']) @@ -483,6 +493,17 @@ class SpendPointsAPIView(APIView): balance_after=user.points, description=description or '消费积分', ) + return user + + try: + user = await sync_to_async(_spend_points_txn)() + + if user is None: + return create_standardized_error_response( + message='积分余额不足', + code=ResponseCode.PARAMETER_ERROR, + status_code=status.HTTP_400_BAD_REQUEST + ) return create_standardized_response( data={ @@ -518,7 +539,7 @@ class SpendCoinsAPIView(APIView): ), responses={200: success_response, 400: error_response, 401: unauthorized_response, 500: error_response}, ) - def post(self, request): + async def post(self, request): amount = request.data.get('amount', 0) description = request.data.get('description', '') @@ -538,16 +559,12 @@ class SpendCoinsAPIView(APIView): status_code=status.HTTP_400_BAD_REQUEST ) - try: + def _spend_coins_txn(): with transaction.atomic(): user = FUser.objects.select_for_update().get(pk=request.user.pk) if user.coins < amount: - return create_standardized_error_response( - message='y币余额不足', - code=ResponseCode.PARAMETER_ERROR, - status_code=status.HTTP_400_BAD_REQUEST - ) + return None user.coins = F('coins') - amount user.save(update_fields=['coins']) @@ -561,6 +578,17 @@ class SpendCoinsAPIView(APIView): balance_after=user.coins, description=description or '消费y币', ) + return user + + try: + user = await sync_to_async(_spend_coins_txn)() + + if user is None: + return create_standardized_error_response( + message='y币余额不足', + code=ResponseCode.PARAMETER_ERROR, + status_code=status.HTTP_400_BAD_REQUEST + ) return create_standardized_response( data={ diff --git a/utils/async_cache.py b/utils/async_cache.py new file mode 100644 index 0000000..203faee --- /dev/null +++ b/utils/async_cache.py @@ -0,0 +1,114 @@ +""" +异步 Redis 缓存封装 —— 100% 异步化基建 +======================================== +Django 的 cache 框架(django-redis)没有异步 API,本模块基于 redis.asyncio 直连, +复用 settings 中的 REDIS 配置,提供与 django cache 对等的异步接口。 + +用法: + from utils.async_cache import aget_cache, aset_cache, adelete_cache + + value = await aget_cache('key') + await aset_cache('key', obj, timeout=3600) + await adelete_cache('key') + +序列化:msgpack(快速紧凑)+ JSON 兜底;与 django-redis 缓存池互不干扰(独立 key 前缀)。 +""" +import asyncio +import json +import logging +from django.conf import settings +from django.core.cache import caches +from asgiref.sync import sync_to_async +import msgpack + +logger = logging.getLogger(__name__) + +_cache_pool = None +_lock = asyncio.Lock() + + +async def _get_pool(): + """懒初始化 redis.asyncio 连接池(复用 settings.REDIS 配置)""" + global _cache_pool + async with _lock: + if _cache_pool is None: + import redis.asyncio as aioredis + _cache_pool = aioredis.Redis( + host=settings.REDIS_HOST, + port=int(settings.REDIS_PORT), + db=int(getattr(settings, 'REDIS_DB', 0)), + password=getattr(settings, 'REDIS_PASSWORD', '') or None, + decode_responses=False, + socket_connect_timeout=5, + socket_timeout=5, + max_connections=50, + ) + return _cache_pool + + +def _pack(value): + try: + return b'async-cache:' + msgpack.packb(value, use_bin_type=True, default=str) + except (TypeError, ValueError): + return b'async-cache-json:' + json.dumps(value, ensure_ascii=False, default=str).encode() + + +def _unpack(raw): + if raw is None: + return None + if raw.startswith(b'async-cache-json:'): + return json.loads(raw[17:].decode()) + return msgpack.unpackb(raw[12:], raw=False) + + +async def aget_cache(key, default=None): + try: + pool = await _get_pool() + raw = await pool.get(f'async:{key}') + if raw is None: + return default + return _unpack(raw) + except Exception: + logger.warning('aget_cache failed for %s, falling back to sync cache', key) + try: + return await sync_to_async(caches['default'].get)(key, default) + except Exception: + return default + + +async def aset_cache(key, value, timeout=300): + try: + pool = await _get_pool() + await pool.set(f'async:{key}', _pack(value), ex=timeout) + return True + except Exception: + logger.warning('aset_cache failed for %s, falling back to sync cache', key) + try: + await sync_to_async(caches['default'].set)(key, value, timeout) + return True + except Exception: + return False + + +async def adelete_cache(key): + try: + pool = await _get_pool() + await pool.delete(f'async:{key}') + return True + except Exception: + logger.warning('adelete_cache failed for %s', key) + try: + await sync_to_async(caches['default'].delete)(key) + return True + except Exception: + return False + + +async def aget_or_set(key, factory, timeout=300): + """原子性不保证(与 django cache 用法一致),value 为 await factory() 结果""" + v = await aget_cache(key) + if v is not None: + return v + v = await factory() if asyncio.iscoroutinefunction(factory) else factory() + await aset_cache(key, v, timeout) + return v diff --git a/weather/views.py b/weather/views.py index b6a5e17..b56fb01 100644 --- a/weather/views.py +++ b/weather/views.py @@ -1,8 +1,11 @@ -import requests +import asyncio from datetime import datetime +import aiohttp +from asgiref.sync import sync_to_async + from rest_framework.permissions import AllowAny -from rest_framework.views import APIView +from adrf.views import APIView from rest_framework.response import Response from rest_framework import status from drf_yasg.utils import swagger_auto_schema @@ -81,7 +84,7 @@ class WeatherView(APIView): ], responses={200: success_response, 400: error_response, 404: error_response, 502: error_response} ) - def get(self, request): + async def get(self, request): city = request.GET.get('city', '').strip() unit = request.GET.get('unit', 'celsius').strip() lang = request.GET.get('lang', 'zh_cn').strip() @@ -98,7 +101,7 @@ class WeatherView(APIView): status=status.HTTP_400_BAD_REQUEST ) - return self._fetch_weather(city, unit, lang) + return await self._fetch_weather(city, unit, lang) @swagger_auto_schema( tags=['Weather'], @@ -115,7 +118,7 @@ class WeatherView(APIView): ), responses={200: success_response, 400: error_response, 404: error_response, 502: error_response} ) - def post(self, request): + async def post(self, request): city = request.data.get('city', '').strip() if isinstance(request.data, dict) else '' unit = request.data.get('unit', 'celsius').strip() if isinstance(request.data, dict) else 'celsius' lang = request.data.get('lang', 'zh_cn').strip() if isinstance(request.data, dict) else 'zh_cn' @@ -132,62 +135,62 @@ class WeatherView(APIView): status=status.HTTP_400_BAD_REQUEST ) - return self._fetch_weather(city, unit, lang) + return await self._fetch_weather(city, unit, lang) - def _fetch_weather(self, city, unit, lang): + async def _fetch_weather(self, city, unit, lang): """ - 调用 Open-Meteo API 获取天气数据 + 调用 Open-Meteo API 获取天气数据(aiohttp 全异步) """ + timeout = aiohttp.ClientTimeout(total=10) try: - # 第一步:通过地理编码API获取城市坐标 - geo_params = { - 'name': city, - 'count': 1, - 'language': 'zh' if lang == 'zh_cn' else 'en', - 'format': 'json' - } - geo_response = requests.get(GEOCODING_URL, params=geo_params, timeout=10) + async with aiohttp.ClientSession(timeout=timeout) as session: + # 第一步:通过地理编码API获取城市坐标 + geo_params = { + 'name': city, + 'count': 1, + 'language': 'zh' if lang == 'zh_cn' else 'en', + 'format': 'json' + } + async with session.get(GEOCODING_URL, params=geo_params) as geo_response: + if geo_response.status != 200: + return Response( + {"code": 502, "message": "地理编码服务异常,请稍后重试", "data": None}, + status=status.HTTP_502_BAD_GATEWAY + ) + geo_data = await geo_response.json() - if geo_response.status_code != 200: - return Response( - {"code": 502, "message": "地理编码服务异常,请稍后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY - ) + results = geo_data.get('results', []) - geo_data = geo_response.json() - results = geo_data.get('results', []) + if not results: + return Response( + {"code": 404, "message": f"未找到城市:{city},请检查城市名称是否正确", "data": None}, + status=status.HTTP_404_NOT_FOUND + ) - if not results: - return Response( - {"code": 404, "message": f"未找到城市:{city},请检查城市名称是否正确", "data": None}, - status=status.HTTP_404_NOT_FOUND - ) + location = results[0] + latitude = location.get('latitude') + longitude = location.get('longitude') + city_name = location.get('name', city) + country = location.get('country', '') - location = results[0] - latitude = location.get('latitude') - longitude = location.get('longitude') - city_name = location.get('name', city) - country = location.get('country', '') + # 第二步:获取天气数据 + temp_unit = 'fahrenheit' if unit == 'fahrenheit' else 'celsius' + weather_params = { + 'latitude': latitude, + 'longitude': longitude, + 'current': 'temperature_2m,relative_humidity_2m,weather_code,wind_speed_10m,pressure_msl', + 'timezone': 'auto', + 'temperature_unit': temp_unit + } - # 第二步:获取天气数据 - temp_unit = 'fahrenheit' if unit == 'fahrenheit' else 'celsius' - weather_params = { - 'latitude': latitude, - 'longitude': longitude, - 'current': 'temperature_2m,relative_humidity_2m,weather_code,wind_speed_10m,pressure_msl', - 'timezone': 'auto', - 'temperature_unit': temp_unit - } + async with session.get(FORECAST_URL, params=weather_params) as weather_response: + if weather_response.status != 200: + return Response( + {"code": 502, "message": "天气服务异常,请稍后重试", "data": None}, + status=status.HTTP_502_BAD_GATEWAY + ) + weather_data = await weather_response.json() - weather_response = requests.get(FORECAST_URL, params=weather_params, timeout=10) - - if weather_response.status_code != 200: - return Response( - {"code": 502, "message": "天气服务异常,请稍后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY - ) - - weather_data = weather_response.json() current = weather_data.get('current', {}) weather_code = current.get('weather_code', 0) @@ -210,16 +213,11 @@ class WeatherView(APIView): status=status.HTTP_200_OK ) - except requests.exceptions.Timeout: + except (aiohttp.ClientError, asyncio.TimeoutError): return Response( {"code": 502, "message": "请求超时,请稍后重试", "data": None}, status=status.HTTP_502_BAD_GATEWAY ) - except requests.exceptions.ConnectionError: - return Response( - {"code": 502, "message": "网络连接异常,请检查网络后重试", "data": None}, - status=status.HTTP_502_BAD_GATEWAY - ) except Exception as e: return Response( {"code": 500, "message": f"服务器内部错误:{str(e)}", "data": None},