diff --git a/rest_framework/relations.py b/rest_framework/relations.py index 4409bce77c..91ed97a062 100644 --- a/rest_framework/relations.py +++ b/rest_framework/relations.py @@ -10,6 +10,7 @@ from django.utils.encoding import smart_str, uri_to_iri from django.utils.translation import gettext_lazy as _ +from rest_framework.exceptions import ValidationError from rest_framework.fields import ( Field, SkipField, empty, get_attribute, is_simple_callable, iter_options ) @@ -246,6 +247,18 @@ def __init__(self, **kwargs): self.pk_field = kwargs.pop('pk_field', None) super().__init__(**kwargs) + @classmethod + def many_init(cls, *args, **kwargs): + if cls is not PrimaryKeyRelatedField: + return super().many_init(*args, **kwargs) + # Use PrimaryKeyManyRelatedField so many=True validates with one + # in_bulk() query. Slug/Hyperlinked keep RelatedField.many_init. + list_kwargs = {'child_relation': cls(*args, **kwargs)} + for key in kwargs: + if key in MANY_RELATION_KWARGS: + list_kwargs[key] = kwargs[key] + return PrimaryKeyManyRelatedField(**list_kwargs) + def use_pk_only_optimization(self): return True @@ -583,3 +596,81 @@ def iter_options(self): cutoff=self.html_cutoff, cutoff_text=self.html_cutoff_text ) + + +class PrimaryKeyManyRelatedField(ManyRelatedField): + """ + Many-related field for PrimaryKeyRelatedField that resolves every pk with + a single `in_bulk()` query instead of one `get()` per item. + + Treated as private API — constructed via PrimaryKeyRelatedField.many_init. + """ + + def to_internal_value(self, data): + if isinstance(data, str) or not hasattr(data, '__iter__'): + self.fail('not_a_list', input_type=type(data).__name__) + if not self.allow_empty and len(data) == 0: + self.fail('empty') + + # Resolve every pk with a single query instead of one `get()` per item. + # Collect per-item errors (incorrect_type / does_not_exist / pk_field) + # keyed by index, matching ListField.run_child_validation. Input + # ordering and duplicates are preserved. + child = self.child_relation + queryset = child.get_queryset() + model_pk = queryset.model._meta.pk + # Each entry is (idx, lookup_key, value): `value` mirrors the per-item + # path (post-`pk_field`) and is used for error details, while + # `lookup_key` is the pk-typed value used to match `in_bulk()` results. + errors = {} + entries = [] + for idx, item in enumerate(data): + try: + value = item + if child.pk_field is not None: + value = child.pk_field.to_internal_value(value) + except ValidationError as exc: + errors[idx] = exc.detail + continue + try: + if isinstance(value, bool): + raise TypeError + # Coerce to the pk's Python type (e.g. "1" -> 1) so the lookup + # below matches the keys returned by `in_bulk()`, exactly as + # `queryset.get(pk=value)` would have. + lookup_key = model_pk.get_prep_value(value) + except (TypeError, ValueError): + try: + child.fail( + 'incorrect_type', data_type=type(value).__name__ + ) + except ValidationError as exc: + errors[idx] = exc.detail + continue + entries.append((idx, lookup_key, value)) + lookup_keys = [lookup_key for _, lookup_key, _ in entries] + try: + objects = queryset.in_bulk(lookup_keys) if lookup_keys else {} + except (TypeError, ValueError): + # queryset doesn't support in_bulk (e.g. distinct/sliced); fall + # back to a collecting per-item loop so mixed lists still report + # every invalid item. + errors = {} + result = [] + for idx, item in enumerate(data): + try: + result.append(child.to_internal_value(item)) + except ValidationError as exc: + errors[idx] = exc.detail + if errors: + raise ValidationError(errors) + return result + for idx, lookup_key, value in entries: + if lookup_key not in objects: + try: + child.fail('does_not_exist', pk_value=value) + except ValidationError as exc: + errors[idx] = exc.detail + if errors: + raise ValidationError(errors) + return [objects[lookup_key] for _, lookup_key, _ in entries] diff --git a/rest_framework/serializers.py b/rest_framework/serializers.py index fc8e83c768..789fb30710 100644 --- a/rest_framework/serializers.py +++ b/rest_framework/serializers.py @@ -61,7 +61,8 @@ ) from rest_framework.relations import ( # NOQA # isort:skip HyperlinkedIdentityField, HyperlinkedRelatedField, ManyRelatedField, - PrimaryKeyRelatedField, RelatedField, SlugRelatedField, StringRelatedField, + PrimaryKeyManyRelatedField, PrimaryKeyRelatedField, RelatedField, + SlugRelatedField, StringRelatedField, ) # Non-field imports, but public API diff --git a/tests/test_relations_pk.py b/tests/test_relations_pk.py index 0769defebd..222656dff6 100644 --- a/tests/test_relations_pk.py +++ b/tests/test_relations_pk.py @@ -1,3 +1,5 @@ +from unittest.mock import patch + import pytest from django.test import TestCase @@ -227,6 +229,150 @@ def test_data_cannot_be_accessed_prior_to_is_valid(self): serializer.data +class PKManyRelatedFieldBulkValidationTests(TestCase): + """`PrimaryKeyRelatedField(many=True)` should resolve all pks in a single + query rather than one query per item (regression test for #9607).""" + + def setUp(self): + self.pks = [ + ManyToManyTarget.objects.create(name='target-%d' % idx).pk + for idx in range(1, 6) + ] + + def _field(self, queryset=None): + if queryset is None: + queryset = ManyToManyTarget.objects.all() + field = serializers.PrimaryKeyRelatedField(queryset=queryset, many=True) + field.bind('targets', serializers.Serializer()) + return field + + def test_validation_uses_single_query(self): + field = self._field() + with self.assertNumQueries(1): + field.run_validation(self.pks) + + def test_order_and_duplicates_preserved(self): + field = self._field() + order = [self.pks[2], self.pks[0], self.pks[0], self.pks[1]] + result = field.run_validation(order) + assert [obj.pk for obj in result] == order + + def test_string_pks_are_accepted(self): + # HTML form input arrives as strings; must match int pks (#9607). + field = self._field() + result = field.run_validation([str(pk) for pk in self.pks]) + assert [obj.pk for obj in result] == self.pks + + def test_does_not_exist_error(self): + field = self._field() + missing = max(self.pks) + 1000 + with pytest.raises(serializers.ValidationError) as exc_info: + field.run_validation([self.pks[0], missing]) + detail = exc_info.value.detail + assert 0 not in detail + assert detail[1][0].code == 'does_not_exist' + + def test_incorrect_type_error(self): + field = self._field() + with pytest.raises(serializers.ValidationError) as exc_info: + field.run_validation(['not-a-pk']) + assert exc_info.value.detail[0][0].code == 'incorrect_type' + + def test_queryset_filtering_is_respected(self): + field = self._field(ManyToManyTarget.objects.exclude(pk=self.pks[1])) + with pytest.raises(serializers.ValidationError) as exc_info: + field.run_validation([self.pks[0], self.pks[1]]) + detail = exc_info.value.detail + assert 0 not in detail + assert detail[1][0].code == 'does_not_exist' + + def test_pk_field_transform_is_applied(self): + field = serializers.PrimaryKeyRelatedField( + queryset=ManyToManyTarget.objects.all(), many=True, + pk_field=serializers.IntegerField()) + field.bind('targets', serializers.Serializer()) + result = field.run_validation([str(self.pks[0]), str(self.pks[1])]) + assert [obj.pk for obj in result] == [self.pks[0], self.pks[1]] + + def test_error_details_match_per_item_with_pk_field(self): + # The bulk path must report the same incorrect_type detail as the + # per-item path, i.e. the type *after* pk_field transformation. + child = serializers.PrimaryKeyRelatedField( + queryset=ManyToManyTarget.objects.all(), + pk_field=serializers.BooleanField()) + child.bind('targets', serializers.Serializer()) + field = serializers.PrimaryKeyRelatedField( + queryset=ManyToManyTarget.objects.all(), many=True, + pk_field=serializers.BooleanField()) + field.bind('targets', serializers.Serializer()) + with pytest.raises(serializers.ValidationError) as per_item: + child.to_internal_value('true') + with pytest.raises(serializers.ValidationError) as bulk: + field.to_internal_value(['true']) + assert bulk.value.detail[0] == per_item.value.detail + assert 'bool' in str(bulk.value.detail[0]) + + def test_many_related_field_with_non_related_child(self): + # Plain ManyRelatedField (not the PK many subclass) still validates + # a non-related child with the per-item loop. + field = serializers.ManyRelatedField( + child_relation=serializers.IntegerField()) + field.bind('values', serializers.Serializer()) + assert field.to_internal_value([1, 2, 3]) == [1, 2, 3] + + def test_many_true_uses_primary_key_many_related_field(self): + field = serializers.PrimaryKeyRelatedField( + queryset=ManyToManyTarget.objects.all(), many=True) + assert isinstance(field, serializers.PrimaryKeyManyRelatedField) + + def test_collects_mixed_errors_in_one_query(self): + field = self._field() + missing = max(self.pks) + 1000 + with self.assertNumQueries(1): + with pytest.raises(serializers.ValidationError) as exc_info: + field.run_validation([missing, 'not-a-pk', self.pks[0]]) + detail = exc_info.value.detail + assert detail[0][0].code == 'does_not_exist' + assert detail[1][0].code == 'incorrect_type' + assert 2 not in detail + + def test_duplicate_invalid_pks_report_each_index(self): + field = self._field() + missing = max(self.pks) + 1000 + with pytest.raises(serializers.ValidationError) as exc_info: + field.run_validation([missing, self.pks[0], missing]) + detail = exc_info.value.detail + assert detail[0][0].code == 'does_not_exist' + assert 1 not in detail + assert detail[2][0].code == 'does_not_exist' + + def test_in_bulk_fallback_collects_errors(self): + # in_bulk() raises TypeError on sliced querysets; inject that so + # the fallback's per-item get() can still resolve valid pks. + field = self._field() + missing = max(self.pks) + 1000 + with patch('django.db.models.query.QuerySet.in_bulk', + side_effect=TypeError): + with pytest.raises(serializers.ValidationError) as exc_info: + field.run_validation([self.pks[0], missing, 'not-a-pk']) + detail = exc_info.value.detail + assert 0 not in detail + assert detail[1][0].code == 'does_not_exist' + assert detail[2][0].code == 'incorrect_type' + + def test_pk_field_validation_error_is_collected(self): + field = serializers.PrimaryKeyRelatedField( + queryset=ManyToManyTarget.objects.all(), many=True, + pk_field=serializers.IntegerField()) + field.bind('targets', serializers.Serializer()) + with self.assertNumQueries(1): + with pytest.raises(serializers.ValidationError) as exc_info: + field.run_validation([self.pks[0], 'not-a-number']) + detail = exc_info.value.detail + assert 0 not in detail + assert 1 in detail + + @pytest.mark.usefixtures("reset_sequences") class PKForeignKeyTests(TestCase): def setUp(self):