Files

233 lines
8.9 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""公共视图组件: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)