baseline: 批次A-D 成果 + membership 半成品(测试红)
This commit is contained in:
@@ -0,0 +1,152 @@
|
||||
"""公共视图组件:async 分页 + 多租户基类 + 创建注入。
|
||||
|
||||
供 apps/catalog, apps/partner 等共享。
|
||||
"""
|
||||
|
||||
from adrf.viewsets import ModelViewSet
|
||||
from adrf.shortcuts import aget_object_or_404
|
||||
from asgiref.sync import sync_to_async
|
||||
from django.db.models import Q
|
||||
|
||||
|
||||
def resolve_tenant(code):
|
||||
from apps.core.models import Tenant
|
||||
try:
|
||||
return Tenant.objects.get(code=code, is_active=True)
|
||||
except Tenant.DoesNotExist:
|
||||
return None
|
||||
|
||||
|
||||
class StandardAsyncPagination:
|
||||
"""手写 async 分页器(兼容 coroutine 或 queryset 入参)。"""
|
||||
|
||||
page_size = 20
|
||||
page_size_query_param = "page_size"
|
||||
max_page_size = 200
|
||||
|
||||
async def paginate_queryset(self, queryset, request, view=None):
|
||||
import asyncio
|
||||
|
||||
if asyncio.iscoroutine(queryset):
|
||||
queryset = await queryset
|
||||
|
||||
try:
|
||||
page_size = int(request.query_params.get(self.page_size_query_param, self.page_size))
|
||||
except (TypeError, ValueError):
|
||||
page_size = self.page_size
|
||||
if page_size < 1: # 0 / 负数会导致空切片或异常
|
||||
page_size = self.page_size
|
||||
page_size = min(page_size, self.max_page_size)
|
||||
|
||||
try:
|
||||
page_number = int(request.query_params.get("page", 1))
|
||||
except (TypeError, ValueError):
|
||||
page_number = 1
|
||||
# 页码必须 ≥1:page=0/-1 会导致负索引切片(Django 会抛 ValueError → 500)
|
||||
if page_number < 1:
|
||||
page_number = 1
|
||||
|
||||
self.total = await sync_to_async(queryset.count)()
|
||||
start = (page_number - 1) * page_size
|
||||
end = start + page_size
|
||||
sliced = await sync_to_async(lambda: list(queryset[start:end]))()
|
||||
return sliced
|
||||
|
||||
def get_paginated_response(self, data):
|
||||
from rest_framework.response import Response
|
||||
return Response({
|
||||
"count": getattr(self, "total", 0),
|
||||
"next": None,
|
||||
"previous": None,
|
||||
"results": data,
|
||||
})
|
||||
|
||||
|
||||
class BaseTenantViewSet(ModelViewSet):
|
||||
"""通用:自动按租户过滤 + 软删除过滤 + 自定义 search + create 注入 tenant。"""
|
||||
|
||||
pagination_class = StandardAsyncPagination
|
||||
|
||||
async def get_tenant(self):
|
||||
tenant = getattr(self.request, "tenant_obj", None)
|
||||
if tenant is not None:
|
||||
return tenant
|
||||
from django.conf import settings
|
||||
|
||||
code = self.request.META.get(
|
||||
"HTTP_X_TENANT_ID", settings.TENANT_DEFAULT
|
||||
)
|
||||
return await sync_to_async(resolve_tenant)(code)
|
||||
|
||||
# 子类可声明关联预取,避免列表接口 N+1
|
||||
# (实测:未声明时 50 张销售单产生 311 条 SQL,声明后降到 3 条)
|
||||
select_related_fields: tuple = ()
|
||||
prefetch_related_fields: tuple = ()
|
||||
|
||||
async def get_queryset(self):
|
||||
tenant = await self.get_tenant()
|
||||
if tenant is None:
|
||||
return self.model.objects.none()
|
||||
qs = self.model.objects.active().filter(tenant=tenant)
|
||||
|
||||
# 关联预取:消除列表页 N+1(嵌套的 prefetch 用 "__" 路径)
|
||||
if self.select_related_fields:
|
||||
qs = qs.select_related(*self.select_related_fields)
|
||||
if self.prefetch_related_fields:
|
||||
qs = qs.prefetch_related(*self.prefetch_related_fields)
|
||||
|
||||
search = self.request.query_params.get("search", "").strip()
|
||||
if search and getattr(self, "search_fields", None):
|
||||
q = Q()
|
||||
for f in self.search_fields:
|
||||
q |= Q(**{f"{f}__icontains": search})
|
||||
qs = qs.filter(q)
|
||||
return qs
|
||||
|
||||
async def aget_object(self):
|
||||
"""修复 adrf aget_object 不支持 async get_queryset 的问题。
|
||||
|
||||
adrf 原实现同步调用 self.get_queryset(),拿到 coroutine 后在
|
||||
aget_object_or_404 里 TypeError → 被吞成 Http404,导致全项目
|
||||
所有 detail 路由(retrieve/PUT/PATCH/DELETE)404。
|
||||
"""
|
||||
queryset = await self.get_queryset()
|
||||
queryset = await self.afilter_queryset(queryset)
|
||||
lookup_url_kwarg = self.lookup_url_kwarg or self.lookup_field
|
||||
filter_kwargs = {self.lookup_field: self.kwargs[lookup_url_kwarg]}
|
||||
obj = await aget_object_or_404(queryset, **filter_kwargs)
|
||||
await sync_to_async(self.check_object_permissions)(self.request, obj)
|
||||
return obj
|
||||
|
||||
async def acreate(self, request, *args, **kwargs):
|
||||
from rest_framework.exceptions import ValidationError
|
||||
from rest_framework import status
|
||||
from rest_framework.response import Response
|
||||
|
||||
tenant = await self.get_tenant()
|
||||
if tenant is None:
|
||||
raise ValidationError({"tenant": "无法识别租户"})
|
||||
|
||||
# 套餐配额拦截(批次 C1):声明了 quota_kind 的 ViewSet 在写入前校验
|
||||
quota_kind = getattr(self, "quota_kind", None)
|
||||
if quota_kind:
|
||||
from apps.billing import quota as billing_quota
|
||||
|
||||
def _check():
|
||||
return billing_quota.check_and_count(tenant, quota_kind, delta=1)
|
||||
|
||||
try:
|
||||
await sync_to_async(_check)()
|
||||
except billing_quota.QuotaExceeded as exc:
|
||||
return Response(exc.as_dict(), status=status.HTTP_403_FORBIDDEN)
|
||||
|
||||
serializer = self.get_serializer(data=request.data)
|
||||
serializer.is_valid(raise_exception=True)
|
||||
validated = serializer.validated_data
|
||||
validated["tenant"] = tenant
|
||||
if request.user.is_authenticated:
|
||||
validated["created_by"] = request.user
|
||||
validated["updated_by"] = request.user
|
||||
|
||||
await sync_to_async(serializer.save)()
|
||||
return Response(serializer.data, status=status.HTTP_201_CREATED)
|
||||
Reference in New Issue
Block a user