diff --git a/api_man/employee_user_binding.md b/api_man/employee_user_binding.md index 3dd49ad..a23f60c 100644 --- a/api_man/employee_user_binding.md +++ b/api_man/employee_user_binding.md @@ -68,6 +68,13 @@ Base URL: `/api_man/user-profiles/` | PATCH | `/api_man/user-profiles/{id}/` | Partially update an existing user profile | | DELETE | `/api_man/user-profiles/{id}/` | Delete a user profile | +### 2.3 User Creation API Endpoint +Base URL: `/api_v1/users/` + +| Method | URL Pattern | Action | +|--------|-------------|--------| +| POST | `/api_v1/users/create/` | Create a new User and associated UserProfile | + ## 3. Binding User to Employee ### 3.1 During Employee Creation @@ -216,7 +223,70 @@ To update an existing UserProfile, use the PUT or PATCH method. } ``` -## 5. Validation Rules +## 5. Creating User with Profile + +### 5.1 User and Profile Creation +This API creates a new Django User and associated UserProfile in a single request. The UserProfile links the User to a specific Merchant. + +#### Request Body Example: +```json +{ + "username": "newuser", + "email": "newuser@example.com", + "password": "securepass123", + "is_staff": false, + "description": "New user account", + "merchant_id": 1 +} +``` + +#### Request Parameters: +- `username` (required): The username for the new user account +- `email` (optional): The email address for the user +- `password` (required): The password for the user account (minimum 6 characters) +- `is_staff` (optional, default: false): Whether the user should have staff privileges +- `description` (optional): Description for the UserProfile +- `merchant_id` (required): ID of the merchant to associate with the user + +#### Response Example: +```json +{ + "user": { + "id": 124, + "username": "newuser", + "email": "newuser@example.com", + "is_staff": false, + "is_active": true, + "date_joined": "2025-11-24T14:00:00Z" + }, + "profile": { + "id": 790, + "user": { + "id": 124, + "username": "newuser", + "email": "newuser@example.com", + "is_staff": false, + "is_active": true, + "date_joined": "2025-11-24T14:00:00Z" + }, + "merchant": 1, + "description": "New user account", + "created_at": "2025-11-24T14:00:00Z", + "updated_at": "2025-11-24T14:00:00Z" + } +} +``` + +### 5.2 Authentication Requirements +- The request must be authenticated with a valid user session +- The authenticated user must have appropriate permissions to create new users + +### 5.3 Validation Rules +- Username must be unique across the system +- Password must be at least 6 characters long +- merchant_id must correspond to an existing merchant in the system + +## 6. Validation Rules ### 5.1 Employee-User Binding Rules 1. **One-to-one constraint**: A User can only be bound to one Employee at a time. @@ -250,6 +320,78 @@ To update an existing UserProfile, use the PUT or PATCH method. } ``` +### 6.1 User Creation API Errors + +#### Username already exists +```json +{ + "username": [ + "用户名已存在" + ] +} +``` + +#### Merchant does not exist +```json +{ + "merchant_id": [ + "商户不存在" + ] +} +``` + +#### Password too short +```json +{ + "password": [ + "Ensure this field has at least 6 characters." + ] +} +``` + +#### Unauthenticated request +```json +{ + "error": "未授权" +} +``` + +#### Server error during creation +```json +{ + "error": "创建用户失败: [detailed error message]" +} +``` + +## 7. Usage Examples + +### 7.1 Creating a Complete User Flow +```bash +# 1. Create a new user with profile +curl -X POST http://localhost/api/v1/users/create/ \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer [token]" \ + -d '{ + "username": "john_doe", + "email": "john@example.com", + "password": "securepass123", + "description": "Store employee", + "merchant_id": 1 + }' + +# Response will contain both user and profile IDs +# Use the user ID to create an Employee record and bind the user +curl -X POST http://localhost/api_man/employees/ \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer [token]" \ + -d '{ + "name": "John Doe", + "position": 1, + "mobile": "1234567890", + "sys_user": 124 # User ID from the previous response + }' +``` + ### Invalid User ID ```json { diff --git a/api_v1/serializers.py b/api_v1/serializers.py index d9013b3..dbe8bae 100644 --- a/api_v1/serializers.py +++ b/api_v1/serializers.py @@ -55,6 +55,51 @@ class CreateStockChangeSerializer(serializers.Serializer): return value +class RelaxedQuantitySerializer(serializers.Serializer): + """宽松模式数量结构""" + value = serializers.DecimalField( + max_digits=10, + decimal_places=2, + min_value=Decimal('0.01'), + help_text="产品总数量" + ) + unit_count = serializers.DecimalField( + max_digits=10, + decimal_places=2, + min_value=Decimal('0.01'), + default=Decimal('1.00'), + help_text="单条(匹)数量,默认 1" + ) + + +class RelaxedProductStockChangeSerializer(serializers.Serializer): + """宽松模式产品序列化器""" + product = serializers.IntegerField(help_text="产品ID") + quantity = RelaxedQuantitySerializer(help_text="数量定义(总数与单条数)") + + def validate_product(self, value: int) -> int: + try: + basic_info_models.Product.objects.get(id=value) + except basic_info_models.Product.DoesNotExist: + raise serializers.ValidationError(f"产品ID {value} 不存在") + return value + + +class CreateStockChangeRelaxedSerializer(serializers.Serializer): + """宽松模式创建库存变动记录""" + products = RelaxedProductStockChangeSerializer(many=True, help_text="产品列表") + + def validate_products(self, value: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + if not value: + raise serializers.ValidationError("产品列表不能为空") + + product_ids = [item['product'] for item in value] + if len(product_ids) != len(set(product_ids)): + raise serializers.ValidationError("产品列表中存在重复的产品ID") + + return value + + class StockChangeRecordResponseSerializer(serializers.ModelSerializer): """库存变动记录响应序列化器""" diff --git a/api_v1/tests.py b/api_v1/tests.py index 7ce503c..76b4010 100644 --- a/api_v1/tests.py +++ b/api_v1/tests.py @@ -1,3 +1,112 @@ from django.test import TestCase +from django.contrib.auth.models import User, Permission +from rest_framework.test import APIClient +from rest_framework import status +from basic_info.models import Merchant, UserProfile -# Create your tests here. + +class UserCreationAPITestCase(TestCase): + """测试用户创建 API""" + + def setUp(self): + """设置测试数据""" + # 创建商户 + self.merchant = Merchant.objects.create( + name='测试商户', + type=1 # 假设1是有效的MerchantType + ) + + # 创建管理员用户 + self.admin_user = User.objects.create_user( + username='admin', + password='adminpass123', + is_staff=True + ) + # 授予创建用户所需的权限 + add_user_perm = Permission.objects.get(codename='add_user') + self.admin_user.user_permissions.add(add_user_perm) + self.admin_user.save() + + # 设置 API 客户端 + self.client = APIClient() + self.client.force_authenticate(user=self.admin_user) + + def test_create_user_with_profile(self): + """测试创建用户和用户资料""" + data = { + 'username': 'testuser', + 'email': 'test@example.com', + 'password': 'testpass123', + 'is_staff': False, + 'description': '测试用户资料', + 'merchant_id': self.merchant.id + } + response = self.client.post('/api/v1/users/create/', data, format='json') + self.assertEqual(response.status_code, status.HTTP_201_CREATED) + + # 验证返回的数据 + self.assertIn('user', response.data) + self.assertIn('profile', response.data) + self.assertEqual(response.data['user']['username'], 'testuser') + self.assertEqual(response.data['user']['email'], 'test@example.com') + self.assertEqual(response.data['user']['is_staff'], False) + self.assertEqual(response.data['profile']['description'], '测试用户资料') + self.assertEqual(response.data['profile']['merchant'], self.merchant.id) + + # 验证数据库中的用户 + user = User.objects.get(username='testuser') + self.assertEqual(user.email, 'test@example.com') + self.assertFalse(user.is_staff) + + # 验证数据库中的用户资料 + profile = UserProfile.objects.get(user=user) + self.assertEqual(profile.merchant, self.merchant) + self.assertEqual(profile.description, '测试用户资料') + + def test_create_user_with_duplicate_username(self): + """测试创建用户时使用重复的用户名""" + # 先创建一个用户 + User.objects.create_user(username='existinguser', password='pass123') + + data = { + 'username': 'existinguser', # 重复的用户名 + 'password': 'testpass123', + 'merchant_id': self.merchant.id + } + response = self.client.post('/api/v1/users/create/', data, format='json') + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn('username', response.data) + + def test_create_user_with_invalid_merchant(self): + """测试创建用户时使用无效的merchant_id""" + data = { + 'username': 'testuser', + 'password': 'testpass123', + 'merchant_id': 999 # 不存在的merchant_id + } + response = self.client.post('/api/v1/users/create/', data, format='json') + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn('merchant_id', response.data) + + def test_create_user_unauthenticated(self): + """测试未认证用户创建用户""" + self.client.force_authenticate(user=None) + + data = { + 'username': 'testuser', + 'password': 'testpass123', + 'merchant_id': self.merchant.id + } + response = self.client.post('/api/v1/users/create/', data, format='json') + self.assertEqual(response.status_code, status.HTTP_401_UNAUTHORIZED) + + def test_create_user_with_short_password(self): + """测试创建用户时密码过短""" + data = { + 'username': 'testuser', + 'password': '123', # 密码过短 + 'merchant_id': self.merchant.id + } + response = self.client.post('/api/v1/users/create/', data, format='json') + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn('password', response.data) diff --git a/api_v1/urls.py b/api_v1/urls.py index bcbd0c4..9ef87c4 100644 --- a/api_v1/urls.py +++ b/api_v1/urls.py @@ -1,11 +1,12 @@ from django.urls import path, include from rest_framework.routers import DefaultRouter -from .views import stock_change_views, user_info, inventory, product_image, stateflow +from .views import stock_change_views, user_info, inventory, product_image, stateflow, users from .views.stock_change_views.snapshot import StockSnapshotListView from .views.printing.views import PrintingOrderViewSet, PrintingJobViewSet, PlateOrderViewSet from .views.upload import UploadFileViewSet from .views.products import ProductQuickViewSet from .views.parameters import StateParameterViewSet +from .views.users import CreateUserWithProfileView # 创建 DRF Router for Stateflow stateflow_router = DefaultRouter() @@ -27,6 +28,7 @@ urlpatterns = [ path('stock-snapshots/', StockSnapshotListView.as_view(), name='list_stock_snapshots'), path('stock-changes/', stock_change_views.list_stock_changes, name='list_stock_changes'), path('stock-change/', stock_change_views.create_full_stock_change, name='create_full_stock_change'), + path('stock-change/relaxed/', stock_change_views.create_relaxed_stock_change, name='create_relaxed_stock_change'), path('stock-change//', stock_change_views.get_stock_change, name='get_stock_change'), path( 'set-merchant-auto-complete-stock-change/', @@ -37,6 +39,9 @@ urlpatterns = [ # 用户信息 API path('user-info/', user_info.user_info, name='user_info'), + # 用户创建 API + path('users/create/', CreateUserWithProfileView.as_view(), name='create_user_with_profile'), + # 库存查询 API path('inventory/', inventory.InventoryAPIView.as_view(), name='inventory'), diff --git a/api_v1/views/__init__.py b/api_v1/views/__init__.py index 84e8878..050a70f 100644 --- a/api_v1/views/__init__.py +++ b/api_v1/views/__init__.py @@ -20,4 +20,5 @@ __all__ = [ 'get_stock_change', 'set_merchant_auto_complete_stock_change', 'stateflow', + 'create_user_with_profile', ] diff --git a/api_v1/views/stock_change_views/README.md b/api_v1/views/stock_change_views/README.md index 36a04d0..ac179af 100644 --- a/api_v1/views/stock_change_views/README.md +++ b/api_v1/views/stock_change_views/README.md @@ -63,6 +63,7 @@ stock_change_views/ ```python # 这些接口保持不变,URL 配置无需修改 create_full_stock_change = CreateStockChangeView.as_view() +create_relaxed_stock_change = CreateStockChangeRelaxedView.as_view() list_stock_changes = ListStockChangesView.as_view() get_stock_change = GetStockChangeView.as_view() set_merchant_auto_complete_stock_change = SetMerchantAutoCompleteView.as_view() @@ -87,6 +88,7 @@ from api_v1.views import stock_change_views urlpatterns = [ path('stock-change/', stock_change_views.create_full_stock_change), + path('stock-change/relaxed/', stock_change_views.create_relaxed_stock_change), path('stock-changes/', stock_change_views.list_stock_changes), path('stock-change//', stock_change_views.get_stock_change), path('set-merchant-auto-complete-stock-change/', stock_change_views.set_merchant_auto_complete_stock_change), @@ -98,6 +100,7 @@ urlpatterns = [ ```python from api_v1.views.stock_change_views import ( CreateStockChangeView, + CreateStockChangeRelaxedView, ListStockChangesView, GetStockChangeView, SetMerchantAutoCompleteView, @@ -105,6 +108,7 @@ from api_v1.views.stock_change_views import ( urlpatterns = [ path('stock-change/', CreateStockChangeView.as_view()), + path('stock-change/relaxed/', CreateStockChangeRelaxedView.as_view()), path('stock-changes/', ListStockChangesView.as_view()), path('stock-change//', GetStockChangeView.as_view()), path('set-merchant-auto-complete-stock-change/', SetMerchantAutoCompleteView.as_view()), diff --git a/api_v1/views/stock_change_views/__init__.py b/api_v1/views/stock_change_views/__init__.py index 3b7658f..f8515e1 100644 --- a/api_v1/views/stock_change_views/__init__.py +++ b/api_v1/views/stock_change_views/__init__.py @@ -5,13 +5,14 @@ """ from .mixins import StockChangeViewMixin -from .create import CreateStockChangeView +from .create import CreateStockChangeView, CreateStockChangeRelaxedView from .list import ListStockChangesView from .detail import GetStockChangeView from .settings import SetMerchantAutoCompleteView # 向后兼容:保持原有的函数式接口 create_full_stock_change = CreateStockChangeView.as_view() +create_relaxed_stock_change = CreateStockChangeRelaxedView.as_view() list_stock_changes = ListStockChangesView.as_view() get_stock_change = GetStockChangeView.as_view() set_merchant_auto_complete_stock_change = SetMerchantAutoCompleteView.as_view() @@ -19,10 +20,12 @@ set_merchant_auto_complete_stock_change = SetMerchantAutoCompleteView.as_view() __all__ = [ 'StockChangeViewMixin', 'CreateStockChangeView', + 'CreateStockChangeRelaxedView', 'ListStockChangesView', 'GetStockChangeView', 'SetMerchantAutoCompleteView', 'create_full_stock_change', + 'create_relaxed_stock_change', 'list_stock_changes', 'get_stock_change', 'set_merchant_auto_complete_stock_change', diff --git a/api_v1/views/stock_change_views/create.py b/api_v1/views/stock_change_views/create.py index e1a4a63..4591692 100644 --- a/api_v1/views/stock_change_views/create.py +++ b/api_v1/views/stock_change_views/create.py @@ -126,3 +126,89 @@ class CreateStockChangeView(StockChangeViewMixin, views.APIView): 'error': '创建库存变动记录失败', 'message': str(e) }, status=status.HTTP_500_INTERNAL_SERVER_ERROR) + + +class CreateStockChangeRelaxedView(StockChangeViewMixin, views.APIView): + """宽松模式:根据总量与单条数量拆分的出入库记录""" + + permission_classes = [IsAuthenticated] + + @extend_schema( + tags=['创建出入库'], + request=serializers.CreateStockChangeRelaxedSerializer, + responses={201: serializers.CreateStockChangeResponseSerializer}, + summary="创建库存变动记录(宽松模式)", + description="根据总数量与单条数量自动拆分明细" + ) + def post(self, request): + if not self.check_employee_permission(request): + return self.permission_error_response('无权限访问') + + record_data = { + 'type': request.data.get('type'), + 'warehouse': request.data.get('warehouse'), + 'source_type': request.data.get('source_type'), + 'source_id': request.data.get('source_id'), + } + + if not all([record_data['type'], record_data['warehouse'], record_data['source_type']]): + return Response({ + 'error': '缺少必要参数', + 'message': '请提供 type, warehouse, source_type' + }, status=status.HTTP_400_BAD_REQUEST) + + products_data = request.data.get('products', []) + if not products_data: + return Response({ + 'error': '产品列表不能为空' + }, status=status.HTTP_400_BAD_REQUEST) + + if not self.validate_warehouse_visibility(record_data['warehouse'], request): + return Response({ + 'error': f'仓库ID {record_data["warehouse"]} 对当前用户不可见' + }, status=status.HTTP_403_FORBIDDEN) + + for p in products_data: + product_id = p.get('product') + if not self.validate_product_visibility(product_id, request): + return Response({ + 'error': f'产品ID {product_id} 对当前用户不可见' + }, status=status.HTTP_403_FORBIDDEN) + + serializer = serializers.CreateStockChangeRelaxedSerializer(data={'products': products_data}) + if not serializer.is_valid(): + return Response({ + 'error': '产品数据验证失败', + 'details': serializer.errors + }, status=status.HTTP_400_BAD_REQUEST) + + try: + stock_change_record, created_details, created_count = stock_services.create_stock_change_record_relaxed( + 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) + + except ValueError as e: + return Response({ + 'error': str(e) + }, status=status.HTTP_400_BAD_REQUEST) + + except Exception as e: + logger.error(f"创建宽松模式库存变动记录失败: {str(e)}", exc_info=True) + return Response({ + 'error': '创建库存变动记录失败', + 'message': str(e) + }, status=status.HTTP_500_INTERNAL_SERVER_ERROR) diff --git a/api_v1/views/users.py b/api_v1/views/users.py new file mode 100644 index 0000000..87de043 --- /dev/null +++ b/api_v1/views/users.py @@ -0,0 +1,123 @@ +from django.contrib.auth.models import User +from rest_framework.decorators import permission_classes +from rest_framework.permissions import IsAuthenticated, DjangoModelPermissions +from rest_framework.response import Response +from rest_framework.pagination import LimitOffsetPagination +from rest_framework import status, serializers +from rest_framework.generics import GenericAPIView +from basic_info.models import UserProfile + + +class UserCreationSerializer(serializers.Serializer): + """用户创建序列化器""" + username = serializers.CharField(max_length=150) + email = serializers.EmailField(required=False) + password = serializers.CharField(min_length=6, write_only=True) + is_staff = serializers.BooleanField(default=False) + description = serializers.CharField(required=False, allow_blank=True) + merchant_id = serializers.IntegerField() + + def validate_username(self, value): + """验证用户名是否已存在""" + if User.objects.filter(username=value).exists(): + raise serializers.ValidationError("用户名已存在") + return value + + def validate_merchant_id(self, value): + """验证merchant_id是否存在""" + from basic_info.models import Merchant + if not Merchant.objects.filter(id=value).exists(): + raise serializers.ValidationError("商户不存在") + return value + + +class UserSerializer(serializers.ModelSerializer): + """用户信息序列化器""" + class Meta: + model = User + fields = ['id', 'username', 'email', 'is_staff', 'is_active', 'date_joined'] + + +class UserProfileDetailSerializer(serializers.ModelSerializer): + """用户资料详细信息序列化器""" + user = UserSerializer(read_only=True) + + class Meta: + model = UserProfile + fields = ['id', 'user', 'merchant', 'description', 'created_at', 'updated_at'] + + +class CreateUserWithProfileView(GenericAPIView): + """创建用户和关联的用户资料信息视图""" + queryset = User.objects.all() + permission_classes = [IsAuthenticated, DjangoModelPermissions] + pagination_class = LimitOffsetPagination + + def post(self, request): + """ + 创建用户和用户资料 + + POST /api/v1/users/create/ + + 请求参数: + - username: 用户名 (必需) + - email: 邮箱 (可选) + - password: 密码 (必需, 最少6位) + - is_staff: 是否为员工 (可选, 默认为False) + - description: 用户资料描述 (可选) + - merchant_id: 商户ID (必需) + + 返回: + - user: 创建的用户信息 + - profile: 创建的用户资料信息 + """ + # 检查当前用户是否有权限(超级用户或有员工身份) + if not request.user.is_authenticated: + return Response({'error': '未授权'}, status=status.HTTP_401_UNAUTHORIZED) + + # 序列化和验证请求数据 + serializer = UserCreationSerializer(data=request.data) + if not serializer.is_valid(): + return Response(serializer.errors, status=status.HTTP_400_BAD_REQUEST) + + username = serializer.validated_data['username'] + email = serializer.validated_data.get('email', '') + password = serializer.validated_data['password'] + is_staff = serializer.validated_data.get('is_staff', False) + description = serializer.validated_data.get('description', '') + merchant_id = serializer.validated_data['merchant_id'] + + try: + # 创建用户 + user = User.objects.create_user( + username=username, + email=email, + password=password, + is_staff=is_staff + ) + + # 创建用户资料 + from basic_info.models import Merchant + merchant = Merchant.objects.get(id=merchant_id) + profile = UserProfile.objects.create( + user=user, + merchant=merchant, + description=description + ) + + # 返回创建的用户和用户资料信息 + user_data = UserSerializer(user).data + profile_data = UserProfileDetailSerializer(profile).data + + return Response({ + 'user': user_data, + 'profile': profile_data + }, status=status.HTTP_201_CREATED) + + except Exception as e: + # 如果创建过程中出现错误,删除已创建的用户 + if 'user' in locals(): + user.delete() + return Response({ + 'error': f'创建用户失败: {str(e)}' + }, status=status.HTTP_500_INTERNAL_SERVER_ERROR) diff --git a/stock/services.py b/stock/services.py index 535233b..837194c 100644 --- a/stock/services.py +++ b/stock/services.py @@ -4,6 +4,7 @@ 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 +from decimal import Decimal, InvalidOperation import logging from typing import List, Dict, Any, Tuple @@ -11,6 +12,26 @@ from typing import List, Dict, Any, Tuple logger = logging.getLogger(__name__) +def _raise_unimplemented_mode(): + raise ValueError('仓库出入库模式为【严进严出】,该模式暂未支持创建出入库记录') + + +def _ensure_strict_mode(warehouse: basic_models.WareHouse): + if warehouse.mode == basic_models.WareHouseModeEnum.RESTRICT_IN: + return + if warehouse.mode == basic_models.WareHouseModeEnum.RESTRICT_IN_OUT: + _raise_unimplemented_mode() + raise ValueError('仓库出入库模式为【宽进宽出】,请使用宽松模式接口创建出入库记录') + + +def _ensure_relaxed_mode(warehouse: basic_models.WareHouse): + if warehouse.mode == basic_models.WareHouseModeEnum.UNRESTRICTED: + return + if warehouse.mode == basic_models.WareHouseModeEnum.RESTRICT_IN_OUT: + _raise_unimplemented_mode() + raise ValueError('仓库出入库模式为【严进宽出】,请使用严谨模式接口创建出入库记录') + + def create_stock_change_record_with_details( *, merchant: basic_models.Merchant, @@ -48,6 +69,7 @@ def create_stock_change_record_with_details( if warehouse.merchant_id != merchant.id: raise ValueError('仓库不属于当前商户') + _ensure_strict_mode(warehouse) created_details: List[models.StockChangeDetail] = [] created_count = 0 @@ -96,6 +118,116 @@ def create_stock_change_record_with_details( return stock_change_record, created_details, created_count +def _to_decimal(value, field_name: str) -> Decimal: + try: + return Decimal(str(value)) + except (InvalidOperation, TypeError): + raise ValueError(f'{field_name} 必须是合法的数值') + + +def _split_quantities(total: Decimal, unit_size: Decimal) -> List[Decimal]: + if total <= 0: + raise ValueError('quantity.value 必须大于 0') + if unit_size <= 0: + raise ValueError('quantity.unit_count 必须大于 0') + + quantities: List[Decimal] = [] + num_full = int(total // unit_size) + remainder = total % unit_size + + if num_full == 0: + quantities.append(total) + else: + quantities.extend([unit_size] * num_full) + if remainder > 0: + quantities.append(remainder) + + return quantities + + +def create_stock_change_record_relaxed( + *, + 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]: + """ + 宽松模式:根据总数量和单条数量自动拆分明细 + """ + + 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('仓库不属于当前商户') + _ensure_relaxed_mode(warehouse) + + created_details: List[models.StockChangeDetail] = [] + + 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') + quantity_data = product_data.get('quantity') or {} + + if not product_id: + raise ValueError('产品ID不能为空') + + total_value = quantity_data.get('value') + unit_size = quantity_data.get('unit_count', 1) + + total_value = _to_decimal(total_value, 'quantity.value') + unit_size = _to_decimal(unit_size, 'quantity.unit_count') + + 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} 不属于当前商户') + + quantities = _split_quantities(total_value, unit_size) + + 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) + + 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, len(created_details) + + 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 4c6e88b..ed9b824 100644 --- a/stock/tests.py +++ b/stock/tests.py @@ -5,7 +5,7 @@ from decimal import Decimal from unittest.mock import patch, MagicMock import logging -from basic_info.models import Product, WareHouse, Supplier, ProductUnitEnum, ProductCategory, Merchant, MerchantTypeEnum +from basic_info.models import Product, WareHouse, Supplier, ProductUnitEnum, ProductCategory, Merchant, MerchantTypeEnum, WareHouseModeEnum from . import models, services @@ -437,6 +437,191 @@ class CreateStockChangeRecordWithDetailsTestCase(StockServicesTestCase): source_id=self.purchase_order.id, products=products, ) + + def test_strict_mode_requires_restrict_in(self): + """严谨模式仅允许在严进宽出仓库中使用""" + self.warehouse_main.mode = WareHouseModeEnum.UNRESTRICTED + self.warehouse_main.save() + + with self.assertRaises(ValueError) as ctx: + 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=[{'product': self.product_fabric_a.id, 'quantity': [Decimal('10.00')]}], + ) + + self.assertIn('宽进宽出', str(ctx.exception)) + + def test_strict_mode_not_implemented_for_restrict_in_out(self): + """严进严出模式暂不支持""" + self.warehouse_main.mode = WareHouseModeEnum.RESTRICT_IN_OUT + self.warehouse_main.save() + + with self.assertRaises(ValueError) as ctx: + 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=[{'product': self.product_fabric_a.id, 'quantity': [Decimal('10.00')]}], + ) + + self.assertIn('严进严出', str(ctx.exception)) + + +class CreateStockChangeRecordRelaxedTestCase(StockServicesTestCase): + """测试宽松模式的库存变动创建""" + + def setUp(self): + super().setUp() + self.warehouse_main.mode = WareHouseModeEnum.UNRESTRICTED + self.warehouse_main.save() + + def test_relaxed_even_split(self): + products = [ + { + 'product': self.product_fabric_a.id, + 'quantity': {'value': Decimal('100.00'), 'unit_count': Decimal('25.00')}, + } + ] + + record, details, created_count = services.create_stock_change_record_relaxed( + 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, + ) + + self.assertEqual(created_count, 4) + self.assertEqual(record.details.count(), 4) + self.assertListEqual( + [detail.quantity for detail in details], + [Decimal('25.00')] * 4, + ) + + def test_relaxed_with_remainder(self): + products = [ + { + 'product': self.product_fabric_a.id, + 'quantity': {'value': Decimal('105.00'), 'unit_count': Decimal('30.00')}, + } + ] + + _, details, created_count = services.create_stock_change_record_relaxed( + 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, + ) + + self.assertEqual(created_count, 4) + self.assertListEqual( + [detail.quantity for detail in details], + [Decimal('30.00'), Decimal('30.00'), Decimal('30.00'), Decimal('15.00')], + ) + + def test_relaxed_invalid_unit_count(self): + products = [ + { + 'product': self.product_fabric_a.id, + 'quantity': {'value': Decimal('50.00'), 'unit_count': Decimal('0')}, + } + ] + + with self.assertRaises(ValueError): + services.create_stock_change_record_relaxed( + 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_relaxed_auto_complete_triggers_completion(self): + self.merchant.auto_complete_stock_change = True + self.merchant.save() + + products = [ + { + 'product': self.product_fabric_a.id, + 'quantity': {'value': Decimal('10.00'), 'unit_count': Decimal('2.00')}, + } + ] + + with patch('stock.services.make_stock_change_completed') as mock_complete: + services.create_stock_change_record_relaxed( + 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, + ) + + mock_complete.assert_called_once() + + def test_relaxed_mode_requires_unrestricted(self): + self.warehouse_main.mode = WareHouseModeEnum.RESTRICT_IN + self.warehouse_main.save() + + products = [ + { + 'product': self.product_fabric_a.id, + 'quantity': {'value': Decimal('10.00'), 'unit_count': Decimal('2.00')}, + } + ] + + with self.assertRaises(ValueError) as ctx: + services.create_stock_change_record_relaxed( + 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, + ) + + self.assertIn('严进宽出', str(ctx.exception)) + + def test_relaxed_mode_not_implemented_for_restrict_in_out(self): + self.warehouse_main.mode = WareHouseModeEnum.RESTRICT_IN_OUT + self.warehouse_main.save() + + products = [ + { + 'product': self.product_fabric_a.id, + 'quantity': {'value': Decimal('10.00'), 'unit_count': Decimal('2.00')}, + } + ] + + with self.assertRaises(ValueError) as ctx: + services.create_stock_change_record_relaxed( + 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, + ) + + self.assertIn('严进严出', str(ctx.exception)) def test_complete_outbound_record(self): """测试完成出库记录"""