Files
dealerhub/backend/apps/core/viewset.py
T

153 lines
5.6 KiB
Python
Raw 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
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)