From 84dbb2ee100a338ab7e782075277532d5915aad9 Mon Sep 17 00:00:00 2001 From: colaftc Date: Fri, 13 Mar 2026 15:57:40 +0800 Subject: [PATCH] fix: merchant not required at customer create api --- api_man/serializers.py | 35 ++++++++ api_man/tests.py | 69 ++++++++++++++- api_man/views.py | 9 ++ api_v1/views/printing/serializers.py | 27 ++++++ .../views/printing/test_printing_job_api.py | 86 +++++++++++++++++++ api_v1/views/printing/views.py | 43 +++++++++- 6 files changed, 267 insertions(+), 2 deletions(-) diff --git a/api_man/serializers.py b/api_man/serializers.py index c70abfe..d84aeaa 100644 --- a/api_man/serializers.py +++ b/api_man/serializers.py @@ -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__' diff --git a/api_man/tests.py b/api_man/tests.py index 3651750..58e4d14 100644 --- a/api_man/tests.py +++ b/api_man/tests.py @@ -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 查询参数)""" diff --git a/api_man/views.py b/api_man/views.py index 124fe64..07b4f5c 100644 --- a/api_man/views.py +++ b/api_man/views.py @@ -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'): diff --git a/api_v1/views/printing/serializers.py b/api_v1/views/printing/serializers.py index e2091de..4e59090 100644 --- a/api_v1/views/printing/serializers.py +++ b/api_v1/views/printing/serializers.py @@ -7,6 +7,7 @@ from rest_framework.fields import empty from api_v1.models import UploadedFile from api_v1.utils.media import build_public_media_url +from api_v1.views.shipment.serializers import SalesItemSerializer from printing import models from .services import PrintingOrderService, PrintingJobService from basic_info.models import Customer, Employee @@ -332,6 +333,7 @@ class PrintingJobListSerializer(serializers.ModelSerializer): last_completed_state = serializers.CharField(read_only=True) business_object_id = serializers.SerializerMethodField() batch_advance_records = serializers.SerializerMethodField() + saleitems = serializers.SerializerMethodField() merchant_id = serializers.IntegerField(source='merchant.id', read_only=True, allow_null=True) class Meta: @@ -344,6 +346,7 @@ class PrintingJobListSerializer(serializers.ModelSerializer): 'status', 'is_completed', 'progress_percentage', 'last_completed_state', 'business_object_id', 'batch_advance_records', + 'saleitems', 'created_at', 'updated_at' ] read_only_fields = [ @@ -396,6 +399,17 @@ class PrintingJobListSerializer(serializers.ModelSerializer): records = list(obj.batch_advance_records.all()) 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): """印染款式明细详情序列化器""" @@ -411,6 +425,7 @@ class PrintingJobDetailSerializer(serializers.ModelSerializer): last_completed_state = serializers.CharField(read_only=True) business_object_id = serializers.SerializerMethodField() batch_advance_records = serializers.SerializerMethodField() + saleitems = serializers.SerializerMethodField() merchant_id = serializers.IntegerField(source='merchant.id', read_only=True, allow_null=True) class Meta: @@ -423,6 +438,7 @@ class PrintingJobDetailSerializer(serializers.ModelSerializer): 'progress_percentage', 'last_completed_state', 'business_object_id', 'batch_advance_records', + 'saleitems', 'created_at', 'updated_at' ] read_only_fields = [ @@ -444,6 +460,17 @@ class PrintingJobDetailSerializer(serializers.ModelSerializer): records = list(obj.batch_advance_records.all()) 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): """印染款式明细创建/更新序列化器""" diff --git a/api_v1/views/printing/test_printing_job_api.py b/api_v1/views/printing/test_printing_job_api.py index 42fc752..76d3a52 100644 --- a/api_v1/views/printing/test_printing_job_api.py +++ b/api_v1/views/printing/test_printing_job_api.py @@ -10,6 +10,7 @@ from django.contrib.auth.models import Permission from django.contrib.contenttypes.models import ContentType from basic_info import models as basic_models from printing import models as printing_models +from shipment import models as shipment_models from stateflow import models as stateflow_models User = get_user_model() @@ -228,6 +229,91 @@ class PrintingJobAPITestCase(TestCase): self.assertEqual(response.status_code, status.HTTP_200_OK) self.assertIn('batch_advance_records', response.data) 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): """测试更新款式明细""" diff --git a/api_v1/views/printing/views.py b/api_v1/views/printing/views.py index 33eeb05..8cfd82e 100644 --- a/api_v1/views/printing/views.py +++ b/api_v1/views/printing/views.py @@ -426,7 +426,48 @@ class PrintingJobViewSet(CustomerVisibilityFilterMixin, LimitedModelViewSet): ) ) 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): """禁用删除操作""" return Response(