v0.163.0
  1from __future__ import annotations
  2
  3from types import NoneType
  4from typing import TYPE_CHECKING, Any
  5
  6from plain.exceptions import ValidationError
  7from plain.postgres.constants import LOOKUP_SEP
  8from plain.postgres.ddl import (
  9    build_include_sql,
 10    compile_expression_sql,
 11    compile_index_expressions_sql,
 12)
 13from plain.postgres.dialect import quote_name
 14from plain.postgres.exceptions import FieldError
 15from plain.postgres.expressions import (
 16    Exists,
 17    F,
 18    OrderBy,
 19    ReplaceableExpression,
 20)
 21from plain.postgres.lookups import Exact
 22from plain.postgres.query_utils import Q
 23
 24if TYPE_CHECKING:
 25    from plain.postgres.base import Model
 26
 27__all__ = ["BaseConstraint", "CheckConstraint", "UniqueConstraint"]
 28
 29
 30ViolationError = str | dict[str, Any] | list[Any] | ValidationError
 31
 32
 33class BaseConstraint:
 34    violation_error: ViolationError | None = None
 35
 36    def __init__(
 37        self,
 38        *,
 39        name: str,
 40        violation_error: ViolationError | None = None,
 41    ) -> None:
 42        self.name = name
 43        self.violation_error = violation_error
 44
 45    @property
 46    def contains_expressions(self) -> bool:
 47        return False
 48
 49    def to_sql(self, model: type[Model]) -> str:
 50        raise NotImplementedError(
 51            "subclasses of BaseConstraint must provide a to_sql() method"
 52        )
 53
 54    def validate(
 55        self, model: type[Model], instance: Model, exclude: set[str] | None = None
 56    ) -> None:
 57        raise NotImplementedError(
 58            "subclasses of BaseConstraint must provide a validate() method"
 59        )
 60
 61    def _build_violation_error(self) -> ValidationError:
 62        if self.violation_error is None:
 63            return ValidationError(f'Constraint "{self.name}" is violated.')
 64        if isinstance(self.violation_error, ValidationError):
 65            return self.violation_error
 66        return ValidationError(self.violation_error)
 67
 68    def _db_violation_error(
 69        self, instance: Model, model: type[Model]
 70    ) -> ValidationError | None:
 71        """The ValidationError to raise when the database reports a violation
 72        of this constraint (mapped from an IntegrityError at the write
 73        boundary), or None if this constraint type can't be mapped — the
 74        caller then re-raises the original IntegrityError.
 75
 76        Subclasses that can describe their own violations override this.
 77        """
 78        return None
 79
 80    def deconstruct(self) -> tuple[str, tuple[Any, ...], dict[str, Any]]:
 81        path = f"{self.__class__.__module__}.{self.__class__.__name__}"
 82        path = path.replace("plain.postgres.constraints", "plain.postgres")
 83        kwargs: dict[str, Any] = {"name": self.name}
 84        if self.violation_error is not None:
 85            kwargs["violation_error"] = self.violation_error
 86        return (path, (), kwargs)
 87
 88    def clone(self) -> BaseConstraint:
 89        _, args, kwargs = self.deconstruct()
 90        return self.__class__(*args, **kwargs)
 91
 92
 93class CheckConstraint(BaseConstraint):
 94    def __init__(
 95        self,
 96        *,
 97        check: Q,
 98        name: str,
 99        violation_error: ViolationError | None = None,
100    ) -> None:
101        self.check = check
102        if not getattr(check, "conditional", False):
103            raise TypeError(
104                "CheckConstraint.check must be a Q instance or boolean expression."
105            )
106        super().__init__(name=name, violation_error=violation_error)
107
108    def to_sql(self, model: type[Model], *, not_valid: bool = False) -> str:
109        """Generate ALTER TABLE ADD CONSTRAINT CHECK SQL as a plain string."""
110        check = compile_expression_sql(model, self.check)
111        table = quote_name(model.model_options.db_table)
112        name = quote_name(self.name)
113        sql = f"ALTER TABLE {table} ADD CONSTRAINT {name} CHECK ({check})"
114        if not_valid:
115            sql += " NOT VALID"
116        return sql
117
118    def referenced_fields(self) -> set[str]:
119        """Top-level model field names referenced by `self.check`.
120
121        Walks lookup keys (`field__regex` → `field`), nested Q nodes, and
122        F-expressions in values or other source expressions.
123        """
124        fields: set[str] = set()
125
126        def visit(node: Any) -> None:
127            if isinstance(node, Q):
128                for child in node.children:
129                    visit(child)
130            elif isinstance(node, tuple) and len(node) == 2:
131                lookup, value = node
132                fields.add(lookup.split(LOOKUP_SEP, 1)[0])
133                visit(value)
134            elif isinstance(node, F):
135                fields.add(node.name.split(LOOKUP_SEP, 1)[0])
136            elif hasattr(node, "get_source_expressions"):
137                for sub in node.get_source_expressions():
138                    visit(sub)
139
140        visit(self.check)
141        return fields
142
143    def validate(
144        self, model: type[Model], instance: Model, exclude: set[str] | None = None
145    ) -> None:
146        against = instance._get_field_value_map(meta=model._model_meta, exclude=exclude)
147        # Skip the check entirely when any field referenced by `self.check` was
148        # excluded — the in-Python pipeline can't resolve a missing field's
149        # annotation, and surfacing a constraint violation here would just
150        # duplicate the field-level error that caused the exclusion.
151        if not self.referenced_fields().issubset(against):
152            return
153        try:
154            if not Q(self.check).check(against):
155                raise self._build_violation_error()
156        except FieldError:
157            pass
158
159    def _db_violation_error(
160        self, instance: Model, model: type[Model]
161    ) -> ValidationError | None:
162        return self._build_violation_error()
163
164    def __repr__(self) -> str:
165        return "<{}: check={} name={}{}>".format(
166            self.__class__.__qualname__,
167            self.check,
168            repr(self.name),
169            (
170                ""
171                if self.violation_error is None
172                else f" violation_error={self.violation_error!r}"
173            ),
174        )
175
176    def __eq__(self, other: object) -> bool:
177        if isinstance(other, CheckConstraint):
178            return (
179                self.name == other.name
180                and self.check == other.check
181                and self.violation_error == other.violation_error
182            )
183        return super().__eq__(other)
184
185    def deconstruct(self) -> tuple[str, tuple[Any, ...], dict[str, Any]]:
186        path, args, kwargs = super().deconstruct()
187        kwargs["check"] = self.check
188        return path, args, kwargs
189
190
191class UniqueConstraint(BaseConstraint):
192    expressions: tuple[ReplaceableExpression, ...]
193
194    def __init__(
195        self,
196        *expressions: str | ReplaceableExpression,
197        fields: tuple[str, ...] | list[str] = (),
198        name: str | None = None,
199        condition: Q | None = None,
200        include: tuple[str, ...] | list[str] | None = None,
201        opclasses: tuple[str, ...] | list[str] = (),
202        violation_error: ViolationError | None = None,
203    ) -> None:
204        if not name:
205            raise ValueError("A unique constraint must be named.")
206        if not expressions and not fields:
207            raise ValueError(
208                "At least one field or expression is required to define a "
209                "unique constraint."
210            )
211        if expressions and fields:
212            raise ValueError(
213                "UniqueConstraint.fields and expressions are mutually exclusive."
214            )
215        if not isinstance(condition, NoneType | Q):
216            raise TypeError("UniqueConstraint.condition must be a Q instance.")
217        if expressions and opclasses:
218            raise ValueError(
219                "UniqueConstraint.opclasses cannot be used with expressions. "
220                "Use a custom OpClass() instead."
221            )
222        if not isinstance(include, NoneType | list | tuple):
223            raise TypeError("UniqueConstraint.include must be a list or tuple.")
224        if not isinstance(opclasses, list | tuple):
225            raise TypeError("UniqueConstraint.opclasses must be a list or tuple.")
226        if opclasses and len(fields) != len(opclasses):
227            raise ValueError(
228                "UniqueConstraint.fields and UniqueConstraint.opclasses must "
229                "have the same number of elements."
230            )
231        self.fields = tuple(fields)
232        self.condition = condition
233        self.include = tuple(include) if include else ()
234        self.opclasses = opclasses
235        self.expressions = tuple(
236            F(expression) if isinstance(expression, str) else expression
237            for expression in expressions
238        )
239        super().__init__(name=name, violation_error=violation_error)
240
241    @property
242    def contains_expressions(self) -> bool:
243        return bool(self.expressions)
244
245    @property
246    def is_partial(self) -> bool:
247        return self.condition is not None
248
249    @property
250    def index_only(self) -> bool:
251        """Whether PostgreSQL can only store this as a unique index, not a constraint.
252
253        PostgreSQL rejects ALTER TABLE ADD CONSTRAINT UNIQUE USING INDEX for
254        partial indexes, expression indexes, and indexes with non-default
255        operator classes.
256        """
257        return bool(self.condition or self.expressions or self.opclasses)
258
259    def to_sql(self, model: type[Model], *, concurrently: bool = False) -> str:
260        """Generate CREATE UNIQUE INDEX or ALTER TABLE ADD CONSTRAINT UNIQUE SQL."""
261        table = quote_name(model.model_options.db_table)
262        name = quote_name(self.name)
263        condition = (
264            compile_expression_sql(model, self.condition)
265            if self.condition is not None
266            else None
267        )
268
269        if self.expressions:
270            columns_sql = compile_index_expressions_sql(model, self.expressions)
271        else:
272            col_parts = []
273            for i, field_name in enumerate(self.fields):
274                field = model._model_meta.get_forward_field(field_name)
275                col = quote_name(field.column)
276                if self.opclasses:
277                    col = f"{col} {self.opclasses[i]}"
278                col_parts.append(col)
279            columns_sql = ", ".join(col_parts)
280
281        include_sql = build_include_sql(model, self.include)
282        condition_sql = f" WHERE ({condition})" if condition else ""
283
284        if concurrently:
285            return f"CREATE UNIQUE INDEX CONCURRENTLY {name} ON {table} ({columns_sql}){include_sql}{condition_sql}"
286        elif condition or self.include or self.opclasses or self.expressions:
287            return f"CREATE UNIQUE INDEX {name} ON {table} ({columns_sql}){include_sql}{condition_sql}"
288        else:
289            return f"ALTER TABLE {table} ADD CONSTRAINT {name} UNIQUE ({columns_sql})"
290
291    def to_attach_sql(self, model: type[Model]) -> str:
292        """Generate ALTER TABLE ADD CONSTRAINT UNIQUE USING INDEX SQL.
293
294        Used after creating the unique index concurrently to attach it
295        as a named constraint.
296        """
297        table = quote_name(model.model_options.db_table)
298        name = quote_name(self.name)
299        return f"ALTER TABLE {table} ADD CONSTRAINT {name} UNIQUE USING INDEX {name}"
300
301    def __repr__(self) -> str:
302        return "<{}:{}{}{}{}{}{}{}>".format(
303            self.__class__.__qualname__,
304            "" if not self.fields else f" fields={self.fields!r}",
305            "" if not self.expressions else f" expressions={self.expressions!r}",
306            f" name={self.name!r}",
307            "" if self.condition is None else f" condition={self.condition}",
308            "" if not self.include else f" include={self.include!r}",
309            "" if not self.opclasses else f" opclasses={self.opclasses!r}",
310            (
311                ""
312                if self.violation_error is None
313                else f" violation_error={self.violation_error!r}"
314            ),
315        )
316
317    def __eq__(self, other: object) -> bool:
318        if isinstance(other, UniqueConstraint):
319            return (
320                self.name == other.name
321                and self.fields == other.fields
322                and self.condition == other.condition
323                and self.include == other.include
324                and self.opclasses == other.opclasses
325                and self.expressions == other.expressions
326                and self.violation_error == other.violation_error
327            )
328        return super().__eq__(other)
329
330    def deconstruct(self) -> tuple[str, tuple[Any, ...], dict[str, Any]]:
331        path, _args, kwargs = super().deconstruct()
332        if self.fields:
333            kwargs["fields"] = self.fields
334        if self.condition:
335            kwargs["condition"] = self.condition
336        if self.include:
337            kwargs["include"] = self.include
338        if self.opclasses:
339            kwargs["opclasses"] = self.opclasses
340        return path, self.expressions, kwargs
341
342    def validate(
343        self, model: type[Model], instance: Model, exclude: set[str] | None = None
344    ) -> None:
345        queryset = model.query
346        if self.fields:
347            lookup_kwargs = {}
348            for field_name in self.fields:
349                if exclude and field_name in exclude:
350                    return
351                field = model._model_meta.get_forward_field(field_name)
352                lookup_value = field.value_from_object(instance)
353                if lookup_value is None:
354                    # A composite constraint containing NULL value cannot cause
355                    # a violation since NULL != NULL in SQL.
356                    return
357                lookup_kwargs[field_name] = lookup_value
358            queryset = queryset.filter(**lookup_kwargs)
359        else:
360            # Ignore constraints with excluded fields.
361            if exclude:
362                for expression in self.expressions:
363                    if hasattr(expression, "flatten"):
364                        for expr in expression.flatten():  # ty: ignore[call-non-callable]
365                            if isinstance(expr, F) and expr.name in exclude:
366                                return
367                    elif isinstance(expression, F) and expression.name in exclude:
368                        return
369            replacements: dict[Any, Any] = {
370                F(field): value
371                for field, value in instance._get_field_value_map(
372                    meta=model._model_meta, exclude=exclude
373                ).items()
374            }
375            expressions = []
376            for expr in self.expressions:
377                # Ignore ordering.
378                if isinstance(expr, OrderBy):
379                    expr = expr.expression
380                expressions.append(Exact(expr, expr.replace_expressions(replacements)))
381            queryset = queryset.filter(*expressions)
382        model_class_id = instance.id
383        if not instance._state.adding and model_class_id is not None:
384            queryset = queryset.exclude(id=model_class_id)
385        if not self.condition:
386            if queryset.exists():
387                raise self._build_unique_violation(instance, model)
388        else:
389            against = instance._get_field_value_map(
390                meta=model._model_meta, exclude=exclude
391            )
392            try:
393                if (self.condition & Exists(queryset.filter(self.condition))).check(
394                    against
395                ):
396                    raise self._build_unique_violation(instance, model)
397            except FieldError:
398                pass
399
400    def _build_unique_violation(
401        self, instance: Model, model: type[Model]
402    ) -> ValidationError:
403        """Build the ValidationError for a unique violation.
404
405        Single-field unique constraints route the error to that field via the
406        dict form so it surfaces under the field rather than NON_FIELD_ERRORS.
407        """
408        single_field = self.fields[0] if len(self.fields) == 1 else None
409
410        if self.violation_error is not None:
411            err = self._build_violation_error()
412            # Only auto-route flat errors. A ValidationError that already has
413            # an error_dict (from dict-form input or a caller-built instance)
414            # already declares its own field routing — don't override it.
415            if single_field and not hasattr(err, "error_dict"):
416                return ValidationError({single_field: [err]})
417            return err
418
419        if self.fields:
420            err = self._unique_error_message(instance, model, self.fields)
421            if single_field:
422                return ValidationError({single_field: [err]})
423            return err
424        return ValidationError(f'Constraint "{self.name}" is violated.')
425
426    def _unique_error_message(
427        self,
428        instance: Model,
429        model: type[Model],
430        unique_check: tuple[str, ...],
431    ) -> ValidationError:
432        """Build the ValidationError describing a violation of `unique_check`,
433        using each field's own `unique_error_message` format string."""
434        meta = model._model_meta
435
436        params: dict[str, Any] = {
437            "model": instance,
438            "model_class": model,
439            "model_name": model.model_options.model_name,
440            "unique_check": unique_check,
441        }
442
443        if len(unique_check) == 1:
444            field = meta.get_forward_field(unique_check[0])
445            params["field_label"] = field.name
446            return ValidationError(
447                message=field.unique_error_message,
448                code="unique",
449                params=params,
450            )
451
452        field_names = [meta.get_forward_field(f).name for f in unique_check]
453        # Put an "and" before the last one.
454        field_names[-1] = f"and {field_names[-1]}"
455        # Comma-join when more than two, otherwise just space-join.
456        sep = ", " if len(field_names) > 2 else " "
457        params["field_label"] = sep.join(field_names)
458
459        # Use the first field's message format.
460        message = meta.get_forward_field(unique_check[0]).unique_error_message
461        return ValidationError(
462            message=message,
463            code="unique",
464            params=params,
465        )
466
467    def _db_violation_error(
468        self, instance: Model, model: type[Model]
469    ) -> ValidationError | None:
470        return self._build_unique_violation(instance, model)