Django一对多模型中Case的total_price字段更新异常求助
问题描述
我是Django新手,在models.py中定义了Case和CaseItem两个模型,二者为一对多关系(一个Case可关联多个CaseItem)。我希望创建CaseItem并关联Case时,自动更新Case的total_price字段,但目前创建Case后该字段始终为默认值0。相关代码如下:
模型代码
class Item(models.Model): name = models.TextField(unique=True) # Price is in cents price = models.PositiveIntegerField( validators=[ MinValueValidator(1) ] ) img = models.TextField() class Case(models.Model): id = models.UUIDField(primary_key=True, default=uuid4) name = models.TextField(unique=True) case_img = models.TextField() total_price = models.PositiveIntegerField(default=0) class CaseItem(models.Model): case = models.ForeignKey(Case, on_delete=models.CASCADE, related_name="items") item = models.ForeignKey(Item, on_delete=models.PROTECT) percentage = models.PositiveIntegerField( validators=[ MaxValueValidator(100), MinValueValidator(1) ] ) # ensure that user cant add the same item to the case class Meta: unique_together = [['case', 'item']] def save(self, *args, **kwargs): self.case.total_price += self.item.price self.case.save() super().save(*args, **kwargs)
创建视图代码
class CreateCustomCase(generics.CreateAPIView): serializer_class = CustomCaseSerializer def create(self, request, *args, **kwargs): data = request.data.copy() serializer = self.get_serializer(data=data, context={'request': request}) if serializer.is_valid(raise_exception=True): serializer.save() return Response(serializer.validated_data, status.HTTP_201_CREATED) return Response(serializer.errors, status.HTTP_400_BAD_REQUEST)
序列化器代码
class CustomCaseSerializer(serializers.Serializer): name = serializers.CharField(max_length=None) case_img = serializers.CharField(max_length=None) items = CustomCaseItemSerializer(many=True) def validate(self, attrs): items = attrs.get("items") case_name = attrs.get("name") try: caseObj = Case.objects.get(name=case_name) except Case.DoesNotExist: caseObj = None if caseObj is not None: raise serializers.ValidationError("Case with the same alredy exists") sumPercentage = sum([item["percentage"] for item in items]) if sumPercentage != 100: raise serializers.ValidationError("Total Sum of percentage should be 100%") itemIdList = [item["id"] for item in items] selectIds = Item.objects.filter(pk__in=itemIdList).values('id') selectIdList = [selectId["id"] for selectId in selectIds] if len(selectIdList) != len(itemIdList): notInSelect = set(itemIdList).difference(selectIdList) raise serializers.ValidationError( f'Did not find the following items with id in the database: {notInSelect}'
错误分析
- 自定义序列化器缺失
create方法:使用serializers.Serializer而非ModelSerializer时,必须手动实现create方法来处理Case和CaseItem的创建逻辑,否则默认save方法不会自动创建关联的CaseItem,自然无法触发CaseItem的save方法更新total_price。 CaseItem.save()逻辑存在漏洞:直接修改内存中的self.case对象累加价格,可能因数据库缓存导致数据不一致;后续若更新CaseItem关联的Item,还会出现重复累加的问题。validate方法代码不完整:提供的validate方法缺少闭合的},且未返回验证后的attrs,会导致代码报错或序列化器丢失有效数据。
解决方案
1. 完善序列化器的create方法
在CustomCaseSerializer中添加create方法,先创建Case,再批量创建关联的CaseItem,确保触发CaseItem的save逻辑:
class CustomCaseSerializer(serializers.Serializer): name = serializers.CharField(max_length=None) case_img = serializers.CharField(max_length=None) items = CustomCaseItemSerializer(many=True) # 补全并修正原validate方法 def validate(self, attrs): items = attrs.get("items") case_name = attrs.get("name") try: Case.objects.get(name=case_name) raise serializers.ValidationError("Case with the same name already exists") except Case.DoesNotExist: pass sum_percentage = sum([item["percentage"] for item in items]) if sum_percentage != 100: raise serializers.ValidationError("Total Sum of percentage should be 100%") item_id_list = [item["id"] for item in items] existing_item_ids = list(Item.objects.filter(pk__in=item_id_list).values_list('id', flat=True)) missing_ids = set(item_id_list) - set(existing_item_ids) if missing_ids: raise serializers.ValidationError( f'Did not find the following items with id in the database: {missing_ids}' ) return attrs def create(self, validated_data): # 提取items数据,先创建Case items_data = validated_data.pop('items') case = Case.objects.create(**validated_data) # 批量创建CaseItem,触发save方法更新total_price for item_data in items_data: item = Item.objects.get(pk=item_data['id']) CaseItem.objects.create( case=case, item=item, percentage=item_data['percentage'] ) # 刷新Case对象,获取最新的total_price case.refresh_from_db() return case
2. 修复CaseItem.save()逻辑
使用F表达式直接在数据库层面操作,避免内存对象与数据库数据不一致的问题,同时支持更新场景的价格修正:
class CaseItem(models.Model): # 保留原有字段和Meta类 def save(self, *args, **kwargs): if self.pk is None: # 新建CaseItem时,累加对应Item的价格 Case.objects.filter(pk=self.case.pk).update( total_price=models.F('total_price') + self.item.price ) else: # 更新CaseItem时,先减去旧Item价格,再加新Item价格 old_item = CaseItem.objects.only('item').get(pk=self.pk).item if old_item != self.item: Case.objects.filter(pk=self.case.pk).update( total_price=models.F('total_price') - old_item.price + self.item.price ) super().save(*args, **kwargs)
内容的提问来源于stack exchange,提问作者Ujjwal Sharma
相关产品推荐
相关产品推荐

