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
160 changes: 160 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,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

Expand Down Expand Up @@ -583,3 +596,150 @@ 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))

objects = self._resolve_objects(queryset, entries)
if objects is None:
return self._collecting_per_item(child, data)

resolved = {}
unmatched = []
for idx, lookup_key, value in entries:
if lookup_key in objects:
resolved[idx] = objects[lookup_key]
else:
unmatched.append((idx, lookup_key, value))

if unmatched:
if queryset.query.is_sliced:
# Allowed set is the slice; do not call get() on a sliced QS.
for idx, lookup_key, value in unmatched:
try:
child.fail('does_not_exist', pk_value=value)
except ValidationError as exc:
errors[idx] = exc.detail
else:
self._recover_unmatched(
queryset, child, unmatched, resolved, errors
)

if errors:
raise ValidationError(errors)
return [resolved[idx] for idx, _, _ in entries]

def _resolve_objects(self, queryset, entries):
"""
Return {pk: obj} for `entries`, or None to signal collecting per-item.

Sliced querysets cannot use in_bulk/get; materialize the slice once.
Other in_bulk TypeErrors (e.g. values()/values_list()) fall back to
collecting per-item to_internal_value.
"""
lookup_keys = [lookup_key for _, lookup_key, _ in entries]
try:
return queryset.in_bulk(lookup_keys) if lookup_keys else {}
except (TypeError, ValueError):
if queryset.query.is_sliced:
try:
return {obj.pk: obj for obj in queryset}
except (TypeError, AttributeError):
return None
return None

def _collecting_per_item(self, child, data):
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

def _recover_unmatched(self, queryset, child, unmatched, resolved, errors):
"""
Recover keys missed by Python equality against in_bulk() results.

One filter(pk__in=...) probe first: empty means true misses. Non-empty
means possible CI-collation / prep mismatch — resolve each unmatched
value with get(pk=value), collecting DoesNotExist like the per-item
path.
"""
probe = list(queryset.filter(
pk__in=[value for _, _, value in unmatched]
))
if not probe:
for idx, lookup_key, value in unmatched:
try:
child.fail('does_not_exist', pk_value=value)
except ValidationError as exc:
errors[idx] = exc.detail
return
for idx, lookup_key, value in unmatched:
try:
resolved[idx] = queryset.get(pk=value)
except ObjectDoesNotExist:
try:
child.fail('does_not_exist', pk_value=value)
except ValidationError as exc:
errors[idx] = exc.detail
except (TypeError, ValueError):
try:
child.fail(
'incorrect_type', data_type=type(value).__name__
)
except ValidationError as exc:
errors[idx] = exc.detail
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
189 changes: 189 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,193 @@ 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_subclass_many_uses_per_item_to_internal_value(self):
# Subclasses keep RelatedField.many_init so overridden
# to_internal_value is still called for many=True (auvipy 79ea3de0).
calls = []

class TenantPKField(serializers.PrimaryKeyRelatedField):
def to_internal_value(self, data):
calls.append(data)
return super().to_internal_value(data)

field = TenantPKField(
queryset=ManyToManyTarget.objects.all(), many=True)
field.bind('targets', serializers.Serializer())
assert isinstance(field, serializers.ManyRelatedField)
assert not isinstance(field, serializers.PrimaryKeyManyRelatedField)
result = field.run_validation([self.pks[0], self.pks[1]])
assert calls == [self.pks[0], self.pks[1]]
assert [obj.pk for obj in result] == [self.pks[0], self.pks[1]]

def test_collects_mixed_errors_with_probe_query(self):
# in_bulk + one filter(pk__in=...) probe for true misses (2 queries).
field = self._field()
missing = max(self.pks) + 1000
with self.assertNumQueries(2):
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_sliced_queryset_materializes_once(self):
# Real sliced queryset: in_bulk raises; materialize the slice once.
allowed = list(ManyToManyTarget.objects.filter(
pk__in=self.pks).order_by('pk')[:3])
allowed_pks = [obj.pk for obj in allowed]
outside = [pk for pk in self.pks if pk not in allowed_pks][0]
field = self._field(
ManyToManyTarget.objects.filter(pk__in=self.pks).order_by('pk')[:3]
)
with self.assertNumQueries(1):
result = field.run_validation(allowed_pks)
assert [obj.pk for obj in result] == allowed_pks
with pytest.raises(serializers.ValidationError) as exc_info:
field.run_validation([allowed_pks[0], outside, '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_in_bulk_key_miss_recovers_via_probe(self):
# Simulate Python key miss after in_bulk (CI collation / prep): empty
# map forces the filter probe + get recovery path.
field = self._field()
missing = max(self.pks) + 1000
with patch.object(
type(ManyToManyTarget.objects.all()),
'in_bulk',
return_value={},
):
result = field.run_validation([self.pks[0]])
assert [obj.pk for obj in result] == [self.pks[0]]
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_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