Django/DRF中如何按HTTP请求方法校验request.data字段合法性?
问题
需要实现以下需求:
- 向
127.0.0.1:8000/api/v1/events/发送POST请求时,校验request.data的正确性 - 向
127.0.0.1:8000/api/v1/events/{pk}发送PATCH请求时,确保request.data中不包含updated_time、event_type字段(禁止更新这两个字段)
已有Event模型:
class Event(models.Model): title = models.CharField(max_length=50) description = models.CharField(max_length=250) created_by = models.ForeignKey(User, related_name='created_events', on_delete=models.CASCADE) event_type = models.CharField(choices=EventType.choices, max_length=10) created_time = models.DateTimeField(auto_now_add=True) updated_time = models.DateTimeField(auto_now=True)
目前使用的校验代码无法区分请求方法的差异,现有视图代码:
def create(self, request): serializer = EventSerializer(data=request.data, context={'request': request}) serializer.is_valid(raise_exception=True) error_response = self._validate_request_data(request.data) if error_response is not None: return error_response user, event_id = request.user, request.data['event'] event = get_object_or_404(Event, pk=event_id) ... ... def partial_update(self, request): error_response = self._validate_request_data(request.data) if error_response is not None: return error_response ... ... def _validate_request_data(self, request_data): request_keys = request_data.keys() actual_keys = self.serializer_class().get_fields().keys() result = [x for x in request_keys if x not in actual_keys] if result: return Response({'error': 'Invalid data'}, status=400) ... ...
问题在于.get_fields()返回序列化器所有字段,无法区分POST和PATCH的校验规则,如何针对不同HTTP方法实现差异化的request.data字段校验?
解决方案
方法1:给校验方法传入请求方法参数
修改_validate_request_data方法,让它接收当前的请求方法,根据不同方法执行不同校验逻辑:
def create(self, request): serializer = EventSerializer(data=request.data, context={'request': request}) serializer.is_valid(raise_exception=True) error_response = self._validate_request_data(request.data, request.method) if error_response is not None: return error_response # 后续业务逻辑... def partial_update(self, request): error_response = self._validate_request_data(request.data, request.method) if error_response is not None: return error_response # 后续业务逻辑... def _validate_request_data(self, request_data, method): request_keys = request_data.keys() actual_keys = self.serializer_class().get_fields().keys() # POST请求:校验所有传入字段都是序列化器允许的字段 if method == 'POST': invalid_keys = [x for x in request_keys if x not in actual_keys] if invalid_keys: return Response({'error': f'Invalid fields: {", ".join(invalid_keys)}'}, status=400) # PATCH请求:禁止传入updated_time和event_type,同时校验其他字段合法性 elif method == 'PATCH': forbidden_keys = ['updated_time', 'event_type'] invalid_forbidden = [x for x in request_keys if x in forbidden_keys] if invalid_forbidden: return Response({'error': f'Cannot update fields: {", ".join(invalid_forbidden)}'}, status=400) allowed_keys = [key for key in actual_keys if key not in forbidden_keys] invalid_other = [x for x in request_keys if x not in allowed_keys] if invalid_other: return Response({'error': f'Invalid fields: {", ".join(invalid_other)}'}, status=400) return None
方法2:使用不同的序列化器类
为POST和PATCH分别定义序列化器,通过视图的get_serializer_class方法根据请求方法返回对应序列化器,更符合DRF的设计理念:
首先定义两个序列化器:
from rest_framework import serializers class EventCreateSerializer(serializers.ModelSerializer): class Meta: model = Event fields = ['title', 'description', 'event_type'] # POST允许提交的字段,created_by后续通过上下文赋值 class EventUpdateSerializer(serializers.ModelSerializer): class Meta: model = Event fields = ['title', 'description'] # 只允许更新这两个字段,自动排除禁止更新的字段
然后在视图中重写get_serializer_class:
def get_serializer_class(self): if self.action == 'create': return EventCreateSerializer elif self.action == 'partial_update': return EventUpdateSerializer return super().get_serializer_class()
视图逻辑简化,直接依赖序列化器的is_valid完成校验:
def create(self, request): serializer = self.get_serializer(data=request.data, context={'request': request}) serializer.is_valid(raise_exception=True) serializer.validated_data['created_by'] = request.user event = serializer.save() # 后续业务逻辑... def partial_update(self, request, pk=None): event = get_object_or_404(Event, pk=pk) serializer = self.get_serializer(event, data=request.data, partial=True) serializer.is_valid(raise_exception=True) serializer.save() # 后续业务逻辑...
方法3:在单个序列化器中根据上下文判断请求方法
如果不想拆分序列化器,可以在序列化器的validate方法中通过上下文获取请求方法,动态添加校验规则:
class EventSerializer(serializers.ModelSerializer): class Meta: model = Event fields = ['title', 'description', 'event_type', 'created_by', 'created_time', 'updated_time'] read_only_fields = ['created_by', 'created_time', 'updated_time'] def validate(self, data): request = self.context.get('request') if request and request.method == 'PATCH': # PATCH请求时检查是否包含禁止更新的字段 forbidden_fields = ['event_type', 'updated_time'] for field in forbidden_fields: if field in data: raise serializers.ValidationError(f"Field '{field}' cannot be updated.") return data
视图中直接调用序列化器的校验即可:
def create(self, request): serializer = EventSerializer(data=request.data, context={'request': request}) serializer.is_valid(raise_exception=True) serializer.validated_data['created_by'] = request.user event = serializer.save() # 后续业务逻辑... def partial_update(self, request, pk=None): event = get_object_or_404(Event, pk=pk) serializer = EventSerializer(event, data=request.data, context={'request': request}, partial=True) serializer.is_valid(raise_exception=True) serializer.save() # 后续业务逻辑...
内容的提问来源于stack exchange,提问作者nope_that_does_not_work
相关产品推荐
相关产品推荐

