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)