forked from erp-dev/erp
221 lines
7.6 KiB
Python
221 lines
7.6 KiB
Python
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 rest_framework.decorators import action
|
||
from rest_framework.response import Response
|
||
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 field and model:
|
||
self.filterset_class = self._build_filterset_class(model, field)
|
||
|
||
qs = super().filter_queryset(queryset)
|
||
|
||
# 统一按创建时间倒序排序(如果模型存在 created_at 字段)
|
||
model = getattr(qs, 'model', model)
|
||
if model:
|
||
field_names = {f.name for f in model._meta.get_fields() if hasattr(f, 'name')}
|
||
if 'created_at' in field_names:
|
||
return qs.order_by('-created_at')
|
||
return qs
|
||
|
||
|
||
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')
|
||
name = self.request.query_params.get('name')
|
||
if group:
|
||
qs = qs.filter(group=group)
|
||
if name:
|
||
qs = qs.filter(name__icontains=name)
|
||
return qs
|
||
|
||
@action(detail=False, methods=['get'], url_path='groups')
|
||
def groups(self, request):
|
||
groups = serializers.basic_models.QuickInput.objects.values_list('group', flat=True).distinct()
|
||
return Response(groups)
|
||
|
||
|
||
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 filter_queryset(self, queryset):
|
||
qs = super().filter_queryset(queryset)
|
||
name = self.request.query_params.get('name')
|
||
if name:
|
||
qs = qs.filter(name__icontains=name)
|
||
return qs
|
||
|
||
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("无权限访问该对象")
|