from datetime import date from django.core.exceptions import ValidationError as DjangoValidationError from django.shortcuts import get_object_or_404 from rest_framework import permissions, serializers, status from rest_framework.parsers import FormParser, JSONParser, MultiPartParser from rest_framework.response import Response from rest_framework.views import APIView from basic_info.models import Merchant from cost import models as cost_models from cost import services as cost_services def _get_employee(request): employee = getattr(request.user, 'employee', None) if employee is None: raise serializers.ValidationError('当前用户未关联员工') return employee def _category_payload(category: cost_models.CostCategory) -> dict: return { 'id': category.id, 'merchant_id': category.merchant_id, 'unique_key': category.unique_key, 'name': category.name, 'parent_id': category.parent_id, 'description': category.description, 'created_at': category.created_at.isoformat().replace('+00:00', 'Z'), 'updated_at': category.updated_at.isoformat().replace('+00:00', 'Z'), } def _entry_payload(entry: cost_models.CostEntry) -> dict: return { 'id': entry.id, 'merchant_id': entry.merchant_id, 'category_id': entry.category_id, 'category_name': entry.category.name if entry.category_id else '', 'amount': str(entry.amount), 'unit_amount': str(entry.unit_amount) if entry.unit_amount is not None else None, 'quantity': str(entry.quantity) if entry.quantity is not None else None, 'unit_name': entry.unit_name or '', 'occurred_at': entry.occurred_at.isoformat(), 'operator_id': entry.operator_id, 'image1': _image_url(entry.image1), 'image2': _image_url(entry.image2), 'source_module': entry.source_module or '', 'source_id': entry.source_id or '', 'remarks': entry.remarks or '', 'created_at': entry.created_at.isoformat().replace('+00:00', 'Z'), 'updated_at': entry.updated_at.isoformat().replace('+00:00', 'Z'), } def _image_url(image_field) -> str: if not image_field: return '' try: return image_field.url except (ValueError, AttributeError): return '' # ==================== Category Write Serializers ==================== class CategoryWriteSerializer(serializers.Serializer): unique_key = serializers.CharField(max_length=100) name = serializers.CharField(max_length=100) parent_id = serializers.IntegerField(required=False, allow_null=True) description = serializers.CharField(required=False, allow_blank=True, default='') class CategoryUpdateSerializer(serializers.Serializer): unique_key = serializers.CharField(required=False, max_length=100) name = serializers.CharField(required=False, max_length=100) parent_id = serializers.IntegerField(required=False, allow_null=True) description = serializers.CharField(required=False, allow_blank=True) # ==================== Entry Write Serializers ==================== class EntryWriteSerializer(serializers.Serializer): category_id = serializers.IntegerField(min_value=1) amount = serializers.DecimalField( required=False, allow_null=True, max_digits=15, decimal_places=2, min_value=0, ) unit_amount = serializers.DecimalField( required=False, allow_null=True, max_digits=15, decimal_places=4, min_value=0, ) quantity = serializers.DecimalField( required=False, allow_null=True, max_digits=12, decimal_places=4, min_value=0, ) unit_name = serializers.CharField(required=False, allow_blank=True, max_length=20, default='') occurred_at = serializers.DateField() image1 = serializers.ImageField(required=False, allow_null=True) image2 = serializers.ImageField(required=False, allow_null=True) source_module = serializers.CharField(required=False, allow_blank=True, default='') source_id = serializers.CharField(required=False, allow_blank=True, default='') remarks = serializers.CharField(required=False, allow_blank=True, default='') def validate(self, attrs): unit_amount = attrs.get('unit_amount') quantity = attrs.get('quantity') has_unit_amount = unit_amount is not None has_quantity = quantity is not None if has_unit_amount != has_quantity: raise serializers.ValidationError('unit_amount 和 quantity 必须同时填写或同时为空') if not has_unit_amount and attrs.get('amount') is None: raise serializers.ValidationError('普通支出必须填写 amount;倍数型支出必须填写 unit_amount 和 quantity') return attrs class EntryUpdateSerializer(serializers.Serializer): category_id = serializers.IntegerField(required=False, min_value=1) amount = serializers.DecimalField( required=False, allow_null=True, max_digits=15, decimal_places=2, min_value=0, ) unit_amount = serializers.DecimalField( required=False, allow_null=True, max_digits=15, decimal_places=4, min_value=0, ) quantity = serializers.DecimalField( required=False, allow_null=True, max_digits=12, decimal_places=4, min_value=0, ) unit_name = serializers.CharField(required=False, allow_blank=True, max_length=20) occurred_at = serializers.DateField(required=False) image1 = serializers.ImageField(required=False, allow_null=True) image2 = serializers.ImageField(required=False, allow_null=True) source_module = serializers.CharField(required=False, allow_blank=True) source_id = serializers.CharField(required=False, allow_blank=True) remarks = serializers.CharField(required=False, allow_blank=True) # ==================== Category Views ==================== class CostCategoryListCreateView(APIView): permission_classes = [permissions.IsAuthenticated] def get(self, request): employee = _get_employee(request) merchant_id = request.query_params.get('merchant_id') queryset = cost_models.CostCategory.objects.filter( merchant_id=merchant_id if merchant_id else employee.merchant_id, ).order_by('name') return Response([_category_payload(c) for c in queryset]) def post(self, request): employee = _get_employee(request) serializer = CategoryWriteSerializer(data=request.data) serializer.is_valid(raise_exception=True) data = serializer.validated_data category = cost_models.CostCategory.objects.create( merchant=employee.merchant, unique_key=data['unique_key'], name=data['name'], parent_id=data.get('parent_id'), description=data.get('description', ''), ) return Response(_category_payload(category), status=status.HTTP_201_CREATED) class CostCategoryDetailView(APIView): permission_classes = [permissions.IsAuthenticated] def get(self, request, category_id): employee = _get_employee(request) category = get_object_or_404( cost_models.CostCategory, id=category_id, merchant_id=employee.merchant_id, ) return Response(_category_payload(category)) def put(self, request, category_id): employee = _get_employee(request) category = get_object_or_404( cost_models.CostCategory, id=category_id, merchant_id=employee.merchant_id, ) serializer = CategoryUpdateSerializer(data=request.data, partial=True) serializer.is_valid(raise_exception=True) data = serializer.validated_data for field in ('unique_key', 'name', 'parent_id', 'description'): if field in data: setattr(category, field, data[field]) category.save(update_fields=[k for k in data if k in data]) return Response(_category_payload(category)) def delete(self, request, category_id): employee = _get_employee(request) category = get_object_or_404( cost_models.CostCategory, id=category_id, merchant_id=employee.merchant_id, ) category.delete() return Response(status=status.HTTP_204_NO_CONTENT) # ==================== Entry Views ==================== class CostEntryListCreateView(APIView): permission_classes = [permissions.IsAuthenticated] parser_classes = [JSONParser, FormParser, MultiPartParser] def get(self, request): employee = _get_employee(request) merchant_id = request.query_params.get('merchant_id') category_id = request.query_params.get('category_id') start = request.query_params.get('start') end = request.query_params.get('end') queryset = cost_models.CostEntry.objects.filter( merchant_id=merchant_id if merchant_id else employee.merchant_id, ).select_related('category').order_by('-occurred_at', '-created_at') if category_id: queryset = queryset.filter(category_id=category_id) if start: queryset = queryset.filter(occurred_at__gte=start) if end: queryset = queryset.filter(occurred_at__lte=end) return Response([_entry_payload(e) for e in queryset]) def post(self, request): employee = _get_employee(request) serializer = EntryWriteSerializer(data=request.data) serializer.is_valid(raise_exception=True) data = serializer.validated_data category = get_object_or_404( cost_models.CostCategory, id=data['category_id'], merchant_id=employee.merchant_id, ) try: entry = cost_services.create_cost_entry( merchant=employee.merchant, category=category, occurred_at=data['occurred_at'], amount=data.get('amount'), unit_amount=data.get('unit_amount'), quantity=data.get('quantity'), unit_name=data.get('unit_name', ''), operator=employee, image1=data.get('image1'), image2=data.get('image2'), source_module=data.get('source_module', ''), source_id=data.get('source_id', ''), remarks=data.get('remarks', ''), ) except DjangoValidationError as exc: raise serializers.ValidationError(exc.message_dict if hasattr(exc, 'message_dict') else exc.messages) return Response(_entry_payload(entry), status=status.HTTP_201_CREATED) class CostEntryDetailView(APIView): permission_classes = [permissions.IsAuthenticated] parser_classes = [JSONParser, FormParser, MultiPartParser] def get(self, request, entry_id): employee = _get_employee(request) entry = get_object_or_404( cost_models.CostEntry.objects.select_related('category'), id=entry_id, merchant_id=employee.merchant_id, ) return Response(_entry_payload(entry)) def put(self, request, entry_id): employee = _get_employee(request) entry = get_object_or_404( cost_models.CostEntry, id=entry_id, merchant_id=employee.merchant_id, ) serializer = EntryUpdateSerializer(data=request.data, partial=True) serializer.is_valid(raise_exception=True) data = serializer.validated_data update_fields = [] if 'category_id' in data: entry.category = get_object_or_404( cost_models.CostCategory, id=data.pop('category_id'), merchant_id=employee.merchant_id, ) update_fields.append('category') for field, value in data.items(): setattr(entry, field, value) update_fields.append(field) if update_fields: try: entry.save(update_fields=update_fields) except DjangoValidationError as exc: raise serializers.ValidationError(exc.message_dict if hasattr(exc, 'message_dict') else exc.messages) entry.refresh_from_db() return Response(_entry_payload(entry)) def delete(self, request, entry_id): employee = _get_employee(request) entry = get_object_or_404( cost_models.CostEntry, id=entry_id, merchant_id=employee.merchant_id, ) entry.delete() return Response(status=status.HTTP_204_NO_CONTENT) # ==================== Summary Views ==================== class CostSummaryByCategoryView(APIView): permission_classes = [permissions.IsAuthenticated] def get(self, request): employee = _get_employee(request) merchant_id = request.query_params.get('merchant_id') start = request.query_params.get('start') end = request.query_params.get('end') if merchant_id: merchant = get_object_or_404(Merchant, id=merchant_id) else: merchant = employee.merchant start_date = date.fromisoformat(start) if start else None end_date = date.fromisoformat(end) if end else None aggregated = cost_services.aggregate_by_category( merchant=merchant, start_date=start_date, end_date=end_date, ) return Response({ 'merchant_id': merchant.id, 'start_date': str(start_date) if start_date else None, 'end_date': str(end_date) if end_date else None, 'results': aggregated, })