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,
|
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:
|
class Meta:
|
||||||
model = basic_models.Customer
|
model = basic_models.Customer
|
||||||
fields = '__all__'
|
fields = '__all__'
|
||||||
|
|||||||
@@ -734,7 +734,7 @@ class SupplierAPITestCase(TestCase):
|
|||||||
|
|
||||||
|
|
||||||
class CustomerAPITestCase(TestCase):
|
class CustomerAPITestCase(TestCase):
|
||||||
"""测试客户 API(支持 name 查询参数)"""
|
"""测试客户 API(支持 name 和 visible_employee_name 查询参数)"""
|
||||||
|
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
self.merchant = Merchant.objects.create(
|
self.merchant = Merchant.objects.create(
|
||||||
@@ -767,6 +767,26 @@ class CustomerAPITestCase(TestCase):
|
|||||||
response = self.client.post('/api/backend/customers/', data, format='json')
|
response = self.client.post('/api/backend/customers/', data, format='json')
|
||||||
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
|
||||||
self.assertEqual(response.data['name'], '测试客户')
|
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):
|
def test_filter_customers_by_name(self):
|
||||||
"""测试通过 name 查询参数过滤客户"""
|
"""测试通过 name 查询参数过滤客户"""
|
||||||
@@ -779,6 +799,53 @@ class CustomerAPITestCase(TestCase):
|
|||||||
self.assertEqual(len(results), 1)
|
self.assertEqual(len(results), 1)
|
||||||
self.assertEqual(results[0]['name'], '客户A')
|
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):
|
class VehicleTypeAPITestCase(TestCase):
|
||||||
"""测试车辆类型 API(支持 name 查询参数)"""
|
"""测试车辆类型 API(支持 name 查询参数)"""
|
||||||
|
|||||||
@@ -172,6 +172,15 @@ class CustomerViewSet(BaseViewSet, BasicInfoFilterMixin):
|
|||||||
if self.can_view_all():
|
if self.can_view_all():
|
||||||
return qs.filter(merchant=self.request.user.employee.merchant)
|
return qs.filter(merchant=self.request.user.employee.merchant)
|
||||||
return CustomerVisibilityService.filter_customers_for_employee(qs, self.request.user)
|
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):
|
def has_permission(self, request, view):
|
||||||
if request.user.is_superuser or request.user.has_perm('basic_info.view_all_customer'):
|
if request.user.is_superuser or request.user.has_perm('basic_info.view_all_customer'):
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from rest_framework.fields import empty
|
|||||||
|
|
||||||
from api_v1.models import UploadedFile
|
from api_v1.models import UploadedFile
|
||||||
from api_v1.utils.media import build_public_media_url
|
from api_v1.utils.media import build_public_media_url
|
||||||
|
from api_v1.views.shipment.serializers import SalesItemSerializer
|
||||||
from printing import models
|
from printing import models
|
||||||
from .services import PrintingOrderService, PrintingJobService
|
from .services import PrintingOrderService, PrintingJobService
|
||||||
from basic_info.models import Customer, Employee
|
from basic_info.models import Customer, Employee
|
||||||
@@ -332,6 +333,7 @@ class PrintingJobListSerializer(serializers.ModelSerializer):
|
|||||||
last_completed_state = serializers.CharField(read_only=True)
|
last_completed_state = serializers.CharField(read_only=True)
|
||||||
business_object_id = serializers.SerializerMethodField()
|
business_object_id = serializers.SerializerMethodField()
|
||||||
batch_advance_records = serializers.SerializerMethodField()
|
batch_advance_records = serializers.SerializerMethodField()
|
||||||
|
saleitems = serializers.SerializerMethodField()
|
||||||
merchant_id = serializers.IntegerField(source='merchant.id', read_only=True, allow_null=True)
|
merchant_id = serializers.IntegerField(source='merchant.id', read_only=True, allow_null=True)
|
||||||
|
|
||||||
class Meta:
|
class Meta:
|
||||||
@@ -344,6 +346,7 @@ class PrintingJobListSerializer(serializers.ModelSerializer):
|
|||||||
'status', 'is_completed', 'progress_percentage', 'last_completed_state',
|
'status', 'is_completed', 'progress_percentage', 'last_completed_state',
|
||||||
'business_object_id',
|
'business_object_id',
|
||||||
'batch_advance_records',
|
'batch_advance_records',
|
||||||
|
'saleitems',
|
||||||
'created_at', 'updated_at'
|
'created_at', 'updated_at'
|
||||||
]
|
]
|
||||||
read_only_fields = [
|
read_only_fields = [
|
||||||
@@ -396,6 +399,17 @@ class PrintingJobListSerializer(serializers.ModelSerializer):
|
|||||||
records = list(obj.batch_advance_records.all())
|
records = list(obj.batch_advance_records.all())
|
||||||
return PrintingJobBatchAdvanceRecordSerializer(records, many=True).data
|
return PrintingJobBatchAdvanceRecordSerializer(records, many=True).data
|
||||||
|
|
||||||
|
def get_saleitems(self, obj):
|
||||||
|
items = getattr(obj, '_saleitems_cache', None)
|
||||||
|
if items is None:
|
||||||
|
from shipment.models import SalesItem
|
||||||
|
items = list(
|
||||||
|
SalesItem.objects.select_related('shipment', 'created_by')
|
||||||
|
.filter(printing_job_id=obj.id)
|
||||||
|
.order_by('id')
|
||||||
|
)
|
||||||
|
return SalesItemSerializer(items, many=True).data
|
||||||
|
|
||||||
|
|
||||||
class PrintingJobDetailSerializer(serializers.ModelSerializer):
|
class PrintingJobDetailSerializer(serializers.ModelSerializer):
|
||||||
"""印染款式明细详情序列化器"""
|
"""印染款式明细详情序列化器"""
|
||||||
@@ -411,6 +425,7 @@ class PrintingJobDetailSerializer(serializers.ModelSerializer):
|
|||||||
last_completed_state = serializers.CharField(read_only=True)
|
last_completed_state = serializers.CharField(read_only=True)
|
||||||
business_object_id = serializers.SerializerMethodField()
|
business_object_id = serializers.SerializerMethodField()
|
||||||
batch_advance_records = serializers.SerializerMethodField()
|
batch_advance_records = serializers.SerializerMethodField()
|
||||||
|
saleitems = serializers.SerializerMethodField()
|
||||||
merchant_id = serializers.IntegerField(source='merchant.id', read_only=True, allow_null=True)
|
merchant_id = serializers.IntegerField(source='merchant.id', read_only=True, allow_null=True)
|
||||||
|
|
||||||
class Meta:
|
class Meta:
|
||||||
@@ -423,6 +438,7 @@ class PrintingJobDetailSerializer(serializers.ModelSerializer):
|
|||||||
'progress_percentage', 'last_completed_state',
|
'progress_percentage', 'last_completed_state',
|
||||||
'business_object_id',
|
'business_object_id',
|
||||||
'batch_advance_records',
|
'batch_advance_records',
|
||||||
|
'saleitems',
|
||||||
'created_at', 'updated_at'
|
'created_at', 'updated_at'
|
||||||
]
|
]
|
||||||
read_only_fields = [
|
read_only_fields = [
|
||||||
@@ -444,6 +460,17 @@ class PrintingJobDetailSerializer(serializers.ModelSerializer):
|
|||||||
records = list(obj.batch_advance_records.all())
|
records = list(obj.batch_advance_records.all())
|
||||||
return PrintingJobBatchAdvanceRecordSerializer(records, many=True).data
|
return PrintingJobBatchAdvanceRecordSerializer(records, many=True).data
|
||||||
|
|
||||||
|
def get_saleitems(self, obj):
|
||||||
|
items = getattr(obj, '_saleitems_cache', None)
|
||||||
|
if items is None:
|
||||||
|
from shipment.models import SalesItem
|
||||||
|
items = list(
|
||||||
|
SalesItem.objects.select_related('shipment', 'created_by')
|
||||||
|
.filter(printing_job_id=obj.id)
|
||||||
|
.order_by('id')
|
||||||
|
)
|
||||||
|
return SalesItemSerializer(items, many=True).data
|
||||||
|
|
||||||
|
|
||||||
class PrintingJobCreateUpdateSerializer(serializers.ModelSerializer):
|
class PrintingJobCreateUpdateSerializer(serializers.ModelSerializer):
|
||||||
"""印染款式明细创建/更新序列化器"""
|
"""印染款式明细创建/更新序列化器"""
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from django.contrib.auth.models import Permission
|
|||||||
from django.contrib.contenttypes.models import ContentType
|
from django.contrib.contenttypes.models import ContentType
|
||||||
from basic_info import models as basic_models
|
from basic_info import models as basic_models
|
||||||
from printing import models as printing_models
|
from printing import models as printing_models
|
||||||
|
from shipment import models as shipment_models
|
||||||
from stateflow import models as stateflow_models
|
from stateflow import models as stateflow_models
|
||||||
|
|
||||||
User = get_user_model()
|
User = get_user_model()
|
||||||
@@ -228,6 +229,91 @@ class PrintingJobAPITestCase(TestCase):
|
|||||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
self.assertIn('batch_advance_records', response.data)
|
self.assertIn('batch_advance_records', response.data)
|
||||||
self.assertEqual(len(response.data['batch_advance_records']), 1)
|
self.assertEqual(len(response.data['batch_advance_records']), 1)
|
||||||
|
|
||||||
|
def test_list_printing_jobs_includes_saleitems(self):
|
||||||
|
"""测试 list 返回包含关联的 shipment.SalesItem 数组"""
|
||||||
|
job1 = printing_models.PrintingJob.objects.create(
|
||||||
|
printing_order=self.printing_order,
|
||||||
|
product=self.product,
|
||||||
|
quantity=100,
|
||||||
|
unit='米',
|
||||||
|
size='50*60',
|
||||||
|
pieces=10
|
||||||
|
)
|
||||||
|
job2 = printing_models.PrintingJob.objects.create(
|
||||||
|
printing_order=self.printing_order,
|
||||||
|
product=self.product,
|
||||||
|
quantity=200,
|
||||||
|
unit='米',
|
||||||
|
size='60*70',
|
||||||
|
pieces=20
|
||||||
|
)
|
||||||
|
|
||||||
|
item1 = shipment_models.SalesItem.objects.create(
|
||||||
|
merchant=self.merchant,
|
||||||
|
name='销售品A',
|
||||||
|
quantity='88.50',
|
||||||
|
unit=shipment_models.UnitChoices.METER,
|
||||||
|
created_by=self.user,
|
||||||
|
printing_job_id=job1.id,
|
||||||
|
customer_id=self.customer.id,
|
||||||
|
position='A1-01',
|
||||||
|
remark='备注A',
|
||||||
|
)
|
||||||
|
item2 = shipment_models.SalesItem.objects.create(
|
||||||
|
merchant=self.merchant,
|
||||||
|
name='销售品B',
|
||||||
|
quantity='12.00',
|
||||||
|
unit=shipment_models.UnitChoices.PIECE,
|
||||||
|
created_by=self.user,
|
||||||
|
printing_job_id=job1.id,
|
||||||
|
customer_id=self.customer.id,
|
||||||
|
position='A1-02',
|
||||||
|
remark='备注B',
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.client.get(f'/api/v1/printing-jobs/?printing_order={self.printing_order.id}')
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
data_by_id = {item['id']: item for item in response.data['results']}
|
||||||
|
self.assertIn('saleitems', data_by_id[job1.id])
|
||||||
|
self.assertEqual(len(data_by_id[job1.id]['saleitems']), 2)
|
||||||
|
self.assertEqual([item['id'] for item in data_by_id[job1.id]['saleitems']], [item1.id, item2.id])
|
||||||
|
self.assertEqual(data_by_id[job1.id]['saleitems'][0]['printing_job_id'], job1.id)
|
||||||
|
self.assertEqual(data_by_id[job1.id]['saleitems'][0]['name'], '销售品A')
|
||||||
|
|
||||||
|
self.assertIn('saleitems', data_by_id[job2.id])
|
||||||
|
self.assertEqual(data_by_id[job2.id]['saleitems'], [])
|
||||||
|
|
||||||
|
def test_retrieve_printing_job_includes_saleitems(self):
|
||||||
|
"""测试 detail 返回包含关联的 shipment.SalesItem 数组"""
|
||||||
|
job = printing_models.PrintingJob.objects.create(
|
||||||
|
printing_order=self.printing_order,
|
||||||
|
product=self.product,
|
||||||
|
quantity=100,
|
||||||
|
unit='米',
|
||||||
|
size='50*60',
|
||||||
|
pieces=10
|
||||||
|
)
|
||||||
|
|
||||||
|
shipment_models.SalesItem.objects.create(
|
||||||
|
merchant=self.merchant,
|
||||||
|
name='详情销售品',
|
||||||
|
quantity='66.00',
|
||||||
|
unit=shipment_models.UnitChoices.METER,
|
||||||
|
created_by=self.user,
|
||||||
|
printing_job_id=job.id,
|
||||||
|
customer_id=self.customer.id,
|
||||||
|
position='B2-01',
|
||||||
|
remark='详情备注',
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.client.get(f'/api/v1/printing-jobs/{job.id}/')
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertIn('saleitems', response.data)
|
||||||
|
self.assertEqual(len(response.data['saleitems']), 1)
|
||||||
|
self.assertEqual(response.data['saleitems'][0]['printing_job_id'], job.id)
|
||||||
|
self.assertEqual(response.data['saleitems'][0]['name'], '详情销售品')
|
||||||
|
|
||||||
def test_update_printing_job(self):
|
def test_update_printing_job(self):
|
||||||
"""测试更新款式明细"""
|
"""测试更新款式明细"""
|
||||||
|
|||||||
@@ -426,7 +426,48 @@ class PrintingJobViewSet(CustomerVisibilityFilterMixin, LimitedModelViewSet):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
return queryset
|
return queryset
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _attach_saleitems_cache(jobs):
|
||||||
|
if not jobs:
|
||||||
|
return
|
||||||
|
|
||||||
|
from shipment.models import SalesItem
|
||||||
|
|
||||||
|
job_ids = [job.id for job in jobs]
|
||||||
|
saleitems = list(
|
||||||
|
SalesItem.objects.select_related('shipment', 'created_by')
|
||||||
|
.filter(printing_job_id__in=job_ids)
|
||||||
|
.order_by('id')
|
||||||
|
)
|
||||||
|
|
||||||
|
items_by_job_id = {job_id: [] for job_id in job_ids}
|
||||||
|
for item in saleitems:
|
||||||
|
items_by_job_id.setdefault(item.printing_job_id, []).append(item)
|
||||||
|
|
||||||
|
for job in jobs:
|
||||||
|
job._saleitems_cache = items_by_job_id.get(job.id, [])
|
||||||
|
|
||||||
|
def list(self, request, *args, **kwargs):
|
||||||
|
queryset = self.filter_queryset(self.get_queryset())
|
||||||
|
|
||||||
|
page = self.paginate_queryset(queryset)
|
||||||
|
if page is not None:
|
||||||
|
self._attach_saleitems_cache(page)
|
||||||
|
serializer = self.get_serializer(page, many=True)
|
||||||
|
return self.get_paginated_response(serializer.data)
|
||||||
|
|
||||||
|
jobs = list(queryset)
|
||||||
|
self._attach_saleitems_cache(jobs)
|
||||||
|
serializer = self.get_serializer(jobs, many=True)
|
||||||
|
return Response(serializer.data)
|
||||||
|
|
||||||
|
def retrieve(self, request, *args, **kwargs):
|
||||||
|
instance = self.get_object()
|
||||||
|
self._attach_saleitems_cache([instance])
|
||||||
|
serializer = self.get_serializer(instance)
|
||||||
|
return Response(serializer.data)
|
||||||
|
|
||||||
def destroy(self, request, *args, **kwargs):
|
def destroy(self, request, *args, **kwargs):
|
||||||
"""禁用删除操作"""
|
"""禁用删除操作"""
|
||||||
return Response(
|
return Response(
|
||||||
|
|||||||
Reference in New Issue
Block a user