优化Django Rest多层关联查询集:解决海量数据库查询问题
Django多级关联查询优化问题
我有一个Django Rest视图,需要获取3级、4级关联对象的数据。视图代码如下:
class OrderViewSet(ModelViewSet): queryset = Order.objects.select_related().prefetch_related().order_by('-id')
需要在序列化器中获取订单关联商品的库存并计算,序列化器代码如下:
class OrderListSerializer(ModelSerializer): shortages = SerializerMethodField() def get_shortages(self, obj): item = obj.order_line.product.item stock = ( item.item_instances.filter(stock_quantity__gt=0) .aggregate(stock=Coalesce(Sum("stock_quantity"), 0, output_field=FloatField())) .get("stock", 0) ) required_stock = self.calculate_required_stock() if stock < required_stock: # 后续逻辑省略 ...
注:原代码中order_line.product.item.item_instances存在变量引用错误,已修正为item.item_instances
上述代码中order_line.product.item.item_instances涉及4个模型,即使在查询集的prefetch_related中添加('order_line__product__item')和('order_line__product__item__item_instances'),仍会产生上千条数据库查询。如何优化该场景?
补充示例模型
class Order(models.Model): order_line = models.ForeignKey( "OrderLine", on_delete=models.CASCADE, related_name="orders", ) type = models.CharField(choices=TYPE_CHOICES, max_length=13, blank=True) start_date = models.DateField(blank=True, null=True) # 其他字段省略 class OrderLine(models.Model): sales_order = models.ForeignKey( "SalesOrder", on_delete=models.CASCADE, related_name="sales_order_lines", ) product = models.ForeignKey("Product", on_delete=models.PROTECT) quantities = models.JSONField(null=True, blank=True) # 其他字段省略 class Product(models.Model): design = models.ForeignKey( Design, on_delete=models.PROTECT, related_name="design_products" ) item = models.ForeignKey( Item, on_delete=models.PROTECT, related_name="item_products" ) # 其他字段省略 class Item(models.Model): shrinkage = models.FloatField(default=0.0) # 其他字段省略 class ItemInstance(models.Model): item = models.ForeignKey( Item, on_delete=models.PROTECT, related_name="item_instances" ) stock_quantity = models.FloatField( default=0.0, blank=True, validators=[MinValueValidator(0)] )
优化方案
核心问题在于序列化器的get_shortages方法中,每个订单都单独执行一次库存聚合查询,导致N+1性能问题。以下是几种有效的优化方式:
1. 子查询预计算库存并注入查询集
通过子查询提前计算每个Item的可用库存,直接附加到Order查询结果中,序列化时直接取值即可。
步骤1:定义库存聚合子查询
from django.db.models import Subquery, OuterRef, Sum, Coalesce, FloatField # 子查询:计算单个Item的可用库存 stock_subquery = ItemInstance.objects.filter( item=OuterRef('pk'), stock_quantity__gt=0 ).aggregate( total_stock=Coalesce(Sum('stock_quantity'), 0, output_field=FloatField()) ).values('total_stock')
步骤2:修改视图查询集
class OrderViewSet(ModelViewSet): queryset = Order.objects.select_related( 'order_line__product__item' ).annotate( # 将子查询结果作为字段注入Order对象 item_stock=Subquery(stock_subquery) ).order_by('-id')
步骤3:序列化器直接使用预计算值
class OrderListSerializer(ModelSerializer): shortages = SerializerMethodField() def get_shortages(self, obj): stock = obj.item_stock required_stock = self.calculate_required_stock() if stock < required_stock: # 执行你的短缺计算逻辑 return required_stock - stock return 0
2. 使用Prefetch传递自定义查询集
通过Prefetch对象为关联的Item传递带库存聚合的查询集,一次性获取所有需要的数据。
from django.db.models import Prefetch, Q # 自定义Item查询集:提前聚合库存 item_queryset = Item.objects.annotate( total_stock=Coalesce( Sum('item_instances__stock_quantity', filter=Q(item_instances__stock_quantity__gt=0)), 0, output_field=FloatField() ) ) class OrderViewSet(ModelViewSet): queryset = Order.objects.select_related( 'order_line__product' ).prefetch_related( # 使用自定义查询集预取Item数据 Prefetch('order_line__product__item', queryset=item_queryset) ).order_by('-id')
此时序列化器中可直接访问obj.order_line.product.item.total_stock,无需额外查询。
3. 批量查询库存并缓存
在序列化器初始化时,批量查询所有关联Item的库存,存入字典缓存,后续直接取值。
class OrderListSerializer(ModelSerializer): shortages = SerializerMethodField() def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) # 批量提取所有需要的Item ID item_ids = [obj.order_line.product.item.id for obj in self.instance] # 批量查询库存并转为字典 item_stocks = Item.objects.filter(id__in=item_ids).annotate( total_stock=Coalesce( Sum('item_instances__stock_quantity', filter=Q(item_instances__stock_quantity__gt=0)), 0, output_field=FloatField() ) ).values_list('id', 'total_stock') self.item_stock_map = dict(item_stocks) def get_shortages(self, obj): item_id = obj.order_line.product.item.id stock = self.item_stock_map.get(item_id, 0) required_stock = self.calculate_required_stock() if stock < required_stock: return required_stock - stock return 0
这种方式将N+1查询压缩为2次查询,性能提升显著。
内容的提问来源于stack exchange,提问作者Sardar Faisal
相关产品推荐
相关产品推荐

