"""公共视图组件: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)