forked from erp-dev/erp
fix: merchant not required at customer create api
This commit is contained in:
@@ -136,6 +136,41 @@ class CustomerSerializer(BaseSerializer):
|
||||
required=False,
|
||||
)
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.validators = [
|
||||
validator
|
||||
for validator in self.validators
|
||||
if not (
|
||||
hasattr(validator, 'fields')
|
||||
and tuple(validator.fields) == ('merchant', 'name')
|
||||
)
|
||||
]
|
||||
|
||||
def _get_request_merchant(self):
|
||||
request = self.context.get('request')
|
||||
try:
|
||||
return request.user.employee.merchant
|
||||
except AttributeError:
|
||||
return None
|
||||
|
||||
def validate(self, attrs):
|
||||
attrs = super().validate(attrs)
|
||||
|
||||
merchant = self._get_request_merchant() or attrs.get('merchant') or getattr(self.instance, 'merchant', None)
|
||||
name = attrs.get('name', getattr(self.instance, 'name', None))
|
||||
|
||||
if merchant and name:
|
||||
queryset = basic_models.Customer.objects.filter(merchant=merchant, name=name)
|
||||
if self.instance is not None:
|
||||
queryset = queryset.exclude(pk=self.instance.pk)
|
||||
if queryset.exists():
|
||||
raise serializers.ValidationError({
|
||||
'name': '该商户下已存在同名客户。'
|
||||
})
|
||||
|
||||
return attrs
|
||||
|
||||
class Meta:
|
||||
model = basic_models.Customer
|
||||
fields = '__all__'
|
||||
|
||||
@@ -734,7 +734,7 @@ class SupplierAPITestCase(TestCase):
|
||||
|
||||
|
||||
class CustomerAPITestCase(TestCase):
|
||||
"""测试客户 API(支持 name 查询参数)"""
|
||||
"""测试客户 API(支持 name 和 visible_employee_name 查询参数)"""
|
||||
|
||||
def setUp(self):
|
||||
self.merchant = Merchant.objects.create(
|
||||
@@ -767,6 +767,26 @@ class CustomerAPITestCase(TestCase):
|
||||
response = self.client.post('/api/backend/customers/', data, format='json')
|
||||
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||
self.assertEqual(response.data['name'], '测试客户')
|
||||
customer = Customer.objects.get(id=response.data['id'])
|
||||
self.assertEqual(customer.merchant, self.merchant)
|
||||
self.assertEqual(customer.created_by, self.employee)
|
||||
|
||||
def test_create_customer_ignores_explicit_merchant_param(self):
|
||||
"""测试创建客户时即使前端显式传 merchant,也以当前用户商户为准"""
|
||||
other_merchant = Merchant.objects.create(
|
||||
name='其他商户',
|
||||
type=MerchantTypeEnum.STORE,
|
||||
)
|
||||
data = {
|
||||
'name': '兼容客户',
|
||||
'merchant': other_merchant.id,
|
||||
}
|
||||
|
||||
response = self.client.post('/api/backend/customers/', data, format='json')
|
||||
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||
customer = Customer.objects.get(id=response.data['id'])
|
||||
self.assertEqual(customer.merchant, self.merchant)
|
||||
self.assertEqual(customer.created_by, self.employee)
|
||||
|
||||
def test_filter_customers_by_name(self):
|
||||
"""测试通过 name 查询参数过滤客户"""
|
||||
@@ -779,6 +799,53 @@ class CustomerAPITestCase(TestCase):
|
||||
self.assertEqual(len(results), 1)
|
||||
self.assertEqual(results[0]['name'], '客户A')
|
||||
|
||||
def test_filter_customers_by_visible_employee_name(self):
|
||||
"""测试通过 visible_employee_name 查询参数过滤客户"""
|
||||
visible_employee = Employee.objects.create(
|
||||
merchant=self.merchant,
|
||||
name='李四业务',
|
||||
)
|
||||
other_employee = Employee.objects.create(
|
||||
merchant=self.merchant,
|
||||
name='王五跟单',
|
||||
)
|
||||
customer_a = Customer.objects.create(
|
||||
merchant=self.merchant,
|
||||
name='客户A',
|
||||
created_by=self.employee,
|
||||
)
|
||||
customer_b = Customer.objects.create(
|
||||
merchant=self.merchant,
|
||||
name='客户B',
|
||||
created_by=self.employee,
|
||||
)
|
||||
customer_a.visible_employees.add(visible_employee)
|
||||
customer_b.visible_employees.add(other_employee)
|
||||
|
||||
response = self.client.get('/api/backend/customers/?visible_employee_name=李四')
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
results = response.data['results'] if isinstance(response.data, dict) else response.data
|
||||
self.assertEqual(len(results), 1)
|
||||
self.assertEqual(results[0]['id'], customer_a.id)
|
||||
|
||||
def test_filter_customers_by_visible_employee_name_returns_distinct_results(self):
|
||||
"""测试 visible_employee_name 过滤不会因多名匹配员工导致客户重复"""
|
||||
customer = Customer.objects.create(
|
||||
merchant=self.merchant,
|
||||
name='客户A',
|
||||
created_by=self.employee,
|
||||
)
|
||||
customer.visible_employees.add(
|
||||
Employee.objects.create(merchant=self.merchant, name='销售张三'),
|
||||
Employee.objects.create(merchant=self.merchant, name='客服张三'),
|
||||
)
|
||||
|
||||
response = self.client.get('/api/backend/customers/?visible_employee_name=张三')
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
results = response.data['results'] if isinstance(response.data, dict) else response.data
|
||||
self.assertEqual(len(results), 1)
|
||||
self.assertEqual(results[0]['id'], customer.id)
|
||||
|
||||
|
||||
class VehicleTypeAPITestCase(TestCase):
|
||||
"""测试车辆类型 API(支持 name 查询参数)"""
|
||||
|
||||
@@ -172,6 +172,15 @@ class CustomerViewSet(BaseViewSet, BasicInfoFilterMixin):
|
||||
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 filter_queryset(self, queryset):
|
||||
qs = super().filter_queryset(queryset)
|
||||
visible_employee_name = self.request.query_params.get('visible_employee_name')
|
||||
if visible_employee_name:
|
||||
qs = qs.filter(
|
||||
visible_employees__name__icontains=visible_employee_name,
|
||||
).distinct()
|
||||
return qs
|
||||
|
||||
def has_permission(self, request, view):
|
||||
if request.user.is_superuser or request.user.has_perm('basic_info.view_all_customer'):
|
||||
|
||||
Reference in New Issue
Block a user