Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 89 additions & 0 deletions rest_framework/relations.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down Expand Up @@ -246,6 +247,16 @@ def __init__(self, **kwargs):
self.pk_field = kwargs.pop('pk_field', None)
super().__init__(**kwargs)

@classmethod
def many_init(cls, *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

Expand Down Expand Up @@ -583,3 +594,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]
3 changes: 2 additions & 1 deletion rest_framework/serializers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
146 changes: 146 additions & 0 deletions tests/test_relations_pk.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
from unittest.mock import patch

import pytest
from django.test import TestCase

Expand Down Expand Up @@ -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):
Expand Down