153 lines
5.6 KiB
Python
153 lines
5.6 KiB
Python
"""公共视图组件: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)
|