如何在Django REST Framework中统计父分类下的子项总数
问题
我是Django新手,现有如下Category模型:
class Category(models.Model): title = models.CharField(max_length=200) parent = models.ForeignKey('self', on_delete=models.CASCADE, null=True, blank=True, related_name='children')
分类结构示例:
- 父分类 - foods
- 子父分类 - Italian food、French food等
- 最终子项:pizza、Polenta、Lasagna(共3个)
- 父分类 - drinks
- 子父分类 - juice、alcohol
- 最终子项:water、tea、lemon tea(共4个)
请问能否在序列化器中统计出foods的子项总数为3,drinks的为4?
解决方案
当然可以实现,以下是几种实用方案:
方法一:ORM递归CTE统计+序列化器字段
这种方式通过数据库递归查询实现统计,性能较好,适合数据量较大的场景:
1. 视图层查询根分类并统计叶子节点数
from django.db.models import IntegerField from django.db.models.expressions import RawSQL from django.db.models.functions import Coalesce from django.db.models import OuterRef # 递归CTE统计每个根分类下的叶子节点数量 root_categories = Category.objects.filter(parent__isnull=True).annotate( leaf_count=Coalesce( RawSQL( """ WITH RECURSIVE category_tree AS ( SELECT id, parent_id, title FROM category WHERE parent_id = %s UNION ALL SELECT c.id, c.parent_id, c.title FROM category c JOIN category_tree ct ON c.parent_id = ct.id ) SELECT COUNT(*) FROM category_tree WHERE NOT EXISTS (SELECT 1 FROM category WHERE parent_id = category_tree.id) """, [OuterRef('id')], output_field=IntegerField() ), 0 ) )
2. 序列化器添加统计字段
from rest_framework import serializers class CategorySerializer(serializers.ModelSerializer): leaf_count = serializers.IntegerField(read_only=True) class Meta: model = Category fields = ['title', 'leaf_count']
方法二:序列化器自定义方法字段
适合数据量较小的场景,实现简单但存在N+1查询问题:
from rest_framework import serializers class CategorySerializer(serializers.ModelSerializer): leaf_count = serializers.SerializerMethodField() def get_leaf_count(self, obj): # 递归统计当前分类下所有无子类的叶子节点 def count_leaves(category): if not category.children.exists(): return 1 total = 0 for child in category.children.all(): total += count_leaves(child) return total return count_leaves(obj) class Meta: model = Category fields = ['title', 'leaf_count']
方法三:反向查询叶子节点并分组统计
先找出所有叶子节点,再追溯其根分类并统计数量:
from django.db.models import Subquery, OuterRef, Count from django.db.models.expressions import RawSQL # 子查询:获取每个节点的根分类ID root_id_subquery = Category.objects.filter( id=OuterRef('id') ).annotate( root_id=RawSQL( """ WITH RECURSIVE cte AS ( SELECT id, parent_id FROM category WHERE id = %s UNION ALL SELECT c.id, c.parent_id FROM category c JOIN cte ct ON c.id = ct.parent_id ) SELECT id FROM cte WHERE parent_id IS NULL """, [OuterRef('id')], ) ).values('root_id') # 统计每个根分类下的叶子节点数 leaf_counts = Category.objects.filter(children__isnull=True).annotate( root_id=Subquery(root_id_subquery) ).values('root_id').annotate(count=Count('id')) # 将统计结果关联到根分类对象 root_categories = Category.objects.filter(parent__isnull=True) for category in root_categories: category.leaf_count = next((item['count'] for item in leaf_counts if item['root_id'] == category.id), 0)
之后使用包含leaf_count字段的序列化器即可返回统计结果。
内容的提问来源于stack exchange,提问作者Xursandbek Ibadullayev
相关产品推荐
相关产品推荐

