"""公共视图组件: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 def iter_unique_together(model): """把 `Model._meta.unique_together` 归一化成字段名元组列表。 Django 允许 `unique_together = ("tenant", "code")`(单组合简写) 与 `(("tenant", "code"), (...))`(多组合)两种写法,这里统一成后者。 """ raw = getattr(model._meta, "unique_together", None) or () if raw and isinstance(raw[0], str): return [tuple(raw)] return [tuple(group) for group in raw] def build_unique_lookup(model, fields, tenant, validated_data): """为一组 unique_together 字段构造预检 lookup。 返回 `(lookup, missing)`:lookup 含 tenant + 其余字段值; missing 为 validated_data 里缺失的字段名(缺字段则跳过预检)。 FK 传入对象时取 pk;软删除模型调用方需另加 is_deleted=False。 """ lookup = {"tenant": tenant} missing = [] for f in fields: if f == "tenant": continue val = validated_data.get(f) if val is None: missing.append(f) continue lookup[f] = getattr(val, "pk", val) return lookup, missing def unique_conflict_message(rest_fields): human = " / ".join(rest_fields) return f"{human} 已存在,请更换" 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 explicit = self.request.META.get("HTTP_X_TENANT_ID") code = explicit if explicit else settings.TENANT_DEFAULT tenant = await sync_to_async(resolve_tenant)(code) if tenant is None and explicit: # The caller explicitly named a tenant that does not exist or is # inactive. Saying "no data" (empty 200) hides the mistake; surface # it as a parameter error so callers can tell 400 from 403. from rest_framework.exceptions import ValidationError raise ValidationError({"tenant": "无法识别租户"}) return tenant # 子类可声明关联预取,避免列表接口 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) # DRF 的 is_valid 会跑 FK(PrimaryKeyRelatedField)的 queryset.get(), # 在 async 上下文里同步查库会抛 SynchronousOnlyOperation(实测: # POST /finance/receipts/ 带 customer 即 500)。包进线程池。 await sync_to_async(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 # P0-4 租户感知的唯一性预检:DRF 因 tenant 不在 Meta.fields 而静默丢弃 # UniqueTogetherValidator(28 模型全中),这里补齐友好 400。 # 这是预检 + DB 约束双保险:预检负责友好报错,DB 约束负责并发兜底。 def _check_unique(): for group in iter_unique_together(self.model): if "tenant" not in group: continue rest = [f for f in group if f != "tenant"] if not rest: continue lookup, missing = build_unique_lookup( self.model, group, tenant, validated ) if missing: continue qs = self.model.objects.filter(**lookup) if hasattr(self.model, "is_deleted"): qs = qs.filter(is_deleted=False) if qs.exists(): raise ValidationError( {rest[0]: unique_conflict_message(rest)} ) await sync_to_async(_check_unique)() # 并发兜底:预检通过后仍可能撞唯一约束(双写竞态),转 400 而非 500。 from django.db import IntegrityError try: await sync_to_async(serializer.save)() except IntegrityError: raise ValidationError( {"code": "编码已存在,请更换(并发写入冲突)"} ) return Response(serializer.data, status=status.HTTP_201_CREATED)