diff --git a/.gitignore b/.gitignore index e0280c51b9..b07f42fc8c 100644 --- a/.gitignore +++ b/.gitignore @@ -13,6 +13,7 @@ /env/ MANIFEST coverage.* +venv/ .coverage .cache/ diff --git a/docs/api-guide/serializers.md b/docs/api-guide/serializers.md index 857a90c24a..1ce913c256 100644 --- a/docs/api-guide/serializers.md +++ b/docs/api-guide/serializers.md @@ -823,6 +823,8 @@ To support multiple updates you'll need to do so explicitly. When writing your m You will need to add an explicit `id` field to the instance serializer. The default implicitly-generated `id` field is marked as `read_only`. This causes it to be removed on updates. Once you declare it explicitly, it will be available in the list serializer's `update` method. +During validation, `ListSerializer` matches each input item to an existing instance using `id` or `pk`. To use another identifier, such as `uuid`, set `lookup_field` on the child serializer's `Meta` class. + Here's an example of how you might choose to implement multiple updates: class BookListSerializer(serializers.ListSerializer): @@ -855,14 +857,13 @@ Here's an example of how you might choose to implement multiple updates: class Meta: list_serializer_class = BookListSerializer + lookup_field = 'id' If the child serializer includes uniqueness validators (`UniqueValidator`, `UniqueTogetherValidator`, or the `UniqueForDateValidator` family), they need to know which object each item in the list is updating, so -that the object itself is not reported as a uniqueness conflict. By default the child serializer's -`.instance` is the whole queryset or list that was passed to the list serializer, so these validators will -raise a `RuntimeError` during a multiple update. To support this, override `run_child_validation()` on -your `ListSerializer` subclass to set the child's `.instance` and `.initial_data` for each item before -validation. For example, if `self.instance` is a queryset: +that the object itself is not reported as a uniqueness conflict. `ListSerializer` sets the child's +`.instance` and `.initial_data` for each matched item before validation. For custom matching behavior, +override `run_child_validation()` on your `ListSerializer` subclass. For example, if `self.instance` is a queryset: class BookListSerializer(serializers.ListSerializer): def run_child_validation(self, data): diff --git a/rest_framework/serializers.py b/rest_framework/serializers.py index fc8e83c768..bf5d8aab9e 100644 --- a/rest_framework/serializers.py +++ b/rest_framework/serializers.py @@ -664,7 +664,46 @@ def run_child_validation(self, data): self.child.initial_data = data return super().run_child_validation(data) """ - return self.child.run_validation(data) + if not hasattr(self.child, 'instance'): + return self.child.run_validation(data) + + if not ( + hasattr(self, '_list_serializer_instance_map') and + isinstance(data, Mapping) + ): + return self.child.run_validation(data) + + lookup_field = getattr(getattr(self.child, 'Meta', None), 'lookup_field', None) + original_instance = self.child.instance + if original_instance is not self.instance: + return self.child.run_validation(data) + + if lookup_field is not None: + data_pk = data.get(lookup_field) + else: + data_pk = data.get('id') + if data_pk is None: + data_pk = data.get('pk') + + child_instance = ( + self._list_serializer_instance_map.get(str(data_pk)) + if data_pk is not None else None + ) + + has_initial_data = hasattr(self.child, 'initial_data') + if has_initial_data: + original_initial_data = self.child.initial_data + + try: + self.child.instance = child_instance + self.child.initial_data = data + return self.child.run_validation(data) + finally: + self.child.instance = original_instance + if has_initial_data: + self.child.initial_data = original_initial_data + elif hasattr(self.child, 'initial_data'): + delattr(self.child, 'initial_data') def to_internal_value(self, data): """ @@ -702,28 +741,69 @@ def to_internal_value(self, data): ret = [] errors = {} - for index, item in enumerate(data): - try: - validated = self.run_child_validation(item) - except ValidationError as exc: - errors[index] = exc.detail + # Build a primary key lookup for instance matching in many=True updates. + instance_map = None + if self.instance is not None: + if isinstance(self.instance, Mapping): + instance_map = {str(k): v for k, v in self.instance.items()} else: - ret.append(validated) + instance_iterable = self.instance + if isinstance(instance_iterable, models.manager.BaseManager): + instance_iterable = instance_iterable.all() + if not isinstance(instance_iterable, (list, tuple, models.query.QuerySet)): + instance_iterable = None + + if instance_iterable is not None: + instance_map = {} + lookup_field = getattr(getattr(self.child, 'Meta', None), 'lookup_field', None) + + for obj in instance_iterable: + if lookup_field is not None: + lookup_values = [getattr(obj, lookup_field, None)] + else: + lookup_values = [ + getattr(obj, 'id', None), + getattr(obj, 'pk', None), + ] + + for lookup_value in lookup_values: + if lookup_value is not None: + instance_map[str(lookup_value)] = obj + + has_instance_map = hasattr(self, '_list_serializer_instance_map') + if has_instance_map: + original_instance_map = self._list_serializer_instance_map + if instance_map is not None: + self._list_serializer_instance_map = instance_map - if errors: - if not api_settings.LIST_SERIALIZER_ERRORS_AS_DICT: - warnings.warn( - 'The list-based error format for `ListSerializer` is ' - 'deprecated and will be removed in DRF 3.20. Set ' - '`REST_FRAMEWORK["LIST_SERIALIZER_ERRORS_AS_DICT"]` to ' - '`True` to use the dictionary-based error format.', - RemovedInDRF320Warning, - stacklevel=4, - ) - errors = [errors.get(index, {}) for index in range(len(data))] - raise ValidationError(errors) + try: + for index, item in enumerate(data): + try: + validated = self.run_child_validation(item) + except ValidationError as exc: + errors[index] = exc.detail + else: + ret.append(validated) + + if errors: + if not api_settings.LIST_SERIALIZER_ERRORS_AS_DICT: + warnings.warn( + 'The list-based error format for `ListSerializer` is ' + 'deprecated and will be removed in DRF 3.20. Set ' + '`REST_FRAMEWORK["LIST_SERIALIZER_ERRORS_AS_DICT"]` to ' + '`True` to use the dictionary-based error format.', + RemovedInDRF320Warning, + stacklevel=4, + ) + errors = [errors.get(index, {}) for index in range(len(data))] + raise ValidationError(errors) - return ret + return ret + finally: + if instance_map is not None and has_instance_map: + self._list_serializer_instance_map = original_instance_map + elif instance_map is not None and hasattr(self, '_list_serializer_instance_map'): + delattr(self, '_list_serializer_instance_map') def to_representation(self, data): """ @@ -758,6 +838,13 @@ def save(self, **kwargs): """ Save and return a list of object instances. """ + assert hasattr(self, '_errors'), ( + 'You must call `.is_valid()` before calling `.save()`.' + ) + assert not self.errors, ( + 'You cannot call `.save()` on a serializer with invalid data.' + ) + # Guard against incorrect use of `serializer.save(commit=False)` assert 'commit' not in kwargs, ( "'commit' is not a valid keyword argument to the 'save()' method. " @@ -765,9 +852,13 @@ def save(self, **kwargs): "inspect 'serializer.validated_data' instead. " "You can also pass additional keyword arguments to 'save()' if you " "need to set extra attributes on the saved model instance. " - "For example: 'serializer.save(owner=request.user)'.'" + "For example: 'serializer.save(owner=request.user)'." + ) + assert not hasattr(self, '_data'), ( + "You cannot call `.save()` after accessing `serializer.data`. " + "If you need to access data before committing to the database then " + "inspect 'serializer.validated_data' instead. " ) - validated_data = [ {**attrs, **kwargs} for attrs in self.validated_data ] diff --git a/tests/test_serializer_lists.py b/tests/test_serializer_lists.py index cbbdeaffd1..0c2d8d49db 100644 --- a/tests/test_serializer_lists.py +++ b/tests/test_serializer_lists.py @@ -207,6 +207,267 @@ def update(self, instance, validated_data): assert updated_instances == expected_output +class TestListSerializerInstanceMatching: + def test_matching_with_default_lookup_field(self): + seen_instances = [] + + class TestSerializer(serializers.Serializer): + pk = serializers.IntegerField() + + def validate(self, attrs): + seen_instances.append(self.instance) + return attrs + + instance = [ + BasicObject(pk=1), + BasicObject(pk=2), + ] + input_data = [ + {'pk': 1}, + {'pk': 2}, + ] + + serializer = TestSerializer(instance, data=input_data, many=True) + assert serializer.is_valid() + assert seen_instances == instance + + def test_matching_with_id_by_default(self): + seen_instances = [] + + class TestSerializer(serializers.Serializer): + id = serializers.IntegerField() + + def validate(self, attrs): + seen_instances.append(self.instance) + return attrs + + instance = [BasicObject(id=1), BasicObject(id=2)] + serializer = TestSerializer( + instance, data=[{'id': 1}, {'id': 2}], many=True + ) + + assert serializer.is_valid() + assert seen_instances == instance + + def test_field_validation_receives_item_initial_data(self): + seen_initial_data = [] + + class TestSerializer(serializers.Serializer): + pk = serializers.IntegerField() + + def validate_pk(self, value): + seen_initial_data.append(self.initial_data) + return value + + instance = [BasicObject(pk=1), BasicObject(pk=2)] + input_data = [{'pk': 1}, {'pk': 2}] + + serializer = TestSerializer(instance, data=input_data, many=True) + assert serializer.is_valid() + assert seen_initial_data == input_data + + def test_object_validation_receives_item_initial_data(self): + seen_initial_data = [] + + class TestSerializer(serializers.Serializer): + pk = serializers.IntegerField() + + def validate(self, attrs): + seen_initial_data.append(self.initial_data) + return attrs + + instance = [BasicObject(pk=1), BasicObject(pk=2)] + input_data = [{'pk': 1}, {'pk': 2}] + + serializer = TestSerializer(instance, data=input_data, many=True) + assert serializer.is_valid() + assert seen_initial_data == input_data + + def test_child_initial_data_state_is_restored_after_validation(self): + class TestSerializer(serializers.Serializer): + pk = serializers.IntegerField() + + instance = [BasicObject(pk=1)] + input_data = [{'pk': 1}] + serializer = TestSerializer(instance, data=input_data, many=True) + original_initial_data = serializer.child.initial_data + + assert serializer.is_valid() + assert serializer.child.initial_data is original_initial_data + + child = TestSerializer() + serializer = serializers.ListSerializer( + child=child, instance=instance, data=input_data + ) + + assert not hasattr(child, 'initial_data') + assert serializer.is_valid() + assert not hasattr(child, 'initial_data') + + def test_mapping_instance_matching(self): + seen_instances = [] + + class TestSerializer(serializers.Serializer): + pk = serializers.IntegerField() + + def validate(self, attrs): + seen_instances.append(self.instance) + return attrs + + obj1 = BasicObject(pk=1) + obj2 = BasicObject(pk=2) + instance = { + '1': obj1, + '2': obj2, + } + input_data = [ + {'pk': 1}, + {'pk': 2}, + ] + + serializer = TestSerializer(instance, data=input_data, many=True) + assert serializer.is_valid() + assert seen_instances == [obj1, obj2] + + def test_unsupported_instance_type_preserves_original_behavior(self): + seen_instances = [] + + class TestSerializer(serializers.Serializer): + pk = serializers.IntegerField() + + def validate(self, attrs): + seen_instances.append(self.instance) + return attrs + + serializer = TestSerializer(instance=123, data=[{'pk': 1}], many=True) + assert serializer.is_valid() + assert seen_instances == [123] + + def test_unmatched_instance_is_none(self): + seen_instances = [] + + class TestSerializer(serializers.Serializer): + id = serializers.IntegerField() + + def validate(self, attrs): + seen_instances.append(self.instance) + return attrs + + serializer = TestSerializer( + [BasicObject(id=1)], data=[{'id': 2}], many=True + ) + + assert serializer.is_valid() + assert seen_instances == [None] + + def test_custom_run_child_validation_instance_is_preserved(self): + seen_instances = [] + + class TestSerializer(serializers.Serializer): + id = serializers.IntegerField() + + def validate(self, attrs): + seen_instances.append(self.instance) + return attrs + + class TestListSerializer(serializers.ListSerializer): + def run_child_validation(self, data): + self.child.instance = 'custom instance' + return super().run_child_validation(data) + + serializer = TestListSerializer( + child=TestSerializer(), + instance=[BasicObject(id=1)], + data=[{'id': 1}], + ) + + assert serializer.is_valid() + assert seen_instances == ['custom instance'] + + @pytest.mark.django_db + def test_manager_instance_matching(self): + seen_instances = [] + + class TestSerializer(serializers.ModelSerializer): + def validate(self, attrs): + seen_instances.append(self.instance) + return attrs + + class Meta: + model = CustomManagerModel + fields = ['id'] + + o2o_target = OneToOneTarget.objects.create(name='target') + instance = CustomManagerModel.objects.create( + text='text', o2o_target=o2o_target + ) + serializer = TestSerializer( + CustomManagerModel.objects, + data=[{'id': instance.pk}], + many=True, + ) + + assert serializer.is_valid() + assert seen_instances == [instance] + + def test_missing_lookup_field_in_data_does_not_assign_instance(self): + seen_instances = [] + + class TestSerializer(serializers.Serializer): + id = serializers.IntegerField(required=False) + + class Meta: + lookup_field = 'uuid' + + def validate(self, attrs): + seen_instances.append(self.instance) + return attrs + + class TestListSerializer(serializers.ListSerializer): + child = TestSerializer() + + serializer = TestListSerializer( + instance=[BasicObject(id=1, uuid='uuid-1')], + data=[{'id': 1}], + ) + assert serializer.is_valid() + assert seen_instances == [None] + + def test_matching_with_configurable_lookup_field(self): + seen_instances = [] + + class TestSerializer(serializers.Serializer): + id = serializers.IntegerField(required=False) + uuid = serializers.CharField() + + class Meta: + lookup_field = 'uuid' + + def validate(self, attrs): + seen_instances.append(self.instance) + return attrs + + obj1 = BasicObject(id=1, uuid='uuid-1') + obj2 = BasicObject(id=2, uuid='uuid-2') + input_data = [{'id': 1, 'uuid': 'uuid-2'}] + + serializer = TestSerializer([obj1, obj2], data=input_data, many=True) + assert serializer.is_valid() + assert seen_instances == [obj2] + + def test_existing_instance_map_is_restored_after_validation(self): + class TestSerializer(serializers.Serializer): + pk = serializers.IntegerField() + + instance = [BasicObject(pk=1)] + original_instance_map = {'sentinel': BasicObject(pk=2)} + serializer = TestSerializer(instance, data=[{'pk': 1}], many=True) + serializer._list_serializer_instance_map = original_instance_map + + assert serializer.is_valid() + assert serializer._list_serializer_instance_map is original_instance_map + + class TestNestedListSerializer: """ Tests for using a ListSerializer as a field. @@ -889,6 +1150,35 @@ def test(self): assert serializer.data +def test_many_true_instance_level_validation_uses_matched_instance(): + class Obj: + def __init__(self, id, valid): + self.id = id + self.valid = valid + + class TestSerializer(serializers.Serializer): + id = serializers.IntegerField() + status = serializers.CharField() + + def validate_status(self, value): + if self.instance is None: + raise serializers.ValidationError("Instance not matched") + if not self.instance.valid: + raise serializers.ValidationError("Invalid instance") + return value + + objs = [Obj(1, True), Obj(2, False)] + serializer = TestSerializer( + instance=objs, + data=[{"id": 1, "status": "ok"}, {"id": 2, "status": "fail"}], + many=True, + partial=True, + ) + + assert not serializer.is_valid() + assert serializer.errors == {1: {'status': ['Invalid instance']}} + + class TestListSerializerErrorBehavior: """ Tests both ListSerializer error formats and consistency with ListField. diff --git a/tests/test_validators.py b/tests/test_validators.py index d75b9a139d..9f9fa7db25 100644 --- a/tests/test_validators.py +++ b/tests/test_validators.py @@ -147,30 +147,22 @@ def test_many_create_validates_uniqueness(self): 0: {'username': ['uniqueness model with this username already exists.']}, } - def test_many_update_requires_child_instance(self): + def test_many_update_matches_child_instance(self): serializer = UniquenessSerializer( instance=UniquenessModel.objects.all(), - data=[{'username': 'existing'}], + data=[{'id': self.instance.pk, 'username': 'existing'}], many=True, ) - message = ( - '`UniqueValidator` cannot determine the current instance during ' - 'a multiple update. Override ' - '`ListSerializer.run_child_validation()` to set `child.instance` ' - 'before validation.' - ) - with pytest.raises(RuntimeError, match=re.escape(message)): - serializer.is_valid() + assert serializer.is_valid() - def test_many_update_with_list_instance_requires_child_instance(self): + def test_many_update_with_list_instance_matches_child_instance(self): instances = [self.instance] serializer = UniquenessSerializer( instance=instances, - data=[{'username': 'existing'}], + data=[{'id': self.instance.pk, 'username': 'existing'}], many=True, ) - with pytest.raises(RuntimeError, match='`UniqueValidator` cannot determine'): - serializer.is_valid() + assert serializer.is_valid() def test_many_update_with_child_instance(self): """ @@ -328,7 +320,7 @@ def test_updated_instance_excluded_from_unique_together(self): 'position': 1 } - def test_many_update_requires_child_instance(self): + def test_many_update_matches_child_instance(self): class ListUpdateSerializer(serializers.ListSerializer): def update(self, instance, validated_data): return instance @@ -348,15 +340,7 @@ class Meta(UniquenessTogetherSerializer.Meta): }], many=True, ) - message = ( - '`UniqueTogetherValidator` cannot determine the current instance ' - 'during a multiple update. Override ' - '`ListSerializer.run_child_validation()` to set `child.instance` ' - 'before validation.' - ) - - with pytest.raises(RuntimeError, match=re.escape(message)): - serializer.is_valid() + assert serializer.is_valid() def test_many_update_with_child_instance(self): """ @@ -1099,20 +1083,17 @@ def test_updated_instance_excluded_from_unique_for_date(self): 'published': datetime.date(2000, 1, 1) } - def test_many_update_requires_child_instance(self): + def test_many_update_matches_child_instance(self): serializer = UniqueForDateSerializer( instance=UniqueForDateModel.objects.all(), - data=[{'slug': 'existing', 'published': '2000-01-01'}], + data=[{ + 'id': self.instance.pk, + 'slug': 'existing', + 'published': '2000-01-01', + }], many=True, ) - message = ( - '`UniqueForDateValidator` cannot determine the current instance ' - 'during a multiple update. Override ' - '`ListSerializer.run_child_validation()` to set `child.instance` ' - 'before validation.' - ) - with pytest.raises(RuntimeError, match=re.escape(message)): - serializer.is_valid() + assert serializer.is_valid() def test_many_update_with_child_instance(self): class ListUpdateSerializer(serializers.ListSerializer):