diff --git a/.claude/rules/plain-postgres.md b/.claude/rules/plain-postgres.md index 576c7b02b6..6b7b0878ce 100644 --- a/.claude/rules/plain-postgres.md +++ b/.claude/rules/plain-postgres.md @@ -111,11 +111,11 @@ Use `Model.query` to build querysets (e.g., `User.query.filter(is_active=True)`) - Use `.annotate(Count(...))` instead of calling `.count()` per row - Fetch all data in the view — templates should never trigger queries - Use `.exists()` not `.count() > 0`, `.count()` not `len(qs)` -- Use `bulk_create`/`bulk_update` for batch ops, `.update()`/`.delete()` for mass ops +- Use `bulk_create`/`bulk_update` for batch ops, `bulk_upsert` for atomic insert-or-update, `.update()`/`.delete()` for mass ops - Use `.values_list()` when you only need specific columns - Wrap multi-step writes in `transaction.atomic()` - Instance writes are `obj.create()` (always INSERT) and `obj.update()` (always UPDATE; `update(fields=[...])` limits the columns) — there is no `save()`, `force_insert`, or `force_update`. Constructing an instance then `create()`-ing it inserts; a hand-set `id` that collides raises `IntegrityError`. -- `create()`/`update()` raise `ValidationError` (not raw `psycopg.IntegrityError`) on a declared unique/check constraint violation or a foreign key pointing at a missing row, even a raced one — the DB enforces it, so inside an open `transaction.atomic()` the violation aborts the transaction (wrap the write in its own `atomic()` to catch and keep using the transaction). Set-based writes (`QuerySet.update()`/`bulk_create()`) and `delete()` blocked by `RESTRICT` raise raw `psycopg.IntegrityError`. Retrying on conflict? `except (psycopg.IntegrityError, ValidationError)`, or `bulk_create(..., update_conflicts=True)` +- `create()`/`update()` raise `ValidationError` (not raw `psycopg.IntegrityError`) on a declared unique/check constraint violation or a foreign key pointing at a missing row, even a raced one — the DB enforces it, so inside an open `transaction.atomic()` the violation aborts the transaction (wrap the write in its own `atomic()` to catch and keep using the transaction). Set-based writes (`QuerySet.update()`/`bulk_create()`) and `delete()` blocked by `RESTRICT` raise raw `psycopg.IntegrityError`. Retrying on conflict? `except (psycopg.IntegrityError, ValidationError)`, or `bulk_upsert(objs, update_fields=[...], unique_fields=[...])` for an atomic insert-or-update - Always paginate list queries — unbounded querysets get slower as data grows Run `uv run plain docs postgres` for full patterns with code examples. diff --git a/plain-cache/plain/cache/core.py b/plain-cache/plain/cache/core.py index 9d6fe6b0cd..65c48cc048 100644 --- a/plain-cache/plain/cache/core.py +++ b/plain-cache/plain/cache/core.py @@ -106,12 +106,12 @@ def set_many( if not mapping: return - # bulk_create fires pre_save, so updated_at's update_now stamps a fresh - # now() at write time on its own. created_at (no update_now) would - # otherwise fall to its DB default, evaluated a hair later -- leaving a - # brand-new row with updated_at < created_at. Stamp created_at from an - # up-front `now` so created_at <= updated_at; it's omitted from - # update_fields, so it's preserved on conflict. + # bulk_upsert fires pre_save and refreshes update_now columns on the + # conflict path, so updated_at looks after itself. created_at (no + # update_now) would otherwise fall to its DB default, evaluated a hair + # later -- leaving a brand-new row with updated_at < created_at. Stamp + # created_at from an up-front `now` so created_at <= updated_at; being + # DB-owned it can't be named in update_fields, so it survives conflicts. now = timezone.now() expires_at = _coerce_expiration(expiration, now=now) items = [] @@ -122,11 +122,11 @@ def set_many( # construction so created_at <= updated_at (see comment above). item.created_at = now items.append(item) - self._model.query.bulk_create( + model = self._model + model.query.bulk_upsert( items, - update_conflicts=True, - update_fields=["value", "expires_at", "updated_at"], - unique_fields=["key"], + update_fields=[model.value, model.expires_at], + unique_fields=[model.key], ) def get_or_set( @@ -241,7 +241,7 @@ def touch(self, key: str, *, expiration: Expiration = None) -> bool: """ # QuerySet.update() issues a direct SQL UPDATE and does NOT fire pre_save, # so updated_at's update_now won't bump on its own -- stamp it by hand. - # (set_many() relies on pre_save instead, since bulk_create does fire it.) + # (set_many() relies on pre_save instead, since bulk_upsert does fire it.) now = timezone.now() updated = ( self._model.query.live() diff --git a/plain-cache/tests/internal/test_set_timestamps.py b/plain-cache/tests/internal/test_set_timestamps.py index 0ef86fad0a..c1adc54dce 100644 --- a/plain-cache/tests/internal/test_set_timestamps.py +++ b/plain-cache/tests/internal/test_set_timestamps.py @@ -1,6 +1,6 @@ """Timestamp invariants for the set-based write paths. -bulk_create fires pre_save (so updated_at's update_now bumps on its own) while +bulk_upsert fires pre_save (so updated_at's update_now bumps on its own) while QuerySet.update() does not -- see core.py for why set_many stamps created_at and touch stamps updated_at. These pin the observable invariants those choices buy. """ diff --git a/plain-postgres/plain/postgres/README.md b/plain-postgres/plain/postgres/README.md index fed3e6c3fc..46e00109e4 100644 --- a/plain-postgres/plain/postgres/README.md +++ b/plain-postgres/plain/postgres/README.md @@ -464,6 +464,64 @@ for name in names: Tag.query.bulk_create([Tag(name=name) for name in names]) ``` +`bulk_create` is insert-only. To insert new rows and update the ones that +already exist in a single statement, use `bulk_upsert` (below). + +#### Use `bulk_upsert` to insert-or-update in one statement + +`bulk_upsert(objs, *, update_fields, unique_fields, batch_size=None)` issues one +`INSERT ... ON CONFLICT (unique_fields) DO UPDATE SET ... RETURNING` per batch. +Rows that don't exist yet are inserted; rows that collide on `unique_fields` have +their `update_fields` overwritten. You get back the objects you passed in, in the +order you passed them (a new list — `objs` itself is never reordered), each with +its DB-generated fields (primary key, DB defaults) populated. + +```python +# Insert new items, refresh `value`/`expires_at` on any existing key. +CachedItem.query.bulk_upsert( + [CachedItem(key=k, value=v, expires_at=exp) for k, v in items], + update_fields=[CachedItem.value, CachedItem.expires_at], + unique_fields=[CachedItem.key], +) +``` + +- `update_fields` and `unique_fields` take field references (`Model.field`), not + strings. A foreign key is named by the relation itself — `Model.tenant`, which + resolves to the `tenant_id` column. (This is the one write API that takes + `Model.fk`. `returning()` refuses it, because there it would be ambiguous with + asking for the whole related object; here a column list can only mean the + column.) +- `unique_fields` must name the **primary key** or a `UniqueConstraint` declared + on the model (no condition, no expressions) — this is the conflict target. A + unique `Index` is not enough; declare a `UniqueConstraint`. +- `update_fields` must be concrete, non-primary-key, must not name the same + column twice (Postgres assigns each column once per statement), and must not + overlap `unique_fields`. A column the database fills in (`create_now`, + `generate=True`, `RandomStringField`) can't be named either — the update would + overwrite the stored value with a freshly evaluated default. +- Every object must have a non-null value for every unique field. `NULL` never + conflicts in Postgres, so it can't be upserted. A database-generated column + (`create_now`, `generate=True`, `RandomStringField`) can't be a unique field + either — your objects never hold its value, so it could never conflict. Nor + can an `update_now=True` column, which is stamped again on every write. +- **Two objects with the same unique key in one batch raise `ValueError`.** + Postgres won't touch a row twice in one statement. Split across batches it's + allowed — the first inserts, the second updates, and the later write wins. +- **`update_now=True` columns are refreshed on a conflict automatically.** You + don't name them in `update_fields`; a row that gets updated gets a fresh + stamp, and the object handed back carries the same one. +- **An `id` you set is kept on the insert path; on a conflict the stored row + wins.** A new row is written with the `id` you gave it. A conflicting one + already has an `id`, and that is the one hydrated back onto your object — the + row in the table is the truth. An `id` that collides with a _different_ row + raises `psycopg.errors.UniqueViolation`, like any set-based write. +- Every object is sorted by its conflict key before anything is sent, so + concurrent `bulk_upsert` calls over overlapping keys lock rows in the same + order and can't deadlock each other. Returned rows are mapped onto the objects + by position, exactly as `bulk_create` does. +- Like `bulk_create`, the write is against the table: a filter on the queryset + you call it from doesn't narrow or exclude anything. + #### Use queryset `.update()` / `.delete()` for mass operations ```python @@ -529,7 +587,7 @@ for row in deleted: - **`returning(Model.field, ...)`** returns a list of dicts with only those columns. Pass field references (`Model.field`), not strings; a many-to-many field or one from another model raises an error at the `returning()` call. - **A foreign key can't be named here.** At class level `Model.fk` is the relation — that is what lets `where()` traverse it, as in `Child.parent.name.equals(...)` — not its column, so `returning(Child.parent)` raises `FieldError`. Foreign key columns come back through no-argument `returning()`, which hands you whole instances. - Without `returning()`, `update()`/`delete()` return an `int` as before. -- `returning()` only applies to `update()` and `delete()`. Any other write on the same queryset — `create()`, `bulk_create()`, `bulk_update()`, `get_or_create()`, `update_or_create()` — raises `TypeError` rather than quietly dropping it. +- `returning()` only applies to `update()` and `delete()`. Any other write on the same queryset — `create()`, `bulk_create()`, `bulk_upsert()`, `bulk_update()`, `get_or_create()`, `update_or_create()` — raises `TypeError` rather than quietly dropping it. - `returning()` keeps the queryset's own class, so a custom `QuerySet` and its methods survive it. Chain your own methods before `returning()` — a type checker sees the returning shape after it, not your subclass. A row lock belongs on the read side of the write, and it composes in either order. The write then needs an open `transaction.atomic()`, and is emitted as a locking sub-select so the lock has somewhere to live — see [Locking a set-based write](#locking-a-set-based-write). @@ -1388,7 +1446,7 @@ except (psycopg.IntegrityError, ValidationError): ... # lost a race — reload and retry, or report it ``` -For a plain insert-or-update with no per-row logic, `bulk_create(..., update_conflicts=True, unique_fields=[...])` is an atomic upsert with no race to catch. +For a plain insert-or-update with no per-row logic, `bulk_upsert(objs, update_fields=[...], unique_fields=[...])` is an atomic upsert with no race to catch. ### Indexes and constraints diff --git a/plain-postgres/plain/postgres/agents/.claude/rules/plain-postgres.md b/plain-postgres/plain/postgres/agents/.claude/rules/plain-postgres.md index 576c7b02b6..6b7b0878ce 100644 --- a/plain-postgres/plain/postgres/agents/.claude/rules/plain-postgres.md +++ b/plain-postgres/plain/postgres/agents/.claude/rules/plain-postgres.md @@ -111,11 +111,11 @@ Use `Model.query` to build querysets (e.g., `User.query.filter(is_active=True)`) - Use `.annotate(Count(...))` instead of calling `.count()` per row - Fetch all data in the view — templates should never trigger queries - Use `.exists()` not `.count() > 0`, `.count()` not `len(qs)` -- Use `bulk_create`/`bulk_update` for batch ops, `.update()`/`.delete()` for mass ops +- Use `bulk_create`/`bulk_update` for batch ops, `bulk_upsert` for atomic insert-or-update, `.update()`/`.delete()` for mass ops - Use `.values_list()` when you only need specific columns - Wrap multi-step writes in `transaction.atomic()` - Instance writes are `obj.create()` (always INSERT) and `obj.update()` (always UPDATE; `update(fields=[...])` limits the columns) — there is no `save()`, `force_insert`, or `force_update`. Constructing an instance then `create()`-ing it inserts; a hand-set `id` that collides raises `IntegrityError`. -- `create()`/`update()` raise `ValidationError` (not raw `psycopg.IntegrityError`) on a declared unique/check constraint violation or a foreign key pointing at a missing row, even a raced one — the DB enforces it, so inside an open `transaction.atomic()` the violation aborts the transaction (wrap the write in its own `atomic()` to catch and keep using the transaction). Set-based writes (`QuerySet.update()`/`bulk_create()`) and `delete()` blocked by `RESTRICT` raise raw `psycopg.IntegrityError`. Retrying on conflict? `except (psycopg.IntegrityError, ValidationError)`, or `bulk_create(..., update_conflicts=True)` +- `create()`/`update()` raise `ValidationError` (not raw `psycopg.IntegrityError`) on a declared unique/check constraint violation or a foreign key pointing at a missing row, even a raced one — the DB enforces it, so inside an open `transaction.atomic()` the violation aborts the transaction (wrap the write in its own `atomic()` to catch and keep using the transaction). Set-based writes (`QuerySet.update()`/`bulk_create()`) and `delete()` blocked by `RESTRICT` raise raw `psycopg.IntegrityError`. Retrying on conflict? `except (psycopg.IntegrityError, ValidationError)`, or `bulk_upsert(objs, update_fields=[...], unique_fields=[...])` for an atomic insert-or-update - Always paginate list queries — unbounded querysets get slower as data grows Run `uv run plain docs postgres` for full patterns with code examples. diff --git a/plain-postgres/plain/postgres/constants.py b/plain-postgres/plain/postgres/constants.py index cec1b9b90f..c1d7bd5356 100644 --- a/plain-postgres/plain/postgres/constants.py +++ b/plain-postgres/plain/postgres/constants.py @@ -9,5 +9,4 @@ class OnConflict(Enum): - IGNORE = "ignore" UPDATE = "update" diff --git a/plain-postgres/plain/postgres/dialect.py b/plain-postgres/plain/postgres/dialect.py index 63e2f330ac..b73a3c6f98 100644 --- a/plain-postgres/plain/postgres/dialect.py +++ b/plain-postgres/plain/postgres/dialect.py @@ -592,8 +592,6 @@ def on_conflict_suffix_sql( update_fields: Iterable[str], unique_fields: Iterable[str], ) -> str: - if on_conflict == OnConflict.IGNORE: - return "ON CONFLICT DO NOTHING" if on_conflict == OnConflict.UPDATE: return "ON CONFLICT({}) DO UPDATE SET {}".format( ", ".join(map(quote_name, unique_fields)), diff --git a/plain-postgres/plain/postgres/options.py b/plain-postgres/plain/postgres/options.py index 2876f63d15..332e8f3d7c 100644 --- a/plain-postgres/plain/postgres/options.py +++ b/plain-postgres/plain/postgres/options.py @@ -207,6 +207,17 @@ def total_unique_constraints(self) -> list[Any]: ) ] + def unique_fields_match_constraint(self, field_names: set[str | None]) -> bool: + """True if field_names names the primary key, or a UniqueConstraint on + the model that has no condition and no expressions.""" + pk_field = self.model._model_meta.get_forward_field("id") + if field_names == {pk_field.name}: + return True + for constraint in self.total_unique_constraints: + if set(constraint.fields) == field_names: + return True + return False + def __repr__(self) -> str: return f"" diff --git a/plain-postgres/plain/postgres/query.py b/plain-postgres/plain/postgres/query.py index 44c6d22c1a..de02e820c4 100644 --- a/plain-postgres/plain/postgres/query.py +++ b/plain-postgres/plain/postgres/query.py @@ -5,9 +5,12 @@ from __future__ import annotations import copy +import datetime +import json import operator import warnings from collections.abc import Callable, Iterator, Sequence +from decimal import Decimal from functools import cached_property from itertools import islice from typing import TYPE_CHECKING, Any, Never, Self, cast, overload @@ -21,6 +24,7 @@ PLAIN_VERSION_PICKLE_KEY, get_connection, ) +from plain.postgres.dialect import get_json_dumps from plain.postgres.exceptions import ( FieldDoesNotExist, FieldError, @@ -32,6 +36,7 @@ PrimaryKeyField, ) from plain.postgres.fields.base import ColumnField +from plain.postgres.fields.json import JSONField from plain.postgres.functions import Cast from plain.postgres.query_utils import Q, condition_origins_of from plain.postgres.sql import ( @@ -55,6 +60,59 @@ from plain.postgres import Model +def conflict_sort_value(field: Field, value: Any) -> str: + """One component of the order bulk_upsert() sends its batches in. + + Concurrent callers only have to agree on an order, not on a meaningful + one, so the requirement is narrow: two callers holding the same logical + key must render it the same way, and comparing the results must never + raise. Everything here serves that. + + The value arrives already through `get_prep_value`, which settles most of + it -- a TimeZoneField's ZoneInfo is its name by then, a UUID string is a + UUID, a naive datetime is aware. What is left is the spellings Postgres + holds equal that `str()` would not: a str subclass that renders itself + some other way, a bytea handed back as a memoryview, the same instant + written at two offsets, a signed zero, a decimal's scale. + """ + if isinstance(field, JSONField): + # Encode with the field's own encoder, which stringifies non-string + # object keys, and only then re-parse and dump with the keys sorted -- + # sorting them first would compare an int key against a str one and + # raise. Two equal objects written with their keys in either order + # then render the same. + return json.dumps( + json.loads(get_json_dumps(field.encoder)(value)), sort_keys=True + ) + if isinstance(value, str): + # A StrEnum member or a SafeString compares equal to the plain string, + # which is all the column holds, but renders itself differently. + # str.__str__ goes around the override. + return str.__str__(value) + if isinstance(value, memoryview | bytearray): + # psycopg hands a bytea column back as a memoryview, whose str() is + # where it happens to sit in memory. + return str(bytes(value)) + if isinstance(value, datetime.datetime) and value.tzinfo is not None: + # timestamptz stores the instant, not the offset it was written at. + return str(value.astimezone(datetime.UTC)) + if isinstance(value, float): + # Postgres holds -0.0 and 0.0 equal; adding zero folds the sign. + return repr(value + 0.0) + if isinstance(value, Decimal): + # numeric holds Decimal("1.0") and Decimal("1.00") equal. normalize() + # gives them one spelling, and abs() folds the negative zero it keeps. + # An exponent too large to normalize is left as it is -- Postgres + # rejects it on write, with the better error. + try: + value = value.normalize() + if value == 0: + value = abs(value) + except ArithmeticError: + pass + return str(value) + + # The maximum number of results to fetch in a get() query. MAX_GET_RESULTS = 21 @@ -716,64 +774,25 @@ def create(self, **kwargs: Any) -> T: obj.create() return obj - def _prepare_for_bulk_create(self, objs: list[T]) -> None: + def _prepare_for_bulk_create(self, objs: list[T], *, operation_name: str) -> None: # The identity PK is the only PK type, so there's no literal Python # default to materialize -- obj.id stays None and the INSERT takes the # DB's DEFAULT path. for obj in objs: - obj._prepare_related_fields_for_save(operation_name="bulk_create") - - def _check_bulk_create_options( - self, - update_conflicts: bool, - update_fields: list[Field] | None, - unique_fields: list[Field] | None, - ) -> OnConflict | None: - if update_conflicts: - if not update_fields: - raise ValueError( - "Fields that will be updated when a row insertion fails " - "on conflicts must be provided." - ) - if not unique_fields: - raise ValueError( - "Unique fields that can trigger the upsert must be provided." - ) - # Updating primary keys and many-to-many fields is forbidden. - from plain.postgres.fields.related import ManyToManyField - - if any(isinstance(f, ManyToManyField) for f in update_fields): - raise ValueError( - "bulk_create() cannot be used with many-to-many fields in " - "update_fields." - ) - if any(f.primary_key for f in update_fields): - raise ValueError( - "bulk_create() cannot be used with primary keys in update_fields." - ) - if unique_fields: - from plain.postgres.fields.related import ManyToManyField - - if any(isinstance(f, ManyToManyField) for f in unique_fields): - raise ValueError( - "bulk_create() cannot be used with many-to-many fields " - "in unique_fields." - ) - return OnConflict.UPDATE - return None + obj._prepare_related_fields_for_save(operation_name=operation_name) def bulk_create( self, objs: Sequence[T], batch_size: int | None = None, - update_conflicts: bool = False, - update_fields: list[str] | None = None, - unique_fields: list[str] | None = None, ) -> list[T]: """ Insert each of the instances into the database. Do *not* call save() on each of the instances. Primary keys are set on the objects via the PostgreSQL RETURNING clause. Multi-table models are not supported. + + This is insert-only -- to insert-or-update on a conflict, use + bulk_upsert(). """ self._reject_returning("bulk_create") if batch_size is not None and batch_size <= 0: @@ -783,23 +802,8 @@ def bulk_create( if not objs: return objs meta = self.model._model_meta - unique_fields_objs: list[Field] | None = None - update_fields_objs: list[Field] | None = None - if unique_fields: - unique_fields_objs = [ - meta.get_forward_field(name) for name in unique_fields - ] - if update_fields: - update_fields_objs = [ - meta.get_forward_field(name) for name in update_fields - ] - on_conflict = self._check_bulk_create_options( - update_conflicts, - update_fields_objs, - unique_fields_objs, - ) fields = meta.fields - self._prepare_for_bulk_create(objs) + self._prepare_for_bulk_create(objs, operation_name="bulk_create") with transaction.atomic(savepoint=False): objs_with_id, objs_without_id = partition(lambda o: o.id is None, objs) if objs_with_id: @@ -807,9 +811,6 @@ def bulk_create( objs_with_id, fields, batch_size, - on_conflict=on_conflict, - update_fields=update_fields_objs, - unique_fields=unique_fields_objs, ) id_field = meta.get_forward_field("id") for obj_with_id, results in zip(objs_with_id, returned_columns): @@ -824,12 +825,13 @@ def bulk_create( objs_without_id, fields, batch_size, - on_conflict=on_conflict, - update_fields=update_fields_objs, - unique_fields=unique_fields_objs, ) - if on_conflict is None: - assert len(returned_columns) == len(objs_without_id) + # Postgres emits one RETURNING row per VALUES row, in order, so + # the rows can be zipped straight onto the objects. bulk_upsert() + # relies on the same guarantee for its ON CONFLICT batches -- the + # two stand or fall together, and + # tests/internal/test_returning_order.py pins it. + assert len(returned_columns) == len(objs_without_id) for obj_without_id, results in zip(objs_without_id, returned_columns): for result, field in zip(results, meta.db_returning_fields): setattr(obj_without_id, field.name, result) @@ -837,6 +839,224 @@ def bulk_create( return objs + def _resolve_bulk_upsert_fields( + self, + update_fields: Sequence[Field[Any] | type[Model]], + unique_fields: Sequence[Field[Any] | type[Model]], + ) -> tuple[list[Field], list[Field]]: + """Check both bulk_upsert() field lists and return the columns they + name, with any `Model.fk` reference resolved to its foreign key + column.""" + object_name = self.model.model_options.object_name + + unique_columns = self._validate_field_refs( + unique_fields, where="bulk_upsert() unique_fields" + ) + update_columns = self._validate_field_refs( + update_fields, where="bulk_upsert() update_fields" + ) + + if not unique_columns: + raise ValueError("bulk_upsert() requires unique_fields.") + # A conflict key the caller doesn't control can never actually conflict, + # so the upsert would silently be an insert every time. + for field in unique_columns: + if field.db_returning and not field.primary_key: + raise ValueError( + f"bulk_upsert() cannot use {object_name}.{field.name} in " + "unique_fields: the database generates its value, so the " + "objects never carry one to conflict on." + ) + if field.auto_fills_on_save: + raise ValueError( + f"bulk_upsert() cannot use {object_name}.{field.name} in " + "unique_fields: it is stamped again on every write, so it " + "can never be a stable conflict key." + ) + if not self.model.model_options.unique_fields_match_constraint( + {f.name for f in unique_columns} + ): + names = [f.name for f in unique_columns] + raise ValueError( + f"bulk_upsert() unique_fields {names} on {object_name} must name " + "the primary key or a UniqueConstraint declared on the model " + "without a condition or expressions." + ) + + if not update_columns: + raise ValueError("bulk_upsert() requires update_fields.") + if any(not isinstance(f, ColumnField) for f in update_columns): + raise ValueError("bulk_upsert() update_fields must be database columns.") + if any(f.primary_key for f in update_columns): + raise ValueError("bulk_upsert() cannot update primary key fields.") + for field in update_columns: + # A database-owned value (create_now, generate=True, + # RandomStringField) isn't the caller's to overwrite: EXCLUDED + # carries a freshly evaluated default, so naming one here would + # reset a creation timestamp on every update. A column that is also + # update_now is exempt -- rewriting it is the whole point. + if field.db_returning and not field.auto_fills_on_save: + raise ValueError( + f"bulk_upsert() cannot update {object_name}.{field.name}: " + "the database generates its value, so the update would " + "overwrite the stored one with a fresh default." + ) + repeated = sorted( + { + field.name + for field in update_columns + if sum(other.name == field.name for other in update_columns) > 1 + } + ) + if repeated: + raise ValueError( + f"bulk_upsert() update_fields names {repeated} more than once; " + "Postgres assigns each column once per statement." + ) + overlap = {f.name for f in update_columns} & {f.name for f in unique_columns} + if overlap: + raise ValueError( + "bulk_upsert() update_fields cannot overlap unique_fields: " + f"{sorted(overlap)}." + ) + + return update_columns, unique_columns + + def bulk_upsert( + self, + objs: Sequence[T], + *, + update_fields: Sequence[Field[Any] | type[Model]], + unique_fields: Sequence[Field[Any] | type[Model]], + batch_size: int | None = None, + ) -> list[T]: + """ + Insert each instance, updating update_fields on any row that already + exists for the unique_fields key. Issues one + INSERT ... ON CONFLICT (unique_fields) DO UPDATE ... RETURNING per batch. + + Both inserted and updated objects come back with their DB-returned + fields (primary key, DB defaults) populated, in the order they were + passed in. update_fields and unique_fields take field references + (`Model.field`); unique_fields must name the primary key or a + UniqueConstraint declared on the model. + + A conflicting row is written with the named update_fields plus every + update_now column on the model, so the stored row and the returned + object agree on when it was last touched. + + bulk_upsert() carries its own RETURNING to populate the objects, so a + prior returning() has nothing to add and is refused. + """ + self._reject_returning("bulk_upsert") + if batch_size is not None and batch_size <= 0: + raise ValueError("Batch size must be a positive integer.") + + update_columns, unique_columns = self._resolve_bulk_upsert_fields( + update_fields, unique_fields + ) + + objs = list(objs) + if not objs: + return objs + + meta = self.model._model_meta + object_name = self.model.model_options.object_name + self._prepare_for_bulk_create(objs, operation_name="bulk_upsert") + + # A NULL conflict key never conflicts in Postgres, so the row would + # always insert and the upsert would quietly be an insert. + sort_keys = [] + for obj in objs: + key = [] + for field in unique_columns: + # Prepared once here and handed to the sort key, rather + # than prepared again inside it. A malformed value is rejected + # at this point, before any statement goes out. + value = field.get_prep_value(field.value_from_object(obj)) + if value is None: + raise ValueError( + f"bulk_upsert() requires a non-null {field.name} on every " + "object; NULL never conflicts in Postgres, so it cannot " + "be upserted." + ) + key.append(conflict_sort_value(field, value)) + sort_keys.append(tuple(key)) + + # An update_now column is stamped by pre_save on the way in, so the + # object already holds a fresh value whether it inserts or updates. + # Setting it from EXCLUDED on the conflict path too is what keeps the + # stored row and the returned object agreeing -- and it's what + # update_now means. The caller doesn't have to name it. + conflict_update_columns = list(update_columns) + for field in meta.fields: + if field.auto_fills_on_save and field not in conflict_update_columns: + conflict_update_columns.append(field) + + # An object that already carries an id inserts with it; one that + # doesn't lets Postgres generate the identity value. bulk_create() + # splits the same way -- an id the caller set is theirs, not ours to + # throw away. When the primary key *is* the conflict target every + # object has one, so every row has the same shape. + fields = meta.fields + fields_without_pk = [f for f in fields if not isinstance(f, PrimaryKeyField)] + pk_is_unique = any(f.primary_key for f in unique_columns) + + # Lock rows in conflict-key order, so two callers touching overlapping + # keys can't deadlock each other. sorted() is stable, so equal keys keep + # their input order and the objects themselves are never compared. objs + # is left alone -- the caller gets its own order back. + # + # The sort has to span *every* object rather than each shape on its own: + # two callers holding the same keys but different ids would otherwise + # lock them in different orders, which is the deadlock this exists to + # avoid. So walk the sorted objects and start a new statement only where + # the shape changes -- an extra statement only where ids interleave. + runs: list[tuple[list[T], Sequence[Field]]] = [] + for position in sorted(range(len(objs)), key=lambda p: sort_keys[p]): + obj = objs[position] + insert_fields = ( + fields if pk_is_unique or obj.id is not None else fields_without_pk + ) + if runs and runs[-1][1] is insert_fields: + runs[-1][0].append(obj) + else: + runs.append(([obj], insert_fields)) + + with transaction.atomic(savepoint=False): + for sent_objs, insert_fields in runs: + try: + returned_rows = self._batched_insert( + sent_objs, + insert_fields, + batch_size, + on_conflict=OnConflict.UPDATE, + update_fields=conflict_update_columns, + unique_fields=unique_columns, + ) + except psycopg.errors.CardinalityViolation as exc: + names = [f.name for f in unique_columns] + raise ValueError( + f"bulk_upsert() sent two {object_name} objects with the " + f"same {names} in one statement, which Postgres refuses " + "-- it can only touch a row once per statement. Collapse " + "the duplicates before calling." + ) from exc + + # Postgres emits one RETURNING row per VALUES row, in order, on + # the DO UPDATE path as much as the insert path, so the rows + # come back in the order the objects were sent. bulk_create() + # maps its rows onto objects by position for the same reason -- + # the two stand or fall together, and + # tests/internal/test_returning_order.py pins it. + assert len(returned_rows) == len(sent_objs) + for obj, row in zip(sent_objs, returned_rows): + for index, field in enumerate(meta.db_returning_fields): + setattr(obj, field.name, row[index]) + obj._state.adding = False + + return objs + def bulk_update( self, objs: Sequence[T], fields: list[str], batch_size: int | None = None ) -> int: @@ -1066,6 +1286,49 @@ def returning(self, *fields: Field[Any]) -> ReturningQuerySet[T, Any]: assert clone._returning_fields, "returning() selected no columns" return cast("ReturningQuerySet[T, Any]", clone) + def _validate_field_refs(self, fields: Sequence[Any], *, where: str) -> list[Field]: + """Require each item to be a Field reference on this queryset's model, + and return the columns those references name. + + `Model.fk` is a ForwardForeignKeyDescriptor rather than a Field -- that + is what lets where() traverse to the related model -- so there is no + other way to name the foreign key column. The write APIs that take + column lists unwrap it here. returning() is the exception and refuses + it outright, which it does before calling this (see + _validated_returning_fields). + + `where` names the call in the error (e.g. "returning()", "bulk_upsert() + unique_fields") so a bad argument points the user at Model.field. + """ + # Local import: related_descriptors imports this module at load time. + from plain.postgres.fields.related_descriptors import ( + ForwardForeignKeyDescriptor, + ) + + object_name = self.model.model_options.object_name + columns = [] + for field in fields: + if isinstance(field, ForwardForeignKeyDescriptor): + field = field._field + if isinstance(field, str): + raise TypeError( + f"{where} takes field references, not strings. " + f"Pass {object_name}.{field} instead of {field!r}." + ) + if not isinstance(field, Field): + raise TypeError( + f"{where} takes field references like {object_name}., " + f"not {field!r}." + ) + if field.model is not self.model: + raise FieldError( + f"{where} cannot use {field.model.model_options.object_name}." + f"{field.name}: it belongs to a different model, not " + f"{object_name}." + ) + columns.append(field) + return columns + def _validated_returning_fields( self, fields: tuple[Field[Any], ...] ) -> list[Field]: @@ -1740,34 +2003,32 @@ def _batched_insert( objs: list[T], fields: Sequence[Field], batch_size: int | None, + *, on_conflict: OnConflict | None = None, update_fields: list[Field] | None = None, unique_fields: list[Field] | None = None, ) -> list[tuple[Any, ...]]: """ - Helper method for bulk_create() to insert objs one batch at a time. + Helper method for bulk_create()/bulk_upsert() to insert objs one batch + at a time, collecting the RETURNING rows from every batch. Pass the + on_conflict kwargs to run each batch as ON CONFLICT DO UPDATE. """ + returning_fields = self.model._model_meta.db_returning_fields max_batch_size = max(len(objs), 1) batch_size = min(batch_size, max_batch_size) if batch_size else max_batch_size - inserted_rows = [] + returned_rows = [] for item in [objs[i : i + batch_size] for i in range(0, len(objs), batch_size)]: - if on_conflict is None: - inserted_rows.extend( - self._insert( # ty: ignore[invalid-argument-type] - item, - fields=fields, - returning_fields=self.model._model_meta.db_returning_fields, - ) - ) - else: - self._insert( + returned_rows.extend( + self._insert( # ty: ignore[invalid-argument-type] item, fields=fields, + returning_fields=returning_fields, on_conflict=on_conflict, update_fields=update_fields, unique_fields=unique_fields, ) - return inserted_rows + ) + return returned_rows def _chain(self) -> Self: """ diff --git a/plain-postgres/tests/app/examples/migrations/0022_upsertitem.py b/plain-postgres/tests/app/examples/migrations/0022_upsertitem.py new file mode 100644 index 0000000000..8779d8e97e --- /dev/null +++ b/plain-postgres/tests/app/examples/migrations/0022_upsertitem.py @@ -0,0 +1,19 @@ +from plain.postgres import migrations + +from plain import postgres + + +class Migration(migrations.Migration): + dependencies = (("examples", "0021_returningevent"),) + + operations = ( + migrations.CreateModel( + name="UpsertItem", + fields=[ + ("id", postgres.PrimaryKeyField()), + ("key", postgres.TextField(max_length=100)), + ("label", postgres.TextField(default="")), + ("value", postgres.IntegerField(default=0)), + ], + ), + ) diff --git a/plain-postgres/tests/app/examples/migrations/0023_upsertpair.py b/plain-postgres/tests/app/examples/migrations/0023_upsertpair.py new file mode 100644 index 0000000000..55fbd897c9 --- /dev/null +++ b/plain-postgres/tests/app/examples/migrations/0023_upsertpair.py @@ -0,0 +1,21 @@ +# Generated by Plain 0.163.1 on 2026-09-19 18:54 + +from plain.postgres import migrations + +from plain import postgres + + +class Migration(migrations.Migration): + dependencies = (("examples", "0022_upsertitem"),) + + operations = ( + migrations.CreateModel( + name="UpsertPair", + fields=[ + ("id", postgres.PrimaryKeyField()), + ("bucket", postgres.TextField(max_length=100)), + ("slug", postgres.TextField(max_length=100)), + ("value", postgres.IntegerField(default=0)), + ], + ), + ) diff --git a/plain-postgres/tests/app/examples/migrations/0024_upserttenant_upsertvaluekey_upsertscoped.py b/plain-postgres/tests/app/examples/migrations/0024_upserttenant_upsertvaluekey_upsertscoped.py new file mode 100644 index 0000000000..8ea4ec1377 --- /dev/null +++ b/plain-postgres/tests/app/examples/migrations/0024_upserttenant_upsertvaluekey_upsertscoped.py @@ -0,0 +1,43 @@ +# Generated by Plain 0.163.1 on 2026-09-19 19:26 + +from plain.postgres import migrations + +from plain import postgres + + +class Migration(migrations.Migration): + dependencies = (("examples", "0023_upsertpair"),) + + operations = ( + migrations.CreateModel( + name="UpsertTenant", + fields=[ + ("id", postgres.PrimaryKeyField()), + ("name", postgres.TextField(max_length=100)), + ], + ), + migrations.CreateModel( + name="UpsertValueKey", + fields=[ + ("id", postgres.PrimaryKeyField()), + ("blob", postgres.BinaryField()), + ("payload", postgres.JSONField()), + ("value", postgres.IntegerField(default=0)), + ("zone", postgres.TimeZoneField()), + ], + ), + migrations.CreateModel( + name="UpsertScoped", + fields=[ + ("id", postgres.PrimaryKeyField()), + ("slug", postgres.TextField(max_length=100)), + ("value", postgres.IntegerField(default=0)), + ( + "tenant", + postgres.ForeignKeyField( + on_delete=postgres.CASCADE, to="examples.upserttenant" + ), + ), + ], + ), + ) diff --git a/plain-postgres/tests/app/examples/migrations/0025_upsertfloatkey.py b/plain-postgres/tests/app/examples/migrations/0025_upsertfloatkey.py new file mode 100644 index 0000000000..f7193c9a41 --- /dev/null +++ b/plain-postgres/tests/app/examples/migrations/0025_upsertfloatkey.py @@ -0,0 +1,20 @@ +# Generated by Plain 0.163.1 on 2026-09-19 19:44 + +from plain.postgres import migrations + +from plain import postgres + + +class Migration(migrations.Migration): + dependencies = (("examples", "0024_upserttenant_upsertvaluekey_upsertscoped"),) + + operations = ( + migrations.CreateModel( + name="UpsertFloatKey", + fields=[ + ("id", postgres.PrimaryKeyField()), + ("score", postgres.FloatField()), + ("value", postgres.IntegerField(default=0)), + ], + ), + ) diff --git a/plain-postgres/tests/app/examples/migrations/0026_upsertdecimalkey.py b/plain-postgres/tests/app/examples/migrations/0026_upsertdecimalkey.py new file mode 100644 index 0000000000..8f1f5d09ef --- /dev/null +++ b/plain-postgres/tests/app/examples/migrations/0026_upsertdecimalkey.py @@ -0,0 +1,20 @@ +# Generated by Plain 0.163.1 on 2026-09-20 02:00 + +from plain.postgres import migrations + +from plain import postgres + + +class Migration(migrations.Migration): + dependencies = (("examples", "0025_upsertfloatkey"),) + + operations = ( + migrations.CreateModel( + name="UpsertDecimalKey", + fields=[ + ("id", postgres.PrimaryKeyField()), + ("amount", postgres.DecimalField(decimal_places=4, max_digits=12)), + ("value", postgres.IntegerField(default=0)), + ], + ), + ) diff --git a/plain-postgres/tests/app/examples/models/__init__.py b/plain-postgres/tests/app/examples/models/__init__.py index 1840a23f28..a79bafaba8 100644 --- a/plain-postgres/tests/app/examples/models/__init__.py +++ b/plain-postgres/tests/app/examples/models/__init__.py @@ -19,4 +19,5 @@ string_conditions, trees, unregistered, + upsert, ) diff --git a/plain-postgres/tests/app/examples/models/upsert.py b/plain-postgres/tests/app/examples/models/upsert.py new file mode 100644 index 0000000000..906152122c --- /dev/null +++ b/plain-postgres/tests/app/examples/models/upsert.py @@ -0,0 +1,118 @@ +"""Test fixtures for QuerySet.bulk_upsert().""" + +from __future__ import annotations + +from decimal import Decimal +from zoneinfo import ZoneInfo + +from plain.postgres import Field, types + +from plain import postgres + + +@postgres.register_model +class UpsertItem(postgres.Model): + key: Field[str] = types.TextField(max_length=100) + value: Field[int] = types.IntegerField(default=0) + label: Field[str] = types.TextField(default="", required=False) + + model_options = postgres.Options( + constraints=[ + postgres.UniqueConstraint(fields=["key"], name="upsertitem_key_unique"), + ] + ) + + +@postgres.register_model +class UpsertPair(postgres.Model): + """Composite conflict key -- unique on (bucket, slug), not on either alone.""" + + bucket: Field[str] = types.TextField(max_length=100) + slug: Field[str] = types.TextField(max_length=100) + value: Field[int] = types.IntegerField(default=0) + + model_options = postgres.Options( + constraints=[ + postgres.UniqueConstraint( + fields=["bucket", "slug"], name="upsertpair_bucket_slug_unique" + ), + ] + ) + + +@postgres.register_model +class UpsertTenant(postgres.Model): + name: Field[str] = types.TextField(max_length=100) + + +@postgres.register_model +class UpsertScoped(postgres.Model): + """A foreign key as part of the conflict key, and as an updated column.""" + + tenant: Field[UpsertTenant] = types.ForeignKeyField( + UpsertTenant, on_delete=postgres.CASCADE + ) + slug: Field[str] = types.TextField(max_length=100) + value: Field[int] = types.IntegerField(default=0) + + model_options = postgres.Options( + constraints=[ + postgres.UniqueConstraint( + fields=["tenant", "slug"], name="upsertscoped_tenant_slug_unique" + ), + ] + ) + + +@postgres.register_model +class UpsertValueKey(postgres.Model): + """A composite conflict key of column types whose Python values are + unhashable (jsonb dicts), unorderable (ZoneInfo), or neither (memoryview). + """ + + payload: Field[object] = types.JSONField() + blob: Field[bytes | memoryview] = types.BinaryField() + zone: Field[ZoneInfo] = types.TimeZoneField() + value: Field[int] = types.IntegerField(default=0) + + model_options = postgres.Options( + constraints=[ + postgres.UniqueConstraint( + fields=["payload", "blob", "zone"], + name="upsertvaluekey_payload_blob_zone_unique", + ), + ] + ) + + +@postgres.register_model +class UpsertFloatKey(postgres.Model): + """A float conflict key -- the column type that can hold NaN.""" + + score: Field[float] = types.FloatField() + value: Field[int] = types.IntegerField(default=0) + + model_options = postgres.Options( + constraints=[ + postgres.UniqueConstraint( + fields=["score"], name="upsertfloatkey_score_unique" + ), + ] + ) + + +@postgres.register_model +class UpsertDecimalKey(postgres.Model): + """A numeric conflict key -- the column type whose scale Python keeps and + Postgres does not.""" + + amount: Field[Decimal] = types.DecimalField(max_digits=12, decimal_places=4) + value: Field[int] = types.IntegerField(default=0) + + model_options = postgres.Options( + constraints=[ + postgres.UniqueConstraint( + fields=["amount"], name="upsertdecimalkey_amount_unique" + ), + ] + ) diff --git a/plain-postgres/tests/internal/test_conflict_lock_order.py b/plain-postgres/tests/internal/test_conflict_lock_order.py new file mode 100644 index 0000000000..007ee8dbce --- /dev/null +++ b/plain-postgres/tests/internal/test_conflict_lock_order.py @@ -0,0 +1,90 @@ +"""bulk_upsert() locks rows in conflict-key order, whatever shape they are. + +Two callers touching overlapping keys deadlock unless they take the locks in +the same order. The order is the sorted conflict key -- and it has to span +every object, not each statement, because a caller whose objects carry ids +sends them in a separate statement from one whose objects don't. Sorting +within each statement would let two callers holding the same four keys lock +them as (a,b),(c,d) and (c,d),(a,b). + +This pins the emitted order rather than racing two sessions: the order is the +whole guarantee, and asserting it directly can't flake. +""" + +from __future__ import annotations + +from app.examples.models.upsert import UpsertItem +from plain.postgres.query import QuerySet + + +def sent_key_runs(monkeypatch, items) -> list[list[str]]: + """The keys bulk_upsert() sends, grouped by statement.""" + runs: list[list[str]] = [] + original = QuerySet._batched_insert + + def recording(self, objs, fields, batch_size, **kwargs): + runs.append([obj.key for obj in objs]) + return original(self, objs, fields, batch_size, **kwargs) + + monkeypatch.setattr(QuerySet, "_batched_insert", recording) + UpsertItem.query.bulk_upsert( + items, update_fields=[UpsertItem.value], unique_fields=[UpsertItem.key] + ) + return runs + + +def test_lock_order_is_the_sorted_key_when_no_object_carries_an_id(db, monkeypatch): + items = [UpsertItem(key=key, value=1) for key in ("d", "b", "c", "a")] + runs = sent_key_runs(monkeypatch, items) + + assert runs == [["a", "b", "c", "d"]] + + +def test_lock_order_is_the_same_when_some_objects_carry_ids(db, monkeypatch): + # The reviewer's scenario: the same four keys, but c and d arrive with ids, + # so they need their own statement. The keys still go out in sorted order. + items = [] + for key in ("d", "b", "c", "a"): + item = UpsertItem(key=key, value=1) + if key in ("c", "d"): + item.id = {"c": 101, "d": 102}[key] + items.append(item) + runs = sent_key_runs(monkeypatch, items) + + assert runs == [["a", "b"], ["c", "d"]] + assert [key for run in runs for key in run] == ["a", "b", "c", "d"] + + +def test_ids_interleaved_through_the_key_order_cost_a_statement_each(db, monkeypatch): + # A new statement starts only where the shape changes, so alternating ids + # is the worst case -- and the key order still holds across all of them. + items = [] + for index, key in enumerate(("a", "b", "c", "d")): + item = UpsertItem(key=key, value=1) + if index % 2 == 0: + item.id = 100 + index + items.append(item) + runs = sent_key_runs(monkeypatch, items) + + assert runs == [["a"], ["b"], ["c"], ["d"]] + + +def test_batch_size_splits_within_a_run_without_disturbing_the_order(db, monkeypatch): + items = [UpsertItem(key=key, value=1) for key in ("d", "b", "c", "a")] + runs: list[list[str]] = [] + original = QuerySet._batched_insert + + def recording(self, objs, fields, batch_size, **kwargs): + runs.append([obj.key for obj in objs]) + return original(self, objs, fields, batch_size, **kwargs) + + monkeypatch.setattr(QuerySet, "_batched_insert", recording) + UpsertItem.query.bulk_upsert( + items, + update_fields=[UpsertItem.value], + unique_fields=[UpsertItem.key], + batch_size=2, + ) + + # batch_size caps the statement inside a run; the run is still one call. + assert runs == [["a", "b", "c", "d"]] diff --git a/plain-postgres/tests/internal/test_conflict_sort_value.py b/plain-postgres/tests/internal/test_conflict_sort_value.py new file mode 100644 index 0000000000..7e08c3b0c1 --- /dev/null +++ b/plain-postgres/tests/internal/test_conflict_sort_value.py @@ -0,0 +1,134 @@ +"""The order bulk_upsert() sends its batches in. + +Concurrent callers avoid deadlocking each other by locking rows in the same +order, which works only if two callers holding the same logical key render it +the same way -- however each of them happened to spell it. And since comparing +the results is the sort, comparing them must never raise. + +Unit-level because `conflict_sort_value` isn't public API; the behavior it buys +is covered end-to-end in tests/public/test_bulk_upsert.py. +""" + +from __future__ import annotations + +import datetime +from decimal import Decimal +from enum import Enum, StrEnum +from zoneinfo import ZoneInfo + +import pytest +from app.examples.models.forms import FormsExample +from app.examples.models.upsert import UpsertItem, UpsertValueKey +from plain.exceptions import ValidationError +from plain.postgres.query import conflict_sort_value + +AMOUNT = FormsExample.amount +RATIO = FormsExample.ratio +MOMENT = FormsExample.event_datetime +KEY = UpsertItem.key + + +def sort_value(field, value): + """Prepare the value the way bulk_upsert() does, then render it.""" + return conflict_sort_value(field, field.get_prep_value(value)) + + +class Shade(StrEnum): + RED = "red" + + +class LegacyShade(str, Enum): + RED = "red" + + +class Shouty(str): + def __str__(self) -> str: # a SafeString-style override + return "SHOUTY" + + +def test_str_subclasses_render_as_the_string_the_column_holds(): + # All of these compare equal to "red" and store as "red", so they have to + # sort as "red" -- str() alone gives "LegacyShade.RED" and "SHOUTY". + assert Shade.RED == LegacyShade.RED == Shouty("red") == "red" + for value in (Shade.RED, LegacyShade.RED, Shouty("red"), "red"): + assert sort_value(KEY, value) == "red" + + +def test_bytes_and_memoryview_of_the_same_content_render_the_same(): + # str(memoryview) is an address, which differs run to run and between two + # callers holding the same bytes. + assert sort_value(UpsertValueKey.blob, memoryview(b"x")) == sort_value( + UpsertValueKey.blob, b"x" + ) + assert sort_value(UpsertValueKey.blob, bytearray(b"x")) == sort_value( + UpsertValueKey.blob, b"x" + ) + assert "memory at" not in sort_value(UpsertValueKey.blob, memoryview(b"x")) + + +def test_one_instant_written_at_two_offsets_renders_once(): + # timestamptz stores the instant, so these are one row. + utc = datetime.datetime(2024, 1, 1, 12, tzinfo=datetime.UTC) + new_york = datetime.datetime(2024, 1, 1, 7, tzinfo=ZoneInfo("America/New_York")) + assert utc == new_york + assert sort_value(MOMENT, utc) == sort_value(MOMENT, new_york) + + +def test_signed_zeros_render_the_same(): + assert -0.0 == 0.0 + assert sort_value(RATIO, -0.0) == sort_value(RATIO, 0.0) + assert sort_value(AMOUNT, Decimal("-0.00")) == sort_value(AMOUNT, Decimal("0.0")) + + +def test_equal_decimals_render_the_same_whatever_their_scale(): + assert Decimal("1.0") == Decimal("1.00") + assert sort_value(AMOUNT, Decimal("1.0")) == sort_value(AMOUNT, Decimal("1.00")) + assert sort_value(AMOUNT, Decimal(100)) == sort_value(AMOUNT, Decimal("1E+2")) + + +def test_different_values_still_render_apart(): + assert sort_value(AMOUNT, Decimal("1.0")) != sort_value(AMOUNT, Decimal("1.5")) + assert sort_value(KEY, "a") != sort_value(KEY, "b") + + +def test_a_decimal_too_large_to_normalize_does_not_raise(): + # normalize() overflows on this; Postgres rejects it on write, and that is + # the error worth surfacing, so the sort leaves it alone. + assert isinstance(sort_value(AMOUNT, Decimal("1E+999999999")), str) + + +def test_a_non_finite_decimal_is_rejected_before_the_sort_key(): + # DecimalField.to_python refuses NaN and infinity, so they never reach the + # sort key. Finding that out before any query is issued is the point. + for value in (Decimal("NaN"), Decimal("sNaN"), Decimal("Infinity")): + with pytest.raises(ValidationError): + sort_value(AMOUNT, value) + + +def test_values_with_no_ordering_of_their_own_render_and_sort(): + keys = [ + sort_value(RATIO, float("nan")), + sort_value(UpsertValueKey.zone, ZoneInfo("America/Chicago")), + sort_value(UpsertValueKey.blob, memoryview(b"x")), + ] + assert all(isinstance(key, str) for key in keys) + assert len(sorted(keys)) == 3 + + +def test_json_objects_render_the_same_whatever_order_their_keys_were_written_in(): + payload = UpsertValueKey.payload + assert sort_value(payload, {"a": 1, "b": 2}) == sort_value( + payload, {"b": 2, "a": 1} + ) + # A non-string object key is valid -- the encoder stringifies it -- and + # sorting the keys before that encode would compare an int to a str. + assert isinstance(sort_value(payload, {1: "a", "2": "b"}), str) + + +def test_a_polymorphic_json_column_sorts_without_comparing_across_types(): + # A jsonb column can hold an object in one row and a number in the next. + payload = UpsertValueKey.payload + keys = [sort_value(payload, value) for value in ({"a": 1}, 7, "s", [1, 2])] + assert all(isinstance(key, str) for key in keys) + assert len(set(keys)) == 4 + assert len(sorted(keys)) == 4 diff --git a/plain-postgres/tests/internal/test_returning_order.py b/plain-postgres/tests/internal/test_returning_order.py new file mode 100644 index 0000000000..317aa06448 --- /dev/null +++ b/plain-postgres/tests/internal/test_returning_order.py @@ -0,0 +1,54 @@ +"""RETURNING comes back in VALUES order, including under ON CONFLICT. + +`bulk_create()` and `bulk_upsert()` both map returned rows onto the objects +they were given by position. That holds only because Postgres processes a +multi-row INSERT sequentially and emits one RETURNING row per VALUES row, in +order -- on the DO UPDATE path as much as the plain insert path. + +Nothing in the SQL standard promises this, so it is pinned here rather than +assumed. If this test ever fails, both call sites are wrong together. +""" + +from __future__ import annotations + +import random + +from app.examples.models.upsert import UpsertItem +from plain.postgres.db import get_connection + +ROWS = 400 + + +def _returned_keys(keys: list[str]) -> list[str]: + """Insert `keys` in one ON CONFLICT statement, return the RETURNING order.""" + table = UpsertItem.model_options.db_table + # Only the placeholder count is interpolated; every value is a parameter. + placeholders = ", ".join(["(%s, %s)"] * len(keys)) + sql = ( + f"INSERT INTO {table} (key, value) VALUES {placeholders} " + "ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value " + "RETURNING key" + ) + params = [param for key in keys for param in (key, 1)] + with get_connection().cursor() as cursor: + cursor.execute(sql, params) + return [row[0] for row in cursor.fetchall()] + + +def test_returning_order_matches_values_order_under_on_conflict(db): + # Seed every other key, so the batch is a shuffled mix of inserts and + # conflicting updates -- the case where the order could plausibly diverge. + for index in range(0, ROWS, 2): + UpsertItem(key=f"k{index}", value=0).create() + + keys = [f"k{index}" for index in range(ROWS)] + random.Random(0).shuffle(keys) + + assert _returned_keys(keys) == keys + + +def test_returning_order_matches_values_order_for_plain_inserts(db): + keys = [f"k{index}" for index in range(ROWS)] + random.Random(1).shuffle(keys) + + assert _returned_keys(keys) == keys diff --git a/plain-postgres/tests/public/test_bulk_upsert.py b/plain-postgres/tests/public/test_bulk_upsert.py new file mode 100644 index 0000000000..3b32594a41 --- /dev/null +++ b/plain-postgres/tests/public/test_bulk_upsert.py @@ -0,0 +1,726 @@ +"""QuerySet.bulk_upsert() inserts new rows and updates conflicting ones. + +One INSERT ... ON CONFLICT (unique_fields) DO UPDATE ... RETURNING per batch. +Every returned object -- inserted or updated -- comes back with its DB-returned +fields (primary key, DB defaults) populated from the row at its own position. +""" + +from __future__ import annotations + +import random +from decimal import Decimal +from zoneinfo import ZoneInfo + +import psycopg +import pytest +from app.examples.models.defaults import DBDefaultsExample +from app.examples.models.indexes import IndexExample +from app.examples.models.mixins import MixinTestModel +from app.examples.models.returning import ReturningEvent +from app.examples.models.upsert import ( + UpsertDecimalKey, + UpsertFloatKey, + UpsertItem, + UpsertPair, + UpsertScoped, + UpsertTenant, + UpsertValueKey, +) +from plain.postgres.exceptions import FieldError + + +def test_bulk_upsert_inserts_new_rows_and_sets_pks(db): + items = [ + UpsertItem(key="a", value=1), + UpsertItem(key="b", value=2), + ] + returned = UpsertItem.query.bulk_upsert( + items, update_fields=[UpsertItem.value], unique_fields=[UpsertItem.key] + ) + + assert [r.id for r in returned] == [item.id for item in items] + assert all(item.id is not None for item in items) + stored = {row.key: row.value for row in UpsertItem.query.all()} + assert stored == {"a": 1, "b": 2} + + +def test_bulk_upsert_mixed_batch_inserts_and_updates(db): + UpsertItem(key="a", value=1).create() + existing_id = UpsertItem.query.get(key="a").id + + items = [ + UpsertItem(key="a", value=10), # conflicts -> update + UpsertItem(key="b", value=20), # new -> insert + ] + UpsertItem.query.bulk_upsert( + items, update_fields=[UpsertItem.value], unique_fields=[UpsertItem.key] + ) + + by_key = {item.key: item for item in items} + # The updated row keeps its original primary key. + assert by_key["a"].id == existing_id + assert by_key["b"].id is not None + assert by_key["b"].id != existing_id + + stored = {row.key: row.value for row in UpsertItem.query.all()} + assert stored == {"a": 10, "b": 20} + + +def test_bulk_upsert_updates_only_named_fields(db): + UpsertItem(key="a", value=1, label="original").create() + + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=99, label="ignored")], + update_fields=[UpsertItem.value], + unique_fields=[UpsertItem.key], + ) + + row = UpsertItem.query.get(key="a") + assert row.value == 99 # named field updated + assert row.label == "original" # field not in update_fields preserved + + +def test_bulk_upsert_hydrates_every_object_when_all_of_them_conflict(db): + # Seed so every input row takes the DO UPDATE path, and pass them out of + # key order so the deadlock sort actually reorders the batch. + for key in ("a", "b", "c"): + UpsertItem(key=key, value=0).create() + seeded_ids = {row.key: row.id for row in UpsertItem.query.all()} + + items = [ + UpsertItem(key="c", value=3), + UpsertItem(key="a", value=1), + UpsertItem(key="b", value=2), + ] + returned = UpsertItem.query.bulk_upsert( + items, update_fields=[UpsertItem.value], unique_fields=[UpsertItem.key] + ) + + # Batches are issued in conflict-key order, but the caller gets its own + # order back. + assert [r.key for r in returned] == ["c", "a", "b"] + + for item in items: + assert item.id == seeded_ids[item.key] + + stored = {row.key: row.value for row in UpsertItem.query.all()} + assert stored == {"a": 1, "b": 2, "c": 3} + + +def test_bulk_upsert_empty_returns_empty(db): + assert ( + UpsertItem.query.bulk_upsert( + [], update_fields=[UpsertItem.value], unique_fields=[UpsertItem.key] + ) + == [] + ) + + +def test_bulk_upsert_batches(db): + items = [UpsertItem(key=f"k{i}", value=i) for i in range(5)] + UpsertItem.query.bulk_upsert( + items, + update_fields=[UpsertItem.value], + unique_fields=[UpsertItem.key], + batch_size=2, + ) + + assert all(item.id is not None for item in items) + assert UpsertItem.query.count() == 5 + + +def test_bulk_upsert_unique_fields_must_match_a_constraint(db): + with pytest.raises(ValueError, match="must name the primary key"): + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=1)], + update_fields=[UpsertItem.value], + unique_fields=[UpsertItem.value], # no unique constraint on value + ) + + +def test_bulk_upsert_null_unique_value_rejected(db): + with pytest.raises(ValueError, match="non-null key"): + UpsertItem.query.bulk_upsert( + # A null unique value is a type error the checker catches; the + # runtime guard is what protects callers that aren't type-checked. + [UpsertItem(key=None, value=1)], # ty: ignore[invalid-argument-type] + update_fields=[UpsertItem.value], + unique_fields=[UpsertItem.key], + ) + + +def test_bulk_upsert_update_fields_cannot_overlap_unique_fields(db): + with pytest.raises(ValueError, match="cannot overlap unique_fields"): + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=1)], + update_fields=[UpsertItem.key], + unique_fields=[UpsertItem.key], + ) + + +def test_bulk_upsert_requires_update_fields(db): + with pytest.raises(ValueError, match="requires update_fields"): + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=1)], + update_fields=[], + unique_fields=[UpsertItem.key], + ) + + +def test_bulk_upsert_string_field_rejected(db): + with pytest.raises(TypeError, match="takes field references, not strings"): + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=1)], + update_fields=["value"], # ty: ignore[invalid-argument-type] + unique_fields=[UpsertItem.key], + ) + + +def test_bulk_upsert_wrong_model_field_rejected(db): + with pytest.raises(FieldError, match="belongs to a different model"): + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=1)], + update_fields=[UpsertItem.value], + unique_fields=[ReturningEvent.label], + ) + + +def test_bulk_create_no_longer_accepts_update_conflicts(db): + # bulk_create is insert-only now; the conflict surface moved to bulk_upsert. + removed_conflict_kwargs: dict[str, object] = { + "update_conflicts": True, + "update_fields": ["value"], + "unique_fields": ["key"], + } + with pytest.raises(TypeError): + UpsertItem.query.bulk_create( + [UpsertItem(key="a", value=1)], + **removed_conflict_kwargs, # ty: ignore[invalid-argument-type] + ) + + +def test_bulk_upsert_composite_unique_fields(db): + UpsertPair(bucket="b1", slug="s1", value=1).create() + seeded_id = UpsertPair.query.get(bucket="b1", slug="s1").id + + items = [ + UpsertPair(bucket="b1", slug="s1", value=10), # conflicts -> update + UpsertPair(bucket="b1", slug="s2", value=20), # same bucket, new slug + UpsertPair(bucket="b2", slug="s1", value=30), # same slug, new bucket + ] + UpsertPair.query.bulk_upsert( + items, + update_fields=[UpsertPair.value], + unique_fields=[UpsertPair.bucket, UpsertPair.slug], + ) + + # The conflicting row is matched back by the whole composite key. + assert items[0].id == seeded_id + assert len({item.id for item in items}) == 3 + + stored = {(row.bucket, row.slug): row.value for row in UpsertPair.query.all()} + assert stored == {("b1", "s1"): 10, ("b1", "s2"): 20, ("b2", "s1"): 30} + + +def test_bulk_upsert_duplicate_keys_in_one_batch_rejected(db): + # Postgres refuses to touch a row twice in one statement. The raw + # CardinalityViolation is re-raised as something that names the problem. + # It aborts the surrounding transaction, so nothing is queried after it. + with pytest.raises(ValueError, match=r"same \['key'\] in one statement"): + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=1), UpsertItem(key="a", value=2)], + update_fields=[UpsertItem.value], + unique_fields=[UpsertItem.key], + ) + + +def test_bulk_upsert_duplicate_keys_in_separate_batches_are_legal(db): + # Split across statements there is no cardinality violation: the first + # inserts the row and the second updates it. + items = [UpsertItem(key="a", value=1), UpsertItem(key="a", value=2)] + UpsertItem.query.bulk_upsert( + items, + update_fields=[UpsertItem.value], + unique_fields=[UpsertItem.key], + batch_size=1, + ) + + assert UpsertItem.query.count() == 1 + row = UpsertItem.query.get(key="a") + # Both objects are hydrated, from their own returned row -- the same row. + assert [item.id for item in items] == [row.id, row.id] + # Equal keys keep their input order, so the later write is the one that lands. + assert row.value == 2 + + +def test_bulk_upsert_database_generated_unique_field_rejected(db): + # db_uuid is generate=True, so the objects hold a DatabaseDefault sentinel + # rather than a value to conflict on. + with pytest.raises(ValueError, match="the database generates its value"): + DBDefaultsExample.query.bulk_upsert( + [DBDefaultsExample(name="a")], + update_fields=[DBDefaultsExample.name], + unique_fields=[DBDefaultsExample.db_uuid], + ) + + +def test_bulk_upsert_validates_arguments_even_when_empty(db): + # An empty objs list is still a bad call if the fields are wrong. + with pytest.raises(TypeError, match="takes field references, not strings"): + UpsertItem.query.bulk_upsert( + [], + update_fields=["value"], # ty: ignore[invalid-argument-type] + unique_fields=[UpsertItem.key], + ) + + +def test_bulk_upsert_refreshes_update_now_columns_without_naming_them(db): + MixinTestModel(name="a").create() + seeded = MixinTestModel.query.get(name="a") + + renamed = MixinTestModel(name="b") + renamed.id = seeded.id + MixinTestModel.query.bulk_upsert( + [renamed], + update_fields=[MixinTestModel.name], # updated_at deliberately absent + unique_fields=[MixinTestModel.id], + ) + + row = MixinTestModel.query.get(id=seeded.id) + assert row.name == "b" + # The row was updated, so its update_now column was too. + assert row.updated_at > seeded.updated_at + # And the object handed back agrees with the row it was hydrated from -- + # pre_save stamps the object, EXCLUDED carries that same stamp to the row. + assert renamed.updated_at == row.updated_at + + +def test_bulk_upsert_cannot_update_a_database_generated_column(db): + # created_at is create_now-only: EXCLUDED would carry a fresh now() and + # reset the creation timestamp on every update. + with pytest.raises(ValueError, match="the database generates its value"): + MixinTestModel.query.bulk_upsert( + [MixinTestModel(name="a")], + update_fields=[MixinTestModel.created_at], + unique_fields=[MixinTestModel.id], + ) + + +def test_bulk_upsert_update_now_unique_field_rejected(db): + with pytest.raises(ValueError, match="stamped again on every write"): + MixinTestModel.query.bulk_upsert( + [MixinTestModel(name="a")], + update_fields=[MixinTestModel.name], + unique_fields=[MixinTestModel.updated_at], + ) + + +def test_bulk_upsert_foreign_key_in_unique_fields(db): + # Model.fk is a relation descriptor, not a Field, but it is the only way to + # name the foreign key column -- so the write-API field lists accept it. + tenant = UpsertTenant(name="t1") + tenant.create() + tenant = UpsertTenant.query.get(name="t1") + UpsertScoped(tenant=tenant, slug="a", value=1).create() + seeded_id = UpsertScoped.query.get(slug="a").id + + items = [ + UpsertScoped(tenant=tenant, slug="a", value=10), # conflicts -> update + UpsertScoped(tenant=tenant, slug="b", value=20), # new -> insert + ] + UpsertScoped.query.bulk_upsert( + items, + update_fields=[UpsertScoped.value], + unique_fields=[UpsertScoped.tenant, UpsertScoped.slug], + ) + + assert items[0].id == seeded_id + assert items[1].id != seeded_id + stored = {row.slug: row.value for row in UpsertScoped.query.all()} + assert stored == {"a": 10, "b": 20} + + +def test_bulk_upsert_foreign_key_in_update_fields(db): + first = UpsertTenant(name="t1") + first.create() + second = UpsertTenant(name="t2") + second.create() + first = UpsertTenant.query.get(name="t1") + second = UpsertTenant.query.get(name="t2") + + UpsertScoped(tenant=first, slug="a", value=1).create() + seeded_id = UpsertScoped.query.get(slug="a").id + + moved = UpsertScoped(tenant=second, slug="a", value=2) + moved.id = seeded_id + UpsertScoped.query.bulk_upsert( + [moved], + update_fields=[UpsertScoped.tenant, UpsertScoped.value], + unique_fields=[UpsertScoped.id], + ) + + row = UpsertScoped.query.get(id=seeded_id) + assert row.tenant.id == second.id + assert row.value == 2 + + +def test_bulk_upsert_keys_that_python_cannot_sort(db): + # ZoneInfo and memoryview have no ordering at all and a jsonb dict has no + # useful one -- the batch still has to be put in a deterministic order. + chicago = ZoneInfo("America/Chicago") + utc = ZoneInfo("UTC") + UpsertValueKey(payload={"a": 1}, blob=b"x", zone=utc, value=1).create() + seeded_id = UpsertValueKey.query.get(value=1).id + + items = [ + # Conflicts: same key, written with its dict keys in the other order. + UpsertValueKey(payload={"a": 1}, blob=b"x", zone=utc, value=10), + UpsertValueKey(payload={"b": 2, "a": 1}, blob=b"y", zone=chicago, value=20), + UpsertValueKey(payload={"a": 1, "b": 2}, blob=b"z", zone=chicago, value=30), + ] + UpsertValueKey.query.bulk_upsert( + items, + update_fields=[UpsertValueKey.value], + unique_fields=[ + UpsertValueKey.payload, + UpsertValueKey.blob, + UpsertValueKey.zone, + ], + ) + + assert items[0].id == seeded_id + assert len({item.id for item in items}) == 3 + stored = {row.value for row in UpsertValueKey.query.all()} + assert stored == {10, 20, 30} + + +def test_bulk_upsert_key_values_that_do_not_compare_to_each_other(db): + # A jsonb conflict key can hold an object in one row and a number in the + # next. Those don't compare, and the batches still have to be ordered. + utc = ZoneInfo("UTC") + items = [ + UpsertValueKey(payload={"a": 1}, blob=b"x", zone=utc, value=1), + UpsertValueKey(payload=7, blob=b"x", zone=utc, value=2), + UpsertValueKey(payload="s", blob=b"x", zone=utc, value=3), + UpsertValueKey(payload=[1, 2], blob=b"x", zone=utc, value=4), + ] + UpsertValueKey.query.bulk_upsert( + items, + update_fields=[UpsertValueKey.value], + unique_fields=[ + UpsertValueKey.payload, + UpsertValueKey.blob, + UpsertValueKey.zone, + ], + ) + + assert len({item.id for item in items}) == 4 + assert {row.value for row in UpsertValueKey.query.all()} == {1, 2, 3, 4} + + +def test_bulk_upsert_json_key_with_non_string_object_keys(db): + # jsonb object keys are always strings -- the encoder stringifies an int + # key on the way in. The sort key has to be canonicalized the same way, or + # sorting the object's keys compares an int against a str and raises. + utc = ZoneInfo("UTC") + UpsertValueKey( + payload={1: "a", "2": "b", "nested": {"z": 1, "y": 2}}, + blob=b"x", + zone=utc, + value=1, + ).create() + seeded_id = UpsertValueKey.query.get(value=1).id + + # The same logical key, written with its keys in another order and the + # int key spelled as a string. + conflicting = UpsertValueKey( + payload={"nested": {"y": 2, "z": 1}, "2": "b", "1": "a"}, + blob=b"x", + zone=utc, + value=99, + ) + UpsertValueKey.query.bulk_upsert( + [conflicting], + update_fields=[UpsertValueKey.value], + unique_fields=[ + UpsertValueKey.payload, + UpsertValueKey.blob, + UpsertValueKey.zone, + ], + ) + + assert conflicting.id == seeded_id + assert UpsertValueKey.query.count() == 1 + assert UpsertValueKey.query.get(id=seeded_id).value == 99 + + +def test_bulk_upsert_nan_key_round_trips(db): + # Postgres holds NaN equal to NaN for uniqueness, so a NaN key really does + # conflict -- and sorting it must not raise the way `<` on a NaN would. + items = [UpsertFloatKey(score=float("nan"), value=1)] + UpsertFloatKey.query.bulk_upsert( + items, + update_fields=[UpsertFloatKey.value], + unique_fields=[UpsertFloatKey.score], + ) + seeded_id = items[0].id + assert seeded_id is not None + + updated = [UpsertFloatKey(score=float("nan"), value=2)] + UpsertFloatKey.query.bulk_upsert( + updated, + update_fields=[UpsertFloatKey.value], + unique_fields=[UpsertFloatKey.score], + ) + + assert updated[0].id == seeded_id + assert UpsertFloatKey.query.count() == 1 + assert UpsertFloatKey.query.get(id=seeded_id).value == 2 + + +def test_bulk_upsert_duplicate_nan_keys_in_one_batch_rejected(db): + # Postgres holds the two NaNs equal, so this is the same row twice. + with pytest.raises(ValueError, match=r"same \['score'\] in one statement"): + UpsertFloatKey.query.bulk_upsert( + [ + UpsertFloatKey(score=float("nan"), value=1), + UpsertFloatKey(score=float("nan"), value=2), + ], + update_fields=[UpsertFloatKey.value], + unique_fields=[UpsertFloatKey.score], + ) + + +def test_bulk_upsert_accepts_any_sequence_of_field_references(db): + # The parameters are Sequence, so a tuple -- or a conflict target hoisted + # into a variable, which a list parameter would reject as invariant -- + # works as well as an inline list. + conflict_target = (UpsertPair.bucket, UpsertPair.slug) + items = [UpsertPair(bucket="b", slug="s", value=1)] + UpsertPair.query.bulk_upsert( + items, + update_fields=(UpsertPair.value,), + unique_fields=conflict_target, + ) + + assert items[0].id is not None + assert UpsertPair.query.get(bucket="b", slug="s").value == 1 + + +def test_bulk_upsert_honors_an_id_the_caller_set(db): + # An explicitly set id is the caller's choice, not ours to discard for a + # generated one -- bulk_create() honors it, and so does this. + item = UpsertItem(key="a", value=1) + item.id = 5 + UpsertItem.query.bulk_upsert( + [item], update_fields=[UpsertItem.value], unique_fields=[UpsertItem.key] + ) + + assert item.id == 5 + assert UpsertItem.query.get(key="a").id == 5 + + +def test_bulk_upsert_mixes_objects_with_and_without_ids(db): + # The two go out as separate statements; both still come back hydrated + # from their own row. + with_id = UpsertItem(key="a", value=1) + with_id.id = 5 + without_id = UpsertItem(key="b", value=2) + UpsertItem.query.bulk_upsert( + [with_id, without_id], + update_fields=[UpsertItem.value], + unique_fields=[UpsertItem.key], + ) + + assert with_id.id == 5 + assert without_id.id is not None + assert without_id.id != 5 + assert {row.key: row.id for row in UpsertItem.query.all()} == { + "a": 5, + "b": without_id.id, + } + + +def test_bulk_upsert_across_several_batches_out_of_key_order(db): + # Every other key already exists, the input is shuffled, and the batches + # are smaller than the input -- so inserts and updates interleave across + # statements and the sort reorders them. Each object still has to come + # back hydrated from its own row. + for index in range(0, 8, 2): + UpsertItem(key=f"k{index}", value=0).create() + seeded = {row.key: row.id for row in UpsertItem.query.all()} + + keys = [f"k{index}" for index in range(8)] + random.Random(0).shuffle(keys) + items = [UpsertItem(key=key, value=int(key[1:])) for key in keys] + + UpsertItem.query.bulk_upsert( + items, + update_fields=[UpsertItem.value], + unique_fields=[UpsertItem.key], + batch_size=2, + ) + + assert [item.key for item in items] == keys # caller's order preserved + for item in items: + if item.key in seeded: + assert item.id == seeded[item.key] + assert item.id is not None + + assert UpsertItem.query.count() == 8 + assert {row.key: row.value for row in UpsertItem.query.all()} == { + f"k{index}": index for index in range(8) + } + + +def test_bulk_upsert_repeated_update_field_rejected(db): + # Postgres assigns each column once per statement; naming one twice is a + # syntax error mid-transaction, so it's caught at the call instead. + with pytest.raises(ValueError, match=r"names \['value'\] more than once"): + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=1)], + update_fields=[UpsertItem.value, UpsertItem.value], + unique_fields=[UpsertItem.key], + ) + + +def test_bulk_upsert_repeated_update_field_rejected_before_any_query(db): + # Validation happens at the call, so an empty objs list is checked too. + with pytest.raises(ValueError, match="more than once"): + UpsertItem.query.bulk_upsert( + [], + update_fields=[UpsertItem.value, UpsertItem.value], + unique_fields=[UpsertItem.key], + ) + + +def test_bulk_upsert_unsaved_related_object_names_bulk_upsert(db): + # The guard is shared with bulk_create(); the message has to name the call + # the user actually made. + with pytest.raises(ValueError, match=r"^bulk_upsert\(\) prohibited"): + UpsertScoped.query.bulk_upsert( + [UpsertScoped(tenant=UpsertTenant(name="unsaved"), slug="s", value=1)], + update_fields=[UpsertScoped.value], + unique_fields=[UpsertScoped.tenant, UpsertScoped.slug], + ) + + +def test_bulk_upsert_decimal_key_conflicts_across_scales(db): + # numeric(12,4) stores 1.0 and 1.00 as the same value, so the second call + # has to find the first one's row rather than insert beside it. + first = [UpsertDecimalKey(amount=Decimal("1.0"), value=1)] + UpsertDecimalKey.query.bulk_upsert( + first, + update_fields=[UpsertDecimalKey.value], + unique_fields=[UpsertDecimalKey.amount], + ) + + second = [UpsertDecimalKey(amount=Decimal("1.00"), value=2)] + UpsertDecimalKey.query.bulk_upsert( + second, + update_fields=[UpsertDecimalKey.value], + unique_fields=[UpsertDecimalKey.amount], + ) + + assert second[0].id == first[0].id + assert UpsertDecimalKey.query.count() == 1 + assert UpsertDecimalKey.query.get(id=first[0].id).value == 2 + + +def test_bulk_upsert_decimal_key_conflicts_across_signed_zero(db): + zero = [UpsertDecimalKey(amount=Decimal("0.0"), value=1)] + UpsertDecimalKey.query.bulk_upsert( + zero, + update_fields=[UpsertDecimalKey.value], + unique_fields=[UpsertDecimalKey.amount], + ) + + negative_zero = [UpsertDecimalKey(amount=Decimal("-0.00"), value=2)] + UpsertDecimalKey.query.bulk_upsert( + negative_zero, + update_fields=[UpsertDecimalKey.value], + unique_fields=[UpsertDecimalKey.amount], + ) + + assert negative_zero[0].id == zero[0].id + assert UpsertDecimalKey.query.count() == 1 + + +def test_bulk_upsert_decimal_keys_at_different_scales_in_one_batch(db): + # Postgres holds these equal, so they are the same row twice in one + # statement -- the sort renders them identically, and the cardinality + # violation is reported as the duplicate it is. + with pytest.raises(ValueError, match=r"same \['amount'\] in one statement"): + UpsertDecimalKey.query.bulk_upsert( + [ + UpsertDecimalKey(amount=Decimal("1.0"), value=1), + UpsertDecimalKey(amount=Decimal("1.000"), value=2), + ], + update_fields=[UpsertDecimalKey.value], + unique_fields=[UpsertDecimalKey.amount], + ) + + +def test_bulk_upsert_conflict_hydrates_the_stored_id_over_the_caller_s(db): + # On the insert path a caller-set id is kept. On the conflict path the + # stored row is the truth: its id is what comes back, not the one passed. + UpsertItem(key="a", value=1).create() + stored_id = UpsertItem.query.get(key="a").id + + item = UpsertItem(key="a", value=2) + item.id = 99 + UpsertItem.query.bulk_upsert( + [item], update_fields=[UpsertItem.value], unique_fields=[UpsertItem.key] + ) + + assert item.id == stored_id + assert UpsertItem.query.count() == 1 + assert UpsertItem.query.get(key="a").value == 2 + + +def test_bulk_upsert_id_colliding_with_another_row_raises(db): + # The conflict target is `key`, so an id that collides with a different + # row is an ordinary primary key violation -- raised raw, like any + # set-based write. + UpsertItem(key="taken", value=1).create() + taken_id = UpsertItem.query.get(key="taken").id + + item = UpsertItem(key="new", value=1) + item.id = taken_id + with pytest.raises(psycopg.errors.UniqueViolation): + UpsertItem.query.bulk_upsert( + [item], update_fields=[UpsertItem.value], unique_fields=[UpsertItem.key] + ) + + +def test_bulk_upsert_unique_index_is_not_a_conflict_target(db): + # A unique Index would work as a Postgres arbiter, but bulk_upsert asks + # for a declared UniqueConstraint so the target is explicit in the model. + with pytest.raises(ValueError, match="must name the primary key"): + IndexExample.query.bulk_upsert( + [IndexExample(name="n", description="d")], + update_fields=[IndexExample.description], + unique_fields=[IndexExample.name], + ) + + +def test_bulk_upsert_cannot_update_the_primary_key(db): + with pytest.raises(ValueError, match="cannot update primary key fields"): + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=1)], + update_fields=[UpsertItem.id], + unique_fields=[UpsertItem.key], + ) + + +def test_bulk_upsert_ignores_queryset_filters(db): + # Like bulk_create, the write is against the table -- a filter on the + # queryset it is called from does not narrow or exclude anything. + UpsertItem(key="a", value=1).create() + + items = [UpsertItem(key="a", value=2), UpsertItem(key="b", value=3)] + UpsertItem.query.filter(key="nothing-matches-this").bulk_upsert( + items, update_fields=[UpsertItem.value], unique_fields=[UpsertItem.key] + ) + + assert {row.key: row.value for row in UpsertItem.query.all()} == {"a": 2, "b": 3} diff --git a/plain-postgres/tests/public/test_returning.py b/plain-postgres/tests/public/test_returning.py index 87cd494012..f98106665f 100644 --- a/plain-postgres/tests/public/test_returning.py +++ b/plain-postgres/tests/public/test_returning.py @@ -290,11 +290,23 @@ def test_returning_instances_carry_foreign_keys(db): [ lambda qs: qs.create(label="x", count=1), lambda qs: qs.bulk_create([ReturningEvent(label="x", count=1)]), + lambda qs: qs.bulk_upsert( + [ReturningEvent(label="x", count=1)], + update_fields=[ReturningEvent.count], + unique_fields=[ReturningEvent.id], + ), lambda qs: qs.bulk_update(list(ReturningEvent.query), ["count"]), lambda qs: qs.get_or_create(label="x", count=1), lambda qs: qs.update_or_create(label="x", defaults={"count": 1}), ], - ids=["create", "bulk_create", "bulk_update", "get_or_create", "update_or_create"], + ids=[ + "create", + "bulk_create", + "bulk_upsert", + "bulk_update", + "get_or_create", + "update_or_create", + ], ) def test_returning_rejects_other_writes(db, write): ReturningEvent(label="seed", count=1).create() diff --git a/plain-postgres/tests/typing/bulk_writes.py b/plain-postgres/tests/typing/bulk_writes.py new file mode 100644 index 0000000000..e742c8a683 --- /dev/null +++ b/plain-postgres/tests/typing/bulk_writes.py @@ -0,0 +1,83 @@ +"""`bulk_create()` inserts and `bulk_upsert()` insert-or-updates. + +Both hand back the model instances they were given. `bulk_upsert()` takes field +references for its conflict target and its update columns -- including a +`Model.fk` reference, which types as the related model class rather than a +Field -- and `bulk_create()` no longer takes a conflict surface at all. +""" + +from __future__ import annotations + +from typing import assert_type + +from app.examples.models.upsert import UpsertItem, UpsertScoped, UpsertTenant + + +def must_accept_bulk_create_as_instances() -> None: + assert_type( + UpsertItem.query.bulk_create([UpsertItem(key="a", value=1)]), + list[UpsertItem], + ) + + +def must_accept_bulk_upsert_as_instances() -> None: + assert_type( + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=1)], + update_fields=[UpsertItem.value], + unique_fields=[UpsertItem.key], + ), + list[UpsertItem], + ) + + +def must_reject_string_field_names() -> None: + # Runtime half: + # tests/public/test_bulk_upsert.py::test_bulk_upsert_string_field_rejected. + UpsertItem.query.bulk_upsert( + [UpsertItem(key="a", value=1)], + update_fields=["value"], # ty: ignore[invalid-argument-type] + unique_fields=[UpsertItem.key], + ) + + +def must_reject_bulk_create_conflict_kwargs() -> None: + # bulk_create is insert-only. Runtime half: tests/public/test_bulk_upsert.py + # ::test_bulk_create_no_longer_accepts_update_conflicts. + UpsertItem.query.bulk_create( + [UpsertItem(key="a", value=1)], + update_conflicts=True, # ty: ignore[unknown-argument] + ) + + +def must_accept_a_foreign_key_reference() -> None: + # Model.fk types as the related model class, not a Field, so the element + # type of these lists is the union. Runtime half: + # tests/public/test_bulk_upsert.py + # ::test_bulk_upsert_foreign_key_in_unique_fields. + UpsertScoped.query.bulk_upsert( + [UpsertScoped(tenant=UpsertTenant(name="t"), slug="s", value=1)], + update_fields=[UpsertScoped.value], + unique_fields=[UpsertScoped.tenant, UpsertScoped.slug], + ) + + +def must_accept_a_hoisted_conflict_target() -> None: + # list is invariant, so a hoisted [M.fk, M.key] infers + # list[type[Scope] | Field[str]] and would not satisfy a list parameter + # even though the same literal written inline does. The parameters are + # Sequence so naming the target is as good as inlining it. + conflict_target = [UpsertScoped.tenant, UpsertScoped.slug] + UpsertScoped.query.bulk_upsert( + [UpsertScoped(tenant=UpsertTenant(name="t"), slug="s", value=1)], + update_fields=[UpsertScoped.value], + unique_fields=conflict_target, + ) + + +def must_accept_tuples() -> None: + UpsertScoped.query.bulk_upsert( + [UpsertScoped(tenant=UpsertTenant(name="t"), slug="s", value=1)], + update_fields=(UpsertScoped.value,), + unique_fields=(UpsertScoped.tenant, UpsertScoped.slug), + )