233 lines
8.9 KiB
Python
233 lines
8.9 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
|
||
|
||
|
||
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)
|