diff --git a/api_v1/views/stock_change_views/create.py b/api_v1/views/stock_change_views/create.py index e207049..e1a4a63 100644 --- a/api_v1/views/stock_change_views/create.py +++ b/api_v1/views/stock_change_views/create.py @@ -5,11 +5,9 @@ from rest_framework import status, views from rest_framework.response import Response from rest_framework.permissions import IsAuthenticated -from django.db import transaction import logging -from stock import models as stock_models -from basic_info import models as basic_info_models +from stock import services as stock_services from drf_spectacular.utils import extend_schema from api_v1 import serializers @@ -99,68 +97,24 @@ class CreateStockChangeView(StockChangeViewMixin, views.APIView): }, status=status.HTTP_400_BAD_REQUEST) try: - with transaction.atomic(): - # 1. 创建库存变动记录 - try: - warehouse = basic_info_models.WareHouse.objects.get(id=record_data['warehouse']) - except basic_info_models.WareHouse.DoesNotExist: - return Response({ - 'error': f'仓库ID {record_data["warehouse"]} 不存在' - }, status=status.HTTP_400_BAD_REQUEST) - - stock_change_record = stock_models.StockChangeRecord.objects.create( - type=record_data['type'], - warehouse=warehouse, - source_type=record_data['source_type'], - source_id=record_data['source_id'], - merchant=request.user.employee.merchant, - created_by=request.user - ) - - logger.info(f"创建新库存变动记录 ID: {stock_change_record.id}") - - # 2. 创建库存变动明细 - created_details = [] - created_count = 0 - - for product_data in products_data: - product_id = product_data['product'] - quantities = product_data['quantity'] - - # 获取产品信息 - try: - product = basic_info_models.Product.objects.get(id=product_id) - except basic_info_models.Product.DoesNotExist: - raise ValueError(f'产品ID {product_id} 不存在') - - # 为每个数量创建明细记录 - for quantity in quantities: - detail = stock_models.StockChangeDetail.objects.create( - stock_change_record=stock_change_record, - product=product, - quantity=quantity, - merchant=request.user.employee.merchant, - unit=product.unit - ) - created_details.append(detail) - created_count += 1 + stock_change_record, created_details, created_count = stock_services.create_stock_change_record_with_details( + merchant=request.user.employee.merchant, + created_by=request.user, + type=record_data['type'], + warehouse_id=record_data['warehouse'], + source_type=record_data['source_type'], + source_id=record_data['source_id'], + products=products_data, + ) + response_serializer = serializers.CreateStockChangeResponseSerializer({ + 'stock_change_record': stock_change_record, + 'details': created_details, + 'message': f'成功创建库存变动记录及 {created_count} 条明细', + 'created_details_count': created_count + }) + + return Response(response_serializer.data, status=status.HTTP_201_CREATED) - if warehouse.merchant.auto_complete_stock_change: - # 自动确认库存变动 - from stock import services as stock_services - stock_services.make_stock_change_completed(stock_change_record) - logger.info(f"自动确认库存变动记录 ID: {stock_change_record.id}") - - # 3. 构建响应 - response_serializer = serializers.CreateStockChangeResponseSerializer({ - 'stock_change_record': stock_change_record, - 'details': created_details, - 'message': f'成功创建库存变动记录及 {created_count} 条明细', - 'created_details_count': created_count - }) - - return Response(response_serializer.data, status=status.HTTP_201_CREATED) - except ValueError as e: return Response({ 'error': str(e) diff --git a/stock/new_version.md b/stock/new_version.md new file mode 100644 index 0000000..fa497b7 --- /dev/null +++ b/stock/new_version.md @@ -0,0 +1,5 @@ +1. 严进严出逻辑实现 +2. 宽进宽出逻辑实现 +3. 调拨单的实现 +4. 采购单的实现 +5. 盘点单的实现 \ No newline at end of file diff --git a/stock/services.py b/stock/services.py index ae8d1dc..535233b 100644 --- a/stock/services.py +++ b/stock/services.py @@ -1,13 +1,101 @@ from . import models +from django.db import transaction from django.db.models import Sum from django.utils import timezone from sse.services import push_simple_message_with_object_id +from basic_info import models as basic_models import logging +from typing import List, Dict, Any, Tuple logger = logging.getLogger(__name__) +def create_stock_change_record_with_details( + *, + merchant: basic_models.Merchant, + created_by, + type: int, + warehouse_id: int, + source_type: int, + source_id: int | None = None, + products: List[Dict[str, Any]] | None = None, +) -> Tuple[models.StockChangeRecord, List[models.StockChangeDetail], int]: + """ + 创建库存变动记录及其明细 + + 参数: + merchant: 当前商户 + created_by: 操作人(可为空) + type: 出入库类型 + warehouse_id: 仓库ID + source_type: 来源类型 + source_id: 来源单据ID,可选 + products: 产品及数量列表 + """ + + if not merchant: + raise ValueError('必须提供商户信息') + if not warehouse_id: + raise ValueError('必须提供仓库ID') + if not products: + raise ValueError('产品列表不能为空') + + try: + warehouse = basic_models.WareHouse.objects.get(id=warehouse_id) + except basic_models.WareHouse.DoesNotExist: + raise ValueError(f'仓库ID {warehouse_id} 不存在') + + if warehouse.merchant_id != merchant.id: + raise ValueError('仓库不属于当前商户') + + created_details: List[models.StockChangeDetail] = [] + created_count = 0 + + with transaction.atomic(): + stock_change_record = models.StockChangeRecord.objects.create( + type=type, + warehouse=warehouse, + source_type=source_type, + source_id=source_id, + merchant=merchant, + created_by=created_by, + ) + logger.info('创建新库存变动记录 ID: %s', stock_change_record.id) + + for product_data in products: + product_id = product_data.get('product') + quantities = product_data.get('quantity', []) + + if not product_id or not quantities: + raise ValueError('产品数据不完整,缺少 product 或 quantity') + + try: + product = basic_models.Product.objects.get(id=product_id) + except basic_models.Product.DoesNotExist: + raise ValueError(f'产品ID {product_id} 不存在') + + if product.merchant_id != merchant.id: + raise ValueError(f'产品ID {product_id} 不属于当前商户') + + for quantity in quantities: + detail = models.StockChangeDetail.objects.create( + stock_change_record=stock_change_record, + product=product, + quantity=quantity, + merchant=merchant, + unit=product.unit, + ) + created_details.append(detail) + created_count += 1 + + if warehouse.merchant.auto_complete_stock_change: + make_stock_change_completed(stock_change_record) + logger.info('自动确认库存变动记录 ID: %s', stock_change_record.id) + + return stock_change_record, created_details, created_count + + def find_inventory(product_id: int, warehouse_id: int) -> models.Inventory | None: """根据产品ID和仓库ID查找库存记录""" diff --git a/stock/tests.py b/stock/tests.py index d540f3c..4c6e88b 100644 --- a/stock/tests.py +++ b/stock/tests.py @@ -371,6 +371,72 @@ class MakeStockChangeCompletedTestCase(StockServicesTestCase): # 验证记录状态 self.stock_change_in.refresh_from_db() self.assertTrue(self.stock_change_in.is_finished) + + +class CreateStockChangeRecordWithDetailsTestCase(StockServicesTestCase): + """测试 create_stock_change_record_with_details 服务""" + + def setUp(self): + super().setUp() + self.products_payload = [ + {'product': self.product_fabric_a.id, 'quantity': [Decimal('10.00'), Decimal('20.50')]}, + {'product': self.product_fabric_b.id, 'quantity': [Decimal('5.00')]} + ] + + def test_create_record_success(self): + """创建库存变动记录并返回所有明细""" + record, details, created_count = services.create_stock_change_record_with_details( + merchant=self.merchant, + created_by=None, + type=models.StockChangeTypeEnum.ADD, + warehouse_id=self.warehouse_main.id, + source_type=models.StockChangeSourceEnum.PURCHASE, + source_id=self.purchase_order.id, + products=self.products_payload, + ) + + self.assertEqual(record.details.count(), 3) + self.assertEqual(len(details), 3) + self.assertEqual(created_count, 3) + self.assertEqual(record.merchant, self.merchant) + self.assertEqual(record.warehouse, self.warehouse_main) + self.assertCountEqual( + [Decimal('10.00'), Decimal('20.50'), Decimal('5.00')], + [detail.quantity for detail in details] + ) + + def test_auto_complete_triggers_completion(self): + """商户开启自动确认时会调用完成逻辑""" + self.merchant.auto_complete_stock_change = True + self.merchant.save() + + with patch('stock.services.make_stock_change_completed') as mock_complete: + services.create_stock_change_record_with_details( + merchant=self.merchant, + created_by=None, + type=models.StockChangeTypeEnum.ADD, + warehouse_id=self.warehouse_main.id, + source_type=models.StockChangeSourceEnum.PURCHASE, + source_id=self.purchase_order.id, + products=self.products_payload, + ) + + mock_complete.assert_called_once() + + def test_invalid_product_raises_error(self): + """无效的产品ID会触发 ValueError""" + products = [{'product': 99999, 'quantity': [Decimal('1.00')]}] + + with self.assertRaises(ValueError): + services.create_stock_change_record_with_details( + merchant=self.merchant, + created_by=None, + type=models.StockChangeTypeEnum.ADD, + warehouse_id=self.warehouse_main.id, + source_type=models.StockChangeSourceEnum.PURCHASE, + source_id=self.purchase_order.id, + products=products, + ) def test_complete_outbound_record(self): """测试完成出库记录"""