baseline: 批次A-D 成果 + membership 半成品(测试红)

This commit is contained in:
agent
2026-09-11 23:11:35 +08:00
commit b3f3095d53
311 changed files with 40540 additions and 0 deletions
+152
View File
@@ -0,0 +1,152 @@
"""公共视图组件: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)