Django REST Framework中序列化过滤当前用户关联字段
问题描述
在项目中定义了三个模型:Group、User和Challenge。每个用户是若干分组的成员,每个挑战面向一个或多个分组开放。
模型定义代码如下:
class Group(TimeStampedModel): name = models.CharField(max_length=255) class User(AbstractUser): user_groups = models.ManyToManyField(Group) class Challenge(TimeStampedModel): ... groups = models.ManyToManyField(Group, null=True, blank=True)
为Challenge模型编写了序列化器,通过GroupSerializer序列化全部挑战数据及关联的分组信息:
class ChallengeSerializer(serializers.ModelSerializer): groups = GroupSerializer(many=True) class Meta: model = Challenge fields = [..., "groups"]
当前用于序列化挑战列表的APIView实现如下:
class ChallengeList(generics.ListAPIView): queryset = Challenge.objects.all() serializer_class = ChallengeSerializer permission_classes = [permissions.IsAuthenticated] pagination_class = PageNumberPagination def get_queryset(self): user_groups = self.request.user.user_groups.all() return Challenge.objects.filter(groups__in=user_groups).distinct()
目前序列化Challenge对象时,会将该对象关联的所有分组全部序列化返回,需要实现仅序列化同时与当前登录用户存在关联关系的相关Group对象。
解决方案
完全可以实现,不需要调整模型层关系,直接在序列化器层做动态过滤即可,两种常用实现方案如下:
方案1:使用SerializerMethodField自定义分组返回逻辑
这是最直观的实现方式,直接替换原有的嵌套序列化器字段,手动过滤符合要求的分组后再序列化:
class ChallengeSerializer(serializers.ModelSerializer): groups = serializers.SerializerMethodField() def get_groups(self, obj): # DRF通用视图会自动把request传入序列化上下文,直接取当前登录用户即可 current_user = self.context['request'].user # 过滤出当前挑战关联、且用户所属的分组 matched_groups = obj.groups.filter( id__in=current_user.user_groups.values_list('id', flat=True) ) return GroupSerializer(matched_groups, many=True).data class Meta: model = Challenge fields = [..., "groups"]
注意:DRF的所有通用视图(包括当前使用的
ListAPIView)默认会将request、view等对象传入序列化器的context属性,不需要额外手动传参。
方案2:重写to_representation方法过滤序列化结果
如果需要保留原有嵌套序列化器的字段定义(比如要复用字段的校验、写入逻辑),可以在序列化最终输出阶段过滤groups内容:
class ChallengeSerializer(serializers.ModelSerializer): groups = GroupSerializer(many=True) class Meta: model = Challenge fields = [..., "groups"] def to_representation(self, instance): rep = super().to_representation(instance) current_user = self.context['request'].user user_group_id_set = set(current_user.user_groups.values_list('id', flat=True)) # 只保留用户所属的分组数据 rep['groups'] = [group for group in rep['groups'] if group['id'] in user_group_id_set] return rep
性能优化建议
为了避免序列化分组时产生N+1查询问题,可以在视图的get_queryset中添加预加载,注意不要在预加载时加过滤条件,避免破坏原有查询的去重逻辑:
def get_queryset(self): user_groups = self.request.user.user_groups.all() return Challenge.objects.filter( groups__in=user_groups ).distinct().prefetch_related('groups')
不要尝试在ORM查询层直接过滤关联的分组,那样会导致Challenge的查询结果出现重复或者关联对象缺失的问题,序列化层过滤是最稳妥的实现方式。
内容的提问来源于stack exchange,提问作者drobilc
相关产品推荐
相关产品推荐

