from rest_framework import viewsets from rest_framework.exceptions import PermissionDenied from rest_framework.pagination import LimitOffsetPagination from rest_framework.permissions import IsAuthenticated, DjangoModelPermissions from django_filters.rest_framework import DjangoFilterBackend from django_filters import rest_framework as dj_filters from . import serializers class BasicInfoFilterMixin: """ 提供基于 name/title 的基础过滤能力,自动判定字段,避免重复配置。 """ filter_backends = [DjangoFilterBackend] _auto_filterset_cache = {} def _detect_filter_field(self): """ 返回 field_name 或 None。 """ model = getattr(getattr(self, 'queryset', None), 'model', None) if model is None: try: qs = super().get_queryset() model = getattr(qs, 'model', None) except Exception: return None field_names = {f.name for f in model._meta.get_fields() if hasattr(f, 'name')} # 优先级:name > title > driver_name if 'name' in field_names: return 'name' if 'title' in field_names: return 'title' if 'driver_name' in field_names: return 'driver_name' return None def _build_filterset_class(self, model, field): """ 动态构建 FilterSet,参数名直接使用字段名,lookup 使用 icontains。 """ cache_key = (model, field) if cache_key in self._auto_filterset_cache: return self._auto_filterset_cache[cache_key] meta_cls = type( 'Meta', (), { 'model': model, 'fields': [field], }, ) auto_filter_cls = type( f'{model.__name__}AutoFilter', (dj_filters.FilterSet,), { field: dj_filters.CharFilter(field_name=field, lookup_expr='icontains'), 'Meta': meta_cls, }, ) self._auto_filterset_cache[cache_key] = auto_filter_cls return auto_filter_cls def filter_queryset(self, queryset): field = self._detect_filter_field() model = getattr(getattr(self, 'queryset', None), 'model', None) if not field or not model: return super().filter_queryset(queryset) self.filterset_class = self._build_filterset_class(model, field) return super().filter_queryset(queryset) class BaseViewSet(BasicInfoFilterMixin, viewsets.ModelViewSet): pagination_class = LimitOffsetPagination def get_queryset(self): qs = super().get_queryset() # 如果是超级用户,返回所有数据 # TODO: 测试结束后要取注释 # if self.request.user.is_superuser: # return qs return qs.filter(merchant=self.request.user.employee.merchant) def perform_create(self, serializer): try: merchant = self.request.user.employee.merchant serializer.save(merchant_id=merchant.id) except AttributeError: raise PermissionDenied("无权限创建该对象") class QuickInputViewSet(BasicInfoFilterMixin, viewsets.ModelViewSet): queryset = serializers.basic_models.QuickInput.objects.all() serializer_class = serializers.QuickInputSerializer pagination_class = LimitOffsetPagination def filter_queryset(self, queryset): qs = super().filter_queryset(queryset) group = self.request.query_params.get('group') if group: qs = qs.filter(group=group) return qs class ProductViewSet(BaseViewSet): queryset = serializers.basic_models.Product.objects serializer_class = serializers.ProductSerializer class WareHouseViewSet(BaseViewSet): queryset = serializers.basic_models.WareHouse.objects serializer_class = serializers.WareHouseSerializer class ProductCategoryViewSet(BaseViewSet): queryset = serializers.basic_models.ProductCategory.objects serializer_class = serializers.ProductCategorySerializer class SupplierViewSet(BaseViewSet): queryset = serializers.basic_models.Supplier.objects serializer_class = serializers.SupplierSerializer class EmployeeViewSet(BaseViewSet): queryset = serializers.basic_models.Employee.objects serializer_class = serializers.EmployeeSerializer class EmployeeTypeViewSet(BaseViewSet): queryset = serializers.basic_models.EmployeeType.objects serializer_class = serializers.EmployeeTypeSerializer class CustomerViewSet(BaseViewSet): queryset = serializers.basic_models.Customer.objects serializer_class = serializers.CustomerSerializer permission_classes = [IsAuthenticated, DjangoModelPermissions] def perform_create(self, serializer): merchant = self.request.user.employee.merchant employee = self.request.user.employee serializer.save(merchant=merchant, created_by=employee) def can_view_all(self) -> bool: return self.request.user.is_superuser or self.request.user.has_perm('basic_info.view_all_customers') def get_queryset(self): qs = super().get_queryset() # 应用可见性过滤 from basic_info.services import CustomerVisibilityService if self.can_view_all(): return qs.filter(merchant=self.request.user.employee.merchant) return CustomerVisibilityService.filter_customers_for_employee(qs, self.request.user) def has_permission(self, request, view): if request.user.is_superuser or request.user.has_perm('basic_info.view_all_customer'): return True return super().has_permission(request, view) class VehicleTypeViewSet(BaseViewSet): queryset = serializers.basic_models.VehicleType.objects serializer_class = serializers.VehicleTypeSerializer class BankAccountViewSet(BaseViewSet): queryset = serializers.basic_models.BankAccount.objects serializer_class = serializers.BankAccountSerializer class DeviceInfoViewSet(BaseViewSet): queryset = serializers.basic_models.DeviceInfo.objects serializer_class = serializers.DeviceInfoSerializer class VehicleTransportRecordViewSet(BaseViewSet): queryset = serializers.basic_models.VehicleTransportRecord.objects serializer_class = serializers.VehicleTransportRecordSerializer class UserProfileViewSet(BaseViewSet): queryset = serializers.basic_models.UserProfile.objects serializer_class = serializers.UserProfileSerializer def filter_queryset(self, queryset): try: merchant = self.request.user.employee.merchant return super().filter_queryset(queryset).filter(merchant=merchant) except AttributeError: raise PermissionDenied("无权限访问该对象")