Skip to content

Commit 2fa19ed

Browse files
committed
Refine ListSerializer instance matching: reduce scope, add fallback behavior, improve tests
1 parent dcb4ad1 commit 2fa19ed

2 files changed

Lines changed: 183 additions & 20 deletions

File tree

rest_framework/serializers.py

Lines changed: 34 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -667,28 +667,32 @@ def run_child_validation(self, data):
667667
if not hasattr(self.child, 'instance'):
668668
return self.child.run_validation(data)
669669

670-
original_instance = self.child.instance
671-
had_initial_data = hasattr(self.child, 'initial_data')
672-
original_initial_data = getattr(self.child, 'initial_data', None)
670+
if not (
671+
hasattr(self, '_list_serializer_instance_map') and
672+
isinstance(data, Mapping)
673+
):
674+
return self.child.run_validation(data)
675+
676+
lookup_field = getattr(getattr(self.child, 'Meta', None), 'lookup_field', None)
677+
data_pk = data.get(lookup_field)
678+
if data_pk is None:
679+
data_pk = data.get('id')
680+
if data_pk is None:
681+
data_pk = data.get('pk')
682+
683+
if data_pk is None:
684+
return self.child.run_validation(data)
685+
686+
child_instance = self._list_serializer_instance_map.get(str(data_pk))
687+
if child_instance is None:
688+
return self.child.run_validation(data)
673689

690+
original_instance = self.child.instance
674691
try:
675-
if (
676-
hasattr(self, '_list_serializer_instance_map') and
677-
isinstance(data, Mapping)
678-
):
679-
data_pk = data.get('id') or data.get('pk')
680-
self.child.instance = (self._list_serializer_instance_map.get(str(data_pk))
681-
if data_pk is not None else None)
682-
683-
self.child.initial_data = data
692+
self.child.instance = child_instance
684693
return self.child.run_validation(data)
685694
finally:
686695
self.child.instance = original_instance
687-
if had_initial_data:
688-
self.child.initial_data = original_initial_data
689-
else:
690-
if hasattr(self.child, 'initial_data'):
691-
del self.child.initial_data
692696

693697
def to_internal_value(self, data):
694698
"""
@@ -731,20 +735,30 @@ def to_internal_value(self, data):
731735
errors = []
732736

733737
# Build a primary key lookup for instance matching in many=True updates.
734-
instance_map = {}
738+
instance_map = None
735739
if self.instance is not None:
736740
if isinstance(self.instance, Mapping):
737741
instance_map = {str(k): v for k, v in self.instance.items()}
738742
elif isinstance(self.instance, (list, tuple, models.query.QuerySet)):
743+
instance_map = {}
744+
lookup_field = getattr(getattr(self.child, 'Meta', None), 'lookup_field', None)
745+
739746
for obj in self.instance:
740-
pk = getattr(obj, 'pk', getattr(obj, 'id', None))
747+
if lookup_field is not None:
748+
pk = getattr(obj, lookup_field, None)
749+
else:
750+
pk = getattr(obj, 'pk', None)
751+
if pk is None:
752+
pk = getattr(obj, 'id', None)
753+
741754
if pk is not None:
742755
key = str(pk)
743756
# If duplicate keys are present, keep the last value,
744757
# matching standard mapping assignment behavior.
745758
instance_map[key] = obj
746759

747-
self._list_serializer_instance_map = instance_map
760+
if instance_map is not None:
761+
self._list_serializer_instance_map = instance_map
748762

749763
try:
750764
for item in data:

tests/test_serializer_lists.py

Lines changed: 149 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -203,6 +203,155 @@ def update(self, instance, validated_data):
203203
assert updated_instances == expected_output
204204

205205

206+
class TestListSerializerInstanceMatching:
207+
def test_matching_with_id(self):
208+
seen_instances = []
209+
210+
class TestSerializer(serializers.Serializer):
211+
id = serializers.IntegerField()
212+
213+
def validate(self, attrs):
214+
seen_instances.append(self.instance)
215+
return attrs
216+
217+
instance = [
218+
BasicObject(id=1),
219+
BasicObject(id=2),
220+
]
221+
input_data = [
222+
{'id': 1},
223+
{'id': 2},
224+
]
225+
226+
serializer = TestSerializer(instance, data=input_data, many=True)
227+
assert serializer.is_valid()
228+
assert seen_instances == instance
229+
230+
def test_matching_with_pk(self):
231+
seen_instances = []
232+
233+
class TestSerializer(serializers.Serializer):
234+
pk = serializers.IntegerField()
235+
236+
def validate(self, attrs):
237+
seen_instances.append(self.instance)
238+
return attrs
239+
240+
instance = [
241+
BasicObject(pk=1),
242+
BasicObject(pk=2),
243+
]
244+
input_data = [
245+
{'pk': 1},
246+
{'pk': 2},
247+
]
248+
249+
serializer = TestSerializer(instance, data=input_data, many=True)
250+
assert serializer.is_valid()
251+
assert seen_instances == instance
252+
253+
def test_matching_with_id_against_object_with_pk_only(self):
254+
seen_instances = []
255+
256+
class TestSerializer(serializers.Serializer):
257+
id = serializers.IntegerField()
258+
259+
def validate(self, attrs):
260+
seen_instances.append(self.instance)
261+
return attrs
262+
263+
instance = [BasicObject(pk=1)]
264+
input_data = [{'id': 1}]
265+
266+
serializer = TestSerializer(instance, data=input_data, many=True)
267+
assert serializer.is_valid()
268+
assert seen_instances == instance
269+
270+
def test_mapping_instance_matching(self):
271+
seen_instances = []
272+
273+
class TestSerializer(serializers.Serializer):
274+
id = serializers.IntegerField()
275+
276+
def validate(self, attrs):
277+
seen_instances.append(self.instance)
278+
return attrs
279+
280+
obj1 = BasicObject(id=1)
281+
obj2 = BasicObject(id=2)
282+
instance = {
283+
'1': obj1,
284+
'2': obj2,
285+
}
286+
input_data = [
287+
{'id': 1},
288+
{'id': 2},
289+
]
290+
291+
serializer = TestSerializer(instance, data=input_data, many=True)
292+
assert serializer.is_valid()
293+
assert seen_instances == [obj1, obj2]
294+
295+
def test_unsupported_instance_type_preserves_original_behavior(self):
296+
seen_instances = []
297+
298+
class TestSerializer(serializers.Serializer):
299+
id = serializers.IntegerField()
300+
301+
def validate(self, attrs):
302+
seen_instances.append(self.instance)
303+
return attrs
304+
305+
serializer = TestSerializer(instance=123, data=[{'id': 1}], many=True)
306+
assert serializer.is_valid()
307+
assert seen_instances == [123]
308+
309+
def test_missing_lookup_field_in_data_does_not_assign_instance(self):
310+
seen_instances = []
311+
312+
class TestSerializer(serializers.Serializer):
313+
id = serializers.IntegerField(required=False)
314+
315+
class Meta:
316+
lookup_field = 'uuid'
317+
318+
def validate(self, attrs):
319+
seen_instances.append(self.instance)
320+
return attrs
321+
322+
class TestListSerializer(serializers.ListSerializer):
323+
child = TestSerializer()
324+
325+
serializer = TestListSerializer(
326+
instance=[BasicObject(id=1, uuid='uuid-1')],
327+
data=[{'id': 1}],
328+
)
329+
assert serializer.is_valid()
330+
assert seen_instances == [None]
331+
332+
def test_matching_with_configurable_lookup_field(self):
333+
seen_instances = []
334+
335+
class TestSerializer(serializers.Serializer):
336+
id = serializers.IntegerField(required=False)
337+
uuid = serializers.CharField()
338+
339+
class Meta:
340+
lookup_field = 'uuid'
341+
342+
def validate(self, attrs):
343+
seen_instances.append(self.instance)
344+
return attrs
345+
346+
obj1 = BasicObject(id=1, uuid='uuid-1')
347+
obj2 = BasicObject(id=2, uuid='uuid-2')
348+
input_data = [{'id': 1, 'uuid': 'uuid-2'}]
349+
350+
serializer = TestSerializer([obj1, obj2], data=input_data, many=True)
351+
assert serializer.is_valid()
352+
assert seen_instances == [obj2]
353+
354+
206355
class TestNestedListSerializer:
207356
"""
208357
Tests for using a ListSerializer as a field.

0 commit comments

Comments
 (0)