diff --git a/docs/api-guide/validators.md b/docs/api-guide/validators.md index f598a7e062..106d6d4f06 100644 --- a/docs/api-guide/validators.md +++ b/docs/api-guide/validators.md @@ -222,6 +222,28 @@ For example: extra_kwargs = {'client': {'required': False}} validators = [] # Remove a default "unique together" constraint. +### UniqueConstraint with conditions + +When using Django's `UniqueConstraint` with conditions that reference other model fields, DRF will automatically use +`UniqueTogetherValidator` instead of field-level `UniqueValidator`. This ensures proper validation behavior when the constraint +effectively involves multiple fields. + +For example, a single-field constraint with a condition becomes a multi-field validation when the condition references other fields. + + class MyModel(models.Model): + name = models.CharField(max_length=100) + status = models.CharField(max_length=20) + + class Meta: + constraints = [ + models.UniqueConstraint( + fields=['name'], + condition=models.Q(status='active'), + name='unique_active_name' + ) + ] + + ### Updating nested serializers When applying an update to an existing instance, uniqueness validators will diff --git a/rest_framework/compat.py b/rest_framework/compat.py index 67f3a72864..25947e2d71 100644 --- a/rest_framework/compat.py +++ b/rest_framework/compat.py @@ -156,3 +156,30 @@ def split_header_value(value, sep=","): SHORT_SEPARATORS = (',', ':') LONG_SEPARATORS = (', ', ': ') INDENT_SEPARATORS = (',', ': ') + + +def get_referenced_base_fields_from_q(q_object): + """ + Return the base field names referenced by a Q object. + This is a compatibility helper for Django versions that may not have + `referenced_base_fields` attribute on Q objects. + """ + if q_object is None: + return set() + + # Prefer Django's built-in implementation when available. + referenced = getattr(q_object, "referenced_base_fields", None) + if referenced is not None: + return set(referenced) + + referenced_fields = set() + for child in q_object.children: + if isinstance(child, tuple): + # child[0] is the field name (e.g., 'status', 'global_id__lte') + # We strip off any lookup part (__lte, __exact, etc.) + field_name = child[0].split('__', 1)[0] + referenced_fields.add(field_name) + else: + # child is another Q object + referenced_fields.update(get_referenced_base_fields_from_q(child)) + return referenced_fields diff --git a/rest_framework/serializers.py b/rest_framework/serializers.py index fc8e83c768..038a6f8c8f 100644 --- a/rest_framework/serializers.py +++ b/rest_framework/serializers.py @@ -27,7 +27,9 @@ from django.utils.functional import cached_property from django.utils.translation import gettext_lazy as _ -from rest_framework.compat import postgres_fields +from rest_framework.compat import ( + get_referenced_base_fields_from_q, postgres_fields +) from rest_framework.deprecation import RemovedInDRF320Warning from rest_framework.exceptions import ErrorDetail, ValidationError from rest_framework.fields import get_error_detail @@ -1469,20 +1471,28 @@ def get_unique_together_constraints(self, model): """ for parent_class in [model] + list(model._meta.parents): for unique_together in parent_class._meta.unique_together: - yield unique_together, model._default_manager, [], None, None + yield unique_together, model._default_manager, [], None, None, None for constraint in parent_class._meta.constraints: - if isinstance(constraint, models.UniqueConstraint) and len(constraint.fields) > 1: + if isinstance(constraint, models.UniqueConstraint): if constraint.condition is None: condition_fields = [] else: - condition_fields = list(constraint.condition.referenced_base_fields) - yield ( - constraint.fields, - model._default_manager, - condition_fields, - constraint.condition, - constraint.nulls_distinct, - ) + condition_fields = list( + get_referenced_base_fields_from_q(constraint.condition) + ) + + # Combine constraint fields and condition fields. If the union + # involves multiple fields, treat as unique-together validation + required_fields = {*constraint.fields, *condition_fields} + if constraint.fields and len(required_fields) > 1: + yield ( + constraint.fields, + model._default_manager, + condition_fields, + constraint.condition, + constraint.nulls_distinct, + constraint, + ) def get_uniqueness_extra_kwargs(self, field_names, declared_fields, extra_kwargs): """ @@ -1515,7 +1525,8 @@ def get_uniqueness_extra_kwargs(self, field_names, declared_fields, extra_kwargs # Include each of the `unique_together` and `UniqueConstraint` field names, # so long as all the field names are included on the serializer. - for unique_together_list, queryset, condition_fields, condition, nulls_distinct in self.get_unique_together_constraints(model): + for unique_together_list, queryset, condition_fields, condition, nulls_distinct, unused_constraint in self.get_unique_together_constraints( + model): unique_together_list_and_condition_fields = set(unique_together_list) | set(condition_fields) if model_fields_names.issuperset(unique_together_list_and_condition_fields): unique_constraint_names |= unique_together_list_and_condition_fields @@ -1624,7 +1635,12 @@ def _get_constraint_violation_error_message(self, constraint): def get_unique_together_validators(self): """ - Determine a default set of validators for any unique_together constraints. + Determine a default set of validators for any unique_together constraints + and UniqueConstraint objects. + + This method now preserves the original constraint object in the yielded + data from get_unique_together_constraints() to ensure custom violation + messages and error codes are correctly propagated to the validators. """ # The field names we're passing though here only include fields # which may map onto a model field. Any dotted field name lookups @@ -1648,17 +1664,11 @@ def get_unique_together_validators(self): for name, source in field_sources.items(): source_map[source].append(name) - unique_constraint_by_fields = { - constraint.fields: constraint - for model_cls in (*self.Meta.model._meta.parents, self.Meta.model) - for constraint in model_cls._meta.constraints - if isinstance(constraint, models.UniqueConstraint) - } - # Note that we make sure to check `unique_together` both on the # base model class, but also on any parent classes. validators = [] - for unique_together, queryset, condition_fields, condition, nulls_distinct in self.get_unique_together_constraints(self.Meta.model): + for unique_together, queryset, condition_fields, condition, nulls_distinct, constraint in self.get_unique_together_constraints( + self.Meta.model): # Skip if serializer does not map to all unique together sources unique_together_and_condition_fields = set(unique_together) | set(condition_fields) if not set(source_map).issuperset(unique_together_and_condition_fields): @@ -1682,8 +1692,9 @@ def get_unique_together_validators(self): field_names = tuple(source_map[f][0] for f in unique_together) - constraint = unique_constraint_by_fields.get(tuple(unique_together)) + # Extract custom violation message and code from the constraint if available violation_error_message = self._get_constraint_violation_error_message(constraint) if constraint else None + violation_error_code = getattr(constraint, 'violation_error_code', None) validator = UniqueTogetherValidator( queryset=queryset, @@ -1691,7 +1702,7 @@ def get_unique_together_validators(self): condition_fields=tuple(source_map[f][0] for f in condition_fields), condition=condition, message=violation_error_message, - code=getattr(constraint, 'violation_error_code', None), + code=violation_error_code, nulls_distinct=nulls_distinct, ) validators.append(validator) diff --git a/rest_framework/utils/field_mapping.py b/rest_framework/utils/field_mapping.py index fd456a08c9..64e7f8f248 100644 --- a/rest_framework/utils/field_mapping.py +++ b/rest_framework/utils/field_mapping.py @@ -8,7 +8,9 @@ from django.db import models from django.utils.text import capfirst -from rest_framework.compat import postgres_fields +from rest_framework.compat import ( + get_referenced_base_fields_from_q, postgres_fields +) from rest_framework.validators import UniqueValidator NUMERIC_FIELD_TYPES = ( @@ -79,10 +81,18 @@ def get_unique_validators(field_name, model_field): unique_error_message = get_unique_error_message(model_field) queryset = model_field.model._default_manager for condition in conditions: - yield UniqueValidator( - queryset=queryset if condition is None else queryset.filter(condition), - message=unique_error_message + condition_fields = ( + get_referenced_base_fields_from_q(condition) + if condition is not None + else set() ) + # Only use UniqueValidator if the union of field and condition fields is 1 + # (i.e. no additional fields referenced in conditions) + if len(field_set | condition_fields) == 1: + yield UniqueValidator( + queryset=queryset if condition is None else queryset.filter(condition), + message=unique_error_message, + ) def get_field_kwargs(field_name, model_field): diff --git a/rest_framework/validators.py b/rest_framework/validators.py index 44f2d3a6c7..d85da13220 100644 --- a/rest_framework/validators.py +++ b/rest_framework/validators.py @@ -173,9 +173,9 @@ def exclude_current_instance(self, attrs, queryset, instance): def __call__(self, attrs, serializer): if ( - serializer.instance is not None and - getattr(serializer.parent, 'many', False) and - not hasattr(serializer.instance, 'pk') + serializer.instance is not None and + getattr(serializer.parent, 'many', False) and + not hasattr(serializer.instance, 'pk') ): raise RuntimeError( '`UniqueTogetherValidator` cannot determine the current ' @@ -189,27 +189,34 @@ def __call__(self, attrs, serializer): queryset = self.filter_queryset(attrs, queryset, serializer) queryset = self.exclude_current_instance(attrs, queryset, serializer.instance) - checked_names = [ - serializer.fields[field_name].source for field_name in self.fields - ] + # Combine constraint fields and condition fields to detect changes + # in either set of fields. This ensures that updates to condition + # fields also trigger revalidation. + checked_names = list({ + serializer.fields[field_name].source for field_name in self.fields + } | { + serializer.fields[field_name].source for field_name in self.condition_fields + }) + # Ignore validation if any field is None if serializer.instance is None: - checked_values = [attrs[field_name] for field_name in checked_names] + checked_values = [attrs.get(field_name) for field_name in checked_names] else: # Ignore validation if all field values are unchanged checked_values = [ - attrs[field_name] + attrs.get(field_name) for field_name in checked_names - if attrs[field_name] != getattr(serializer.instance, field_name) + if attrs.get(field_name) != getattr(serializer.instance, field_name, None) ] condition_sources = (serializer.fields[field_name].source for field_name in self.condition_fields) condition_kwargs = { - source: attrs[source] + source: attrs.get(source) if source in attrs - else getattr(serializer.instance, source) + else getattr(serializer.instance, source, None) for source in condition_sources } + if checked_values: # Skip validation for None values unless nulls_distinct is False if self.nulls_distinct is not False and None in checked_values: diff --git a/tests/test_validators.py b/tests/test_validators.py index 82181f746f..27960005db 100644 --- a/tests/test_validators.py +++ b/tests/test_validators.py @@ -170,6 +170,24 @@ class Meta: unique_together = ('race_name', 'position') +class ConditionUniquenessTogetherModel(models.Model): + """ + Used to ensure that unique constraints with single fields but at least one other + distinct condition field are included when checking unique_together constraints. + """ + race_name = models.CharField(max_length=100) + position = models.IntegerField() + + class Meta: + constraints = [ + models.UniqueConstraint( + name="condition_uniqueness_together_model_race_name", + fields=('race_name',), + condition=models.Q(position__lte=1) + ) + ] + + class UniquenessTogetherSerializer(serializers.ModelSerializer): class Meta: model = UniquenessTogetherModel @@ -182,6 +200,12 @@ class Meta: fields = '__all__' +class ConditionUniquenessTogetherSerializer(serializers.ModelSerializer): + class Meta: + model = ConditionUniquenessTogetherModel + fields = '__all__' + + class TestUniquenessTogetherValidation(TestCase): def setUp(self): self.instance = UniquenessTogetherModel.objects.create( @@ -222,6 +246,22 @@ def test_is_not_unique_together(self): ] } + def test_is_not_unique_together_condition_based(self): + """ + Failing unique together validation should result in non-field errors when a condition-based + unique together constraint is violated. + """ + ConditionUniquenessTogetherModel.objects.create(race_name='example', position=1) + + data = {'race_name': 'example', 'position': 1} + serializer = ConditionUniquenessTogetherSerializer(data=data) + assert not serializer.is_valid() + assert serializer.errors == { + 'non_field_errors': [ + 'The fields race_name must make a unique set.' + ] + } + def test_is_unique_together(self): """ In a unique together validation, one field may be non-unique @@ -235,6 +275,36 @@ def test_is_unique_together(self): 'position': 2 } + def test_is_unique_together_condition_based(self): + """ + In a condition-based unique together validation, data is valid when + the constrained field differs when the condition applies. + """ + ConditionUniquenessTogetherModel.objects.create(race_name='example', position=1) + + data = {'race_name': 'other', 'position': 1} + serializer = ConditionUniquenessTogetherSerializer(data=data) + assert serializer.is_valid() + assert serializer.validated_data == { + 'race_name': 'other', + 'position': 1 + } + + def test_is_unique_together_when_condition_does_not_apply(self): + """ + In a condition-based unique together validation, data is valid when + the condition does not apply, even if constrained fields match existing records. + """ + ConditionUniquenessTogetherModel.objects.create(race_name='example', position=1) + + data = {'race_name': 'example', 'position': 2} + serializer = ConditionUniquenessTogetherSerializer(data=data) + assert serializer.is_valid() + assert serializer.validated_data == { + 'race_name': 'example', + 'position': 2 + } + def test_updated_instance_excluded_from_unique_together(self): """ When performing an update, the existing instance does not count @@ -278,6 +348,21 @@ class Meta(UniquenessTogetherSerializer.Meta): with pytest.raises(RuntimeError, match=re.escape(message)): serializer.is_valid() + def test_updated_instance_excluded_from_unique_together_condition_based(self): + """ + When performing an update, the existing instance does not count + as a match against uniqueness. + """ + instance = ConditionUniquenessTogetherModel.objects.create(race_name='example', position=1) + + data = {'race_name': 'example', 'position': 0} + serializer = ConditionUniquenessTogetherSerializer(instance, data=data) + assert serializer.is_valid() + assert serializer.validated_data == { + 'race_name': 'example', + 'position': 0 + } + def test_unique_together_is_required(self): """ In a unique together validation, all fields are required. @@ -583,6 +668,39 @@ class Meta: """) assert repr(serializer) == expected + def test_condition_field_change_triggers_revalidation(self): + """ + When only a condition field changes, the validator should still + recheck uniqueness. Previously, the validator would skip validation + because it only checked changes in constraint fields. + """ + # Create an initial record that satisfies the condition + instance = ConditionUniquenessTogetherModel.objects.create( + race_name='example', + position=1 # This matches the condition (position__lte=1) + ) + # Create another record with a different race_name + ConditionUniquenessTogetherModel.objects.create( + race_name='other', + position=1 + ) + + # Now update the first record's position from 1 to 2 + # This changes the condition field (position) but not the constraint field (race_name) + data = {'race_name': 'example', 'position': 2} + serializer = ConditionUniquenessTogetherSerializer(instance, data=data) + assert serializer.is_valid() + assert serializer.validated_data == { + 'race_name': 'example', + 'position': 2 + } + + # Now create another record that would violate the constraint + # if the condition applied (but it shouldn't, because position=2 doesn't match) + new_data = {'race_name': 'example', 'position': 2} + new_serializer = ConditionUniquenessTogetherSerializer(data=new_data) + assert new_serializer.is_valid() + class UniqueConstraintModel(models.Model): race_name = models.CharField(max_length=100) @@ -786,22 +904,22 @@ class Meta: def test_single_field_uniq_validators(self): """ UniqueConstraint with single field must be transformed into - field's UniqueValidator + field's UniqueValidator if no distinct condition fields exist (else UniqueTogetherValidator) """ # Backends like PostgreSQL add Min/Max validators for IntegerField; # SQLite does not because it has no fixed integer range. has_int_range = connection.ops.integer_field_range('IntegerField')[0] is not None extra_validators_qty = 2 if has_int_range else 0 serializer = UniqueConstraintSerializer() - assert len(serializer.validators) == 2 + assert len(serializer.validators) == 4 validators = serializer.fields['global_id'].validators assert len(validators) == 1 + extra_validators_qty assert validators[0].queryset == UniqueConstraintModel.objects + ids_in_qs = {frozenset(v.queryset.values_list('id', flat=True)) for v in validators if hasattr(v, "queryset")} + assert ids_in_qs == {frozenset({1, 2, 3})} validators = serializer.fields['fancy_conditions'].validators - assert len(validators) == 2 + extra_validators_qty - ids_in_qs = {frozenset(v.queryset.values_list('id', flat=True)) for v in validators if hasattr(v, "queryset")} - assert ids_in_qs == {frozenset([1]), frozenset([3])} + assert len(validators) == extra_validators_qty def test_nullable_unique_constraint_fields_are_not_required(self): serializer = UniqueConstraintNullableSerializer(data={'title': 'Bob'})