From 5e8237d0b03123fc18d1199cd6f0adeb1db0d1d7 Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Fri, 14 Aug 2026 11:39:58 +0200 Subject: [PATCH 01/11] feat: versioned models --- docs/reference/api.md | 4 + docs/reference/types.md | 32 +++ docs/usage.md | 108 ++++++++++ sqlargon/__init__.py | 22 +- sqlargon/mixins.py | 81 +++++++- sqlargon/orm.py | 33 ++- sqlargon/repository.py | 191 +++++++++++++++++- tests/e2e/conftest.py | 26 ++- tests/e2e/models.py | 37 +++- tests/e2e/test_versioned.py | 192 ++++++++++++++++++ tests/test_versioned.py | 386 ++++++++++++++++++++++++++++++++++++ 11 files changed, 1102 insertions(+), 10 deletions(-) create mode 100644 tests/e2e/test_versioned.py create mode 100644 tests/test_versioned.py diff --git a/docs/reference/api.md b/docs/reference/api.md index bab540f..ceacfd9 100644 --- a/docs/reference/api.md +++ b/docs/reference/api.md @@ -8,8 +8,12 @@ Generated from the source. See [Usage](../usage.md) for a narrative introduction ::: sqlargon.repository.SoftDeleteRepository +::: sqlargon.repository.VersionedRepository + ::: sqlargon.repository.DeletedRowExistsError +::: sqlargon.repository.ConcurrentModificationError + ::: sqlargon.functools.atomic ## Unit of work diff --git a/docs/reference/types.md b/docs/reference/types.md index 7450ab1..96b385b 100644 --- a/docs/reference/types.md +++ b/docs/reference/types.md @@ -141,6 +141,9 @@ Both accept `sa_column_type=` to store in something other than `JSON` (e.g. `sa. | `UUIDV7ModelMixin` | `id` — `GUID` primary key, `uuid7` client default, `GenerateUUIDV7()` server default | | | `CreatedUpdatedMixin` | `created_at`, `updated_at` — `Timestamp`, `now()` server defaults, `onupdate` | `is_new` hybrid property | | `SoftDeleteMixin` | `tombstone` — boolean, defaults to false | `not_deleted` and `is_deleted` hybrid properties | +| `VersionedMixin` | *(abstract marker — no columns)* | | +| `UUIDVersionedMixin` | `version_id` — `GUID`, `uuid4` default, `GenerateUUID()` server default | `__mapper_args__` with `version_id_col` + UUID generator | +| `XminVersionedMixin` | `xmin` — PostgreSQL system column, `String`, `system=True`, `FetchedValue()` | `__mapper_args__` with `version_id_col` + `version_id_generator=False` | ```python from sqlargon.mixins import CreatedUpdatedMixin, SoftDeleteMixin, UUIDV7ModelMixin @@ -181,3 +184,32 @@ class User(UUIDV7ModelMixin, CreatedUpdatedMixin, SoftDeleteBase): ``` `SoftDeleteModel` is the matching type variable, bound to `SoftDeleteBase`. + +`VersionedMixin` is an abstract marker — use one of its concrete subclasses: + +- `UUIDVersionedMixin` — backend-agnostic, a `GUID` version column with a fresh UUID on + every update. `VersionedBase` combines it with `Base`: + +```python +from sqlargon import VersionedBase + + +class User(UUIDV7ModelMixin, VersionedBase): + name: Mapped[str] = mapped_column(sa.Unicode(255)) +``` + +- `XminVersionedMixin` — PostgreSQL only, uses the `xmin` system column (server-managed, + changes on every UPDATE). `XminVersionedBase` combines it with `Base`: + +```python +from sqlargon import XminVersionedBase + + +class User(UUIDV7ModelMixin, XminVersionedBase): + name: Mapped[str] = mapped_column(sa.Unicode(255)) +``` + +Both set `__mapper_args__` with `version_id_col`, enabling SQLAlchemy's ORM-level +versioning when using `AsyncSession` directly. `VersionedModel` is the matching type +variable, bound to `VersionedBase`. See [Versioned models](../usage.md#versioned-models) +for the repository API. diff --git a/docs/usage.md b/docs/usage.md index 260b3ce..288ea2b 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -289,6 +289,114 @@ class AuditedRepository(SoftDeleteRepository[SoftDeleteModel], abstract=True): class UserRepository(AuditedRepository[User]): ... ``` +## Versioned models + +`VersionedRepository` adds optimistic concurrency control: every update bumps a version +column, and `update_if_match` / `delete_if_match` check that the row's version matches +the one the caller loaded — if it does not, the row was modified by a concurrent +transaction and the call returns `None` (or raises `ConcurrentModificationError`). + +Two versioning strategies are provided: + +- **`UUIDVersionedMixin`** (backend-agnostic) — a `GUID` version column with a fresh UUID + on every update. Declare the model on `VersionedBase`, which combines the mixin with + `Base`: + +```python +from sqlargon import VersionedBase, VersionedRepository +from sqlargon.mixins import UUIDModelMixin + + +class User(UUIDModelMixin, VersionedBase): + name: Mapped[str] = mapped_column(sa.Unicode(255)) + + +class UserRepository(VersionedRepository[User]): ... + + +users = UserRepository() +``` + +- **`XminVersionedMixin`** (PostgreSQL only) — uses PostgreSQL's `xmin` system column, + which the server changes automatically on every UPDATE. Declare the model on + `XminVersionedBase`: + +```python +from sqlargon import XminVersionedBase + + +class User(UUIDModelMixin, XminVersionedBase): + name: Mapped[str] = mapped_column(sa.Unicode(255)) +``` + +The same `VersionedRepository` works with both strategies — it reads the version column +and generator from the SQLAlchemy mapper, so for `xmin` models it skips the client-side +increment (the server does it) and only adds the `WHERE` guard. + +### Auto-increment + +Every `update_one`, `update_many` and `bulk_update` call through a `VersionedRepository` +automatically sets a new version on the matched rows: + +```python +user = await users.get(id=user_id) +# user.version_id == "abc-123..." + +updated = await users.update_one({"name": "Jane"}, User.id == user_id) +# updated.version_id == "def-456..." (a fresh UUID) +``` + +The version column is excluded from the default `ON CONFLICT DO UPDATE` set, so an upsert +(`create_or_update`, `bulk_create_or_update`) never silently clobbers a version. + +### Optimistic concurrency check + +`update_if_match` and `delete_if_match` add `WHERE version_col = expected_version` to the +statement. If the row's version no longer matches, zero rows are touched and the call +returns `None`: + +```python +user = await users.get(id=user_id) + +updated = await users.update_if_match( + {"name": "Jane"}, + User.id == user_id, + expected_version=user.version_id, +) +if updated is None: + # someone else modified the row first — reload and retry +``` + +Pass `raise_on_mismatch=True` to raise `ConcurrentModificationError` instead of returning +`None`: + +```python +updated = await users.update_if_match( + {"name": "Jane"}, + User.id == user_id, + expected_version=user.version_id, + raise_on_mismatch=True, +) +``` + +`delete_if_match` works the same way: + +```python +deleted = await users.delete_if_match( + User.id == user_id, + expected_version=user.version_id, +) +``` + +You can also add a manual `Model.version_id == expected` filter to any regular method — +`update_one` will return `None` when the version does not match, without the convenience +of `update_if_match`'s `raise_on_mismatch` flag. + +`VersionedRepository` is generic over `VersionedModel`, a type variable bound to +`VersionedBase`. `XminVersionedBase` models work at runtime (the looser `VersionedMixin` +is enough) but need a `# type: ignore[type-var]` for the static bound, just as +`SoftDeleteMixin`-by-hand models do for `SoftDeleteModel`. + ## Building queries `select`, `insert`, `upsert`, `update`, `delete`, `where`/`filter`, `join` and `load` return diff --git a/sqlargon/__init__.py b/sqlargon/__init__.py index a740fb2..bab5718 100644 --- a/sqlargon/__init__.py +++ b/sqlargon/__init__.py @@ -3,12 +3,24 @@ from .cluster import AnyDatabase, DatabaseCluster from .database import BaseDatabase, Database, ReadOnlyDatabase, ReadOnlyError from .functools import atomic -from .orm import Base, Model, ORMModel, SoftDeleteBase, SoftDeleteModel +from .mixins import UUIDVersionedMixin, VersionedMixin, XminVersionedMixin +from .orm import ( + Base, + Model, + ORMModel, + SoftDeleteBase, + SoftDeleteModel, + VersionedBase, + VersionedModel, + XminVersionedBase, +) from .registry import get_default_database, set_default_database from .repository import ( + ConcurrentModificationError, DeletedRowExistsError, SoftDeleteRepository, SQLAlchemyRepository, + VersionedRepository, ) from .routing import ( DefaultRouter, @@ -32,6 +44,7 @@ "AnyDatabase", "Base", "BaseDatabase", + "ConcurrentModificationError", "Database", "DatabaseCluster", "DefaultRouter", @@ -52,6 +65,13 @@ "SoftDeleteBase", "SoftDeleteModel", "SoftDeleteRepository", + "UUIDVersionedMixin", + "VersionedBase", + "VersionedMixin", + "VersionedModel", + "VersionedRepository", + "XminVersionedBase", + "XminVersionedMixin", "__version__", "atomic", "get_default_database", diff --git a/sqlargon/mixins.py b/sqlargon/mixins.py index a8dfa3d..d35de7d 100644 --- a/sqlargon/mixins.py +++ b/sqlargon/mixins.py @@ -1,11 +1,12 @@ from datetime import datetime -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from uuid import UUID from weakref import WeakKeyDictionary import sqlalchemy as sa +from sqlalchemy import FetchedValue from sqlalchemy.ext.hybrid import hybrid_property -from sqlalchemy.orm import Mapped, mapped_column +from sqlalchemy.orm import Mapped, declared_attr, mapped_column from uuid_utils.compat import uuid4, uuid7 from .types import GUID, GenerateUUID, GenerateUUIDV7, Timestamp, now @@ -96,3 +97,79 @@ def is_deleted(self) -> bool: @classmethod def _is_deleted_expression(cls) -> sa.ColumnElement[bool]: return cls.tombstone.is_(sa.true()) + + +def _generate_version_uuid(_version: Any) -> UUID: + """Generate a fresh UUID for the version column, ignoring the prior value.""" + return uuid4() + + +class VersionedMixin: + """Marker mixin for optimistically versioned models. + + Use :class:`UUIDVersionedMixin` (backend-agnostic) or + :class:`XminVersionedMixin` (PostgreSQL-only) — this base only + exists so :class:`~sqlargon.repository.VersionedRepository` can + validate its model at subclass time, the same role + :class:`SoftDeleteMixin` plays for + :class:`~sqlargon.repository.SoftDeleteRepository`. + """ + + +class UUIDVersionedMixin(VersionedMixin): + """Backend-agnostic versioning via a UUID version column. + + Every UPDATE through + :class:`~sqlargon.repository.VersionedRepository` replaces + ``version_id`` with a fresh UUID; ``update_if_match`` / + ``delete_if_match`` add a ``WHERE version_id = :expected`` guard so + a stale row matches zero rows. + + The column uses the portable :class:`~sqlargon.types.GUID` type + (native UUID on PostgreSQL, CHAR(36) elsewhere) with both a + client-side ``default`` and a ``server_default`` so raw SQL inserts + also get a version. + """ + + version_id: Mapped[UUID] = mapped_column( + GUID(), + nullable=False, + default=uuid4, + server_default=GenerateUUID(), + ) + + @declared_attr.directive + def __mapper_args__(cls) -> dict[str, Any]: + return { + "eager_defaults": True, + "version_id_col": cls.version_id, + "version_id_generator": _generate_version_uuid, + } + + +class XminVersionedMixin(VersionedMixin): + """PostgreSQL-native versioning via the ``xmin`` system column. + + The ``xmin`` column is implicit (not in DDL); PostgreSQL changes it + on every UPDATE, so the repository does **not** set it — only + ``update_if_match`` / ``delete_if_match`` use it as a guard. + + The value is an ``xid``, which the driver decodes as it sees fit -- + asyncpg yields an ``int`` -- so the attribute is left untyped and the + guard compares the column as text. + + Only works on PostgreSQL. On other backends the column does not + exist and queries will fail. + """ + + xmin: Mapped[Any] = mapped_column( + "xmin", sa.String, system=True, server_default=FetchedValue() + ) + + @declared_attr.directive + def __mapper_args__(cls) -> dict[str, Any]: + return { + "eager_defaults": True, + "version_id_col": cls.xmin, + "version_id_generator": False, + } diff --git a/sqlargon/orm.py b/sqlargon/orm.py index 8f9a234..9492182 100644 --- a/sqlargon/orm.py +++ b/sqlargon/orm.py @@ -4,7 +4,11 @@ from sqlalchemy import MetaData from sqlalchemy.orm import DeclarativeBase, declared_attr -from .mixins import SoftDeleteMixin +from .mixins import ( + SoftDeleteMixin, + UUIDVersionedMixin, + XminVersionedMixin, +) camel_to_snake = re.compile(r"(? None: + super().__init_subclass__(abstract=abstract, **kwargs) + if abstract: + return + if not issubclass(cls.model, VersionedMixin): + msg = ( + f"{cls.model.__name__} must inherit from VersionedMixin " + f"(UUIDVersionedMixin or XminVersionedMixin) to be used " + f"with {cls.__name__}" + ) + raise TypeError(msg) + if cls._version_col() is None: + msg = ( + f"{cls.model.__name__} maps no version column: its " + f"__mapper_args__ carry no version_id_col, so " + f"{cls.__name__} could not tell a stale row from a fresh one" + ) + raise TypeError(msg) + + @classmethod + def _version_col(cls) -> Any: + """The version column from the mapper, or ``None``.""" + return getattr(cls.model.__mapper__, "version_id_col", None) + + @classmethod + def _version_generator(cls) -> Any: + """The version generator: ``False`` (server-side), ``True`` + (integer increment), or a callable.""" + return getattr(cls.model.__mapper__, "version_id_generator", True) + + @classmethod + def _is_server_versioned(cls) -> bool: + return cls._version_generator() is False + + @classmethod + def _get_default_set(cls) -> set[str]: + col = cls._version_col() + excluded = {col.name} if col is not None else set() + return super()._get_default_set() - excluded + + def _with_version_increment(self, values: Values) -> Values: + """Return ``values`` with the version column bumped.""" + col = self._version_col() + generator = self._version_generator() + if col is None or generator is False: + return values + name = col.name + if callable(generator): + if isinstance(values, MappingABC): + return {**values, name: generator(None)} + return [{**row, name: generator(None)} for row in values] + # generator is True — integer increment via SQL expression + if isinstance(values, MappingABC): + return {**values, name: col + 1} + # integer increment is not supported for executemany + return values + + def _version_filter(self, expected: Any) -> Any: + """The guard matching ``expected`` against the version column. + + A server managed version is PostgreSQL's ``xmin``, an ``xid`` no + driver binds natively: asyncpg decodes it into an ``int`` while a + version round tripped through a client comes back as a ``str``, and + ``xid = varchar`` is not an operator PostgreSQL has. Comparing the + column as text accepts either. + """ + col = self._version_col() + if self._is_server_versioned(): + return cast(col, Text) == str(expected) + return col == expected + + def update(self, values: Values, *, return_results: bool = False) -> Self: + values = self._with_version_increment(values) + return super().update(values, return_results=return_results) + + async def bulk_update( + self, + values: MultipleValues, + *args: Any, + on_: set[str] | None = None, + **kwargs: Any, + ) -> None: + col = self._version_col() + generator = self._version_generator() + if col is not None and callable(generator): + name = col.name + values = [{**row, name: generator(None)} for row in values] + await super().bulk_update(values, *args, on_=on_, **kwargs) + + async def update_if_match( + self, + values: SingleValue, + *args: _ColumnExpressionArgument[bool], + expected_version: Any, + raise_on_mismatch: bool = False, + **kwargs: Any, + ) -> VersionedModel | None: + """Update with ``WHERE version_col = expected_version``. + + Returns the updated model, or ``None`` if no row matched (the + version was stale or the row is gone). Raises + :class:`ConcurrentModificationError` when ``raise_on_mismatch`` is + ``True`` and no row matched. + """ + filters = (*args, self._version_filter(expected_version)) + result = await self._update_returning(values, *filters, **kwargs) + row = result.one_or_none() + if row is None and raise_on_mismatch: + msg = ( + f"{self.model.__name__} with version {expected_version!r} " + "was modified or deleted by a concurrent transaction" + ) + raise ConcurrentModificationError(msg) + return row + + async def delete_if_match( + self, + *args: _ColumnExpressionArgument[bool], + expected_version: Any, + raise_on_mismatch: bool = False, + **kwargs: Any, + ) -> VersionedModel | None: + """Delete with ``WHERE version_col = expected_version``. + + Returns the deleted model, or ``None`` if no row matched. Raises + :class:`ConcurrentModificationError` when ``raise_on_mismatch`` is + ``True`` and no row matched. + """ + filters = (*args, self._version_filter(expected_version)) + result = await self._delete_returning(*filters, **kwargs) + row = result.one_or_none() + if row is None and raise_on_mismatch: + msg = ( + f"{self.model.__name__} with version {expected_version!r} " + "was modified or deleted by a concurrent transaction" + ) + raise ConcurrentModificationError(msg) + return row diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 0e065b2..b292190 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -20,9 +20,12 @@ from .models import ( SERVER_DEFAULT_TABLES, TABLES, + XMIN_TABLES, DocumentRepository, SoftUserRepository, UserRepository, + VersionedUserRepository, + XminUserRepository, ) if TYPE_CHECKING: @@ -64,9 +67,12 @@ def database_url( @pytest.fixture(scope="session") def tables(backend: Backend) -> tuple[sa.Table, ...]: """The e2e tables this backend can hold.""" + result = TABLES if backend.server_side_uuid: - return TABLES + SERVER_DEFAULT_TABLES - return TABLES + result = result + SERVER_DEFAULT_TABLES + if backend.dialect == "postgresql": + result = result + XMIN_TABLES + return result @pytest.fixture(scope="session") @@ -146,6 +152,12 @@ def needs_json_key_operators(backend: Backend) -> None: pytest.skip(f"the {backend.name} JSON key operators match values, not keys") +@pytest.fixture +def needs_xmin(backend: Backend) -> None: + if backend.dialect != "postgresql": + pytest.skip(f"{backend.name} has no xmin system column") + + @pytest.fixture def users() -> UserRepository: return UserRepository() @@ -159,3 +171,13 @@ def documents() -> DocumentRepository: @pytest.fixture def soft_users() -> SoftUserRepository: return SoftUserRepository() + + +@pytest.fixture +def versioned_users() -> VersionedUserRepository: + return VersionedUserRepository() + + +@pytest.fixture +def xmin_users() -> XminUserRepository: + return XminUserRepository() diff --git a/tests/e2e/models.py b/tests/e2e/models.py index edfe03f..a10b49d 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -13,7 +13,15 @@ from pydantic import BaseModel from sqlalchemy.orm import Mapped, mapped_column -from sqlargon import Base, SoftDeleteBase, SoftDeleteRepository, SQLAlchemyRepository +from sqlargon import ( + Base, + SoftDeleteBase, + SoftDeleteRepository, + SQLAlchemyRepository, + VersionedBase, + VersionedRepository, + XminVersionedBase, +) from sqlargon.mixins import CreatedUpdatedMixin, UUIDModelMixin from sqlargon.types import GUID, JSON, GenerateUUID, GenerateUUIDV7, Timestamp, now from sqlargon.types.pydantic import Pydantic @@ -108,13 +116,38 @@ def on_conflict(self) -> OnConflictOptions: return {"index_elements": {"name"}, "set_": {"name"}} +class VersionedUser(UUIDModelMixin, VersionedBase): + __tablename__ = "e2e_versioned_user" + + name: Mapped[str] = mapped_column(sa.Unicode(64), unique=True) + + +class XminUser(UUIDModelMixin, XminVersionedBase): + """PostgreSQL-only model versioned via the ``xmin`` system column.""" + + __tablename__ = "e2e_xmin_user" + + name: Mapped[str] = mapped_column(sa.Unicode(64), unique=True) + + +class VersionedUserRepository(VersionedRepository[VersionedUser]): + default_order_by = VersionedUser.name + + +class XminUserRepository(VersionedRepository[XminUser]): # type: ignore[type-var] + default_order_by = XminUser.name + + def _tables(*models: type[Base]) -> tuple[sa.Table, ...]: return tuple(Base.metadata.tables[model.__tablename__] for model in models) #: Tables every backend can hold; the only ones the e2e suite creates. -TABLES: tuple[sa.Table, ...] = _tables(User, Document, SoftUser) +TABLES: tuple[sa.Table, ...] = _tables(User, Document, SoftUser, VersionedUser) #: Tables whose DDL carries a server side UUID default, which not every #: backend accepts -- see :attr:`~tests.e2e.backends.Backend.server_side_uuid`. SERVER_DEFAULT_TABLES: tuple[sa.Table, ...] = _tables(ServerDefaults) + +#: Tables that only PostgreSQL can hold (the ``xmin`` system column). +XMIN_TABLES: tuple[sa.Table, ...] = _tables(XminUser) diff --git a/tests/e2e/test_versioned.py b/tests/e2e/test_versioned.py new file mode 100644 index 0000000..c0d1c5f --- /dev/null +++ b/tests/e2e/test_versioned.py @@ -0,0 +1,192 @@ +"""E2E tests for VersionedRepository across real database backends. + +UUID-based versioning runs on every backend; ``xmin`` tests are scoped to +PostgreSQL via the ``needs_xmin`` fixture. +""" + +from __future__ import annotations + +import pytest + +from sqlargon import ConcurrentModificationError + +from .models import VersionedUser, XminUser + +NAMES = ("Andrew", "John", "Vincent") + + +async def seed(versioned_users, *names: str) -> None: + await versioned_users.bulk_create( + [{"name": name} for name in names], return_results=False + ) + + +# --- UUID versioning (all backends) --- + + +async def test_create_sets_initial_version(versioned_users): + user = await versioned_users.create(name="John") + + assert user is not None + assert user.version_id is not None + + +async def test_update_increments_version(versioned_users): + user = await versioned_users.create(name="John") + original = user.version_id + + updated = await versioned_users.update_one( + {"name": "Jane"}, VersionedUser.id == user.id + ) + + assert updated is not None + assert updated.version_id != original + + +async def test_update_if_match_matching_version(versioned_users): + user = await versioned_users.create(name="John") + + updated = await versioned_users.update_if_match( + {"name": "Jane"}, + VersionedUser.id == user.id, + expected_version=user.version_id, + ) + + assert updated is not None + assert updated.name == "Jane" + + +async def test_update_if_match_mismatched_version(versioned_users): + user = await versioned_users.create(name="John") + from uuid_utils.compat import uuid4 as _uuid4 # noqa: PLC0415 + + updated = await versioned_users.update_if_match( + {"name": "Jane"}, + VersionedUser.id == user.id, + expected_version=_uuid4(), + ) + + assert updated is None + # the row is unchanged + stored = await versioned_users.get(id=user.id) + assert stored is not None + assert stored.name == "John" + + +async def test_update_if_match_raises_on_mismatch(versioned_users): + user = await versioned_users.create(name="John") + from uuid_utils.compat import uuid4 as _uuid4 # noqa: PLC0415 + + with pytest.raises(ConcurrentModificationError, match="was modified or deleted"): + await versioned_users.update_if_match( + {"name": "Jane"}, + VersionedUser.id == user.id, + expected_version=_uuid4(), + raise_on_mismatch=True, + ) + + +async def test_delete_if_match_matching_version(versioned_users): + user = await versioned_users.create(name="John") + + deleted = await versioned_users.delete_if_match( + VersionedUser.id == user.id, + expected_version=user.version_id, + ) + + assert deleted is not None + assert deleted.name == "John" + assert await versioned_users.count() == 0 + + +async def test_delete_if_match_mismatched_version(versioned_users): + user = await versioned_users.create(name="John") + from uuid_utils.compat import uuid4 as _uuid4 # noqa: PLC0415 + + deleted = await versioned_users.delete_if_match( + VersionedUser.id == user.id, + expected_version=_uuid4(), + ) + + assert deleted is None + assert await versioned_users.count() == 1 + + +async def test_upsert_preserves_version(versioned_users): + user = await versioned_users.create(name="John") + original = user.version_id + + await versioned_users.create_or_update(id=user.id, name="Johnny") + + stored = await versioned_users.get(id=user.id) + assert stored is not None + assert stored.name == "Johnny" + assert stored.version_id == original + + +async def test_bulk_update_increments_version(versioned_users): + await seed(versioned_users, *NAMES) + rows = await versioned_users.list() + original = {row.id: row.version_id for row in rows} + + await versioned_users.bulk_update( + [{"id": row.id, "name": f"renamed {row.name}"} for row in rows] + ) + + updated = {row.id: row.version_id for row in await versioned_users.list()} + assert all(updated[rid] != original[rid] for rid in original) + + +# --- xmin versioning (PostgreSQL only) --- + + +@pytest.mark.usefixtures("needs_xmin") +async def test_xmin_changes_on_update(xmin_users): + user = await xmin_users.create(name="John") + original = user.xmin + + updated = await xmin_users.update_one({"name": "Jane"}, XminUser.id == user.id) + + assert updated is not None + assert updated.xmin != original + + +@pytest.mark.usefixtures("needs_xmin") +async def test_xmin_update_if_match_matching(xmin_users): + user = await xmin_users.create(name="John") + + updated = await xmin_users.update_if_match( + {"name": "Jane"}, + XminUser.id == user.id, + expected_version=user.xmin, + ) + + assert updated is not None + assert updated.name == "Jane" + + +@pytest.mark.usefixtures("needs_xmin") +async def test_xmin_update_if_match_mismatched(xmin_users): + user = await xmin_users.create(name="John") + + updated = await xmin_users.update_if_match( + {"name": "Jane"}, + XminUser.id == user.id, + expected_version="999999999", + ) + + assert updated is None + + +@pytest.mark.usefixtures("needs_xmin") +async def test_xmin_delete_if_match_matching(xmin_users): + user = await xmin_users.create(name="John") + + deleted = await xmin_users.delete_if_match( + XminUser.id == user.id, + expected_version=user.xmin, + ) + + assert deleted is not None + assert deleted.name == "John" + assert await xmin_users.count() == 0 diff --git a/tests/test_versioned.py b/tests/test_versioned.py new file mode 100644 index 0000000..7d7ccf6 --- /dev/null +++ b/tests/test_versioned.py @@ -0,0 +1,386 @@ +"""Unit tests for VersionedRepository and the versioning mixins. + +Mirrors ``test_soft_delete.py``: in-memory SQLite via the shared ``db`` +fixture, module-level models to survive ``--count=3`` re-registration. +""" + +from uuid import UUID + +import pytest +import sqlalchemy as sa + +from sqlargon import ( + Base, + ConcurrentModificationError, + Database, + SQLAlchemyRepository, + UUIDVersionedMixin, + VersionedBase, + VersionedMixin, + VersionedRepository, +) +from sqlargon.mixins import UUIDModelMixin +from sqlargon.types import GUID, GenerateUUID + +# Models defined at module level to avoid re-registration with --count=3 + + +class VersionedArticle(UUIDModelMixin, VersionedBase): + __tablename__ = "test_versioned_article" + name = sa.Column(sa.Unicode(255), nullable=True) + + +class HandVersioned(UUIDVersionedMixin, Base): + """The mixin combined with ``Base`` by hand, rather than VersionedBase.""" + + __tablename__ = "test_versioned_hand" + id = sa.Column(sa.Integer, primary_key=True) + name = sa.Column(sa.Unicode(255), nullable=True) + + +class PlainModel(Base): + __tablename__ = "test_versioned_plain" + id = sa.Column(sa.Integer, primary_key=True) + + +class MarkerOnly(VersionedMixin, Base): + """The abstract marker, with none of the columns its subclasses map.""" + + __tablename__ = "test_versioned_marker" + id = sa.Column(sa.Integer, primary_key=True) + + +class VersionedArticleRepository(VersionedRepository[VersionedArticle]): + default_order_by = VersionedArticle.id + + +class RawArticleRepository(SQLAlchemyRepository[VersionedArticle]): + """Unscoped view of the same table, to observe what is physically stored.""" + + default_order_by = VersionedArticle.id + + +class HandVersionedRepository(VersionedRepository[HandVersioned]): # type: ignore[type-var] + pass + + +# --- fixtures --- + + +@pytest.fixture(autouse=True) +async def tables(db: Database): + async with db.engine.begin() as conn: + await conn.run_sync(VersionedArticle.__table__.create, checkfirst=True) + yield + async with db.engine.begin() as conn: + await conn.run_sync(VersionedArticle.__table__.drop, checkfirst=True) + + +@pytest.fixture +def repository(): + return VersionedArticleRepository() + + +@pytest.fixture +def raw(): + return RawArticleRepository() + + +@pytest.fixture +async def article(repository): + return await repository.create(name="john") + + +# --- model validation --- + + +def test_model_without_versioned_mixin_is_rejected(): + with pytest.raises(TypeError, match="PlainModel must inherit from VersionedMixin"): + + class BadRepository(VersionedRepository[PlainModel]): # type: ignore[type-var] + pass + + +def test_model_without_version_column_is_rejected(): + # the marker alone maps nothing to compare, so every guard would match + # zero rows and report a concurrent modification that never happened + with pytest.raises(TypeError, match="MarkerOnly maps no version column"): + + class BadRepository(VersionedRepository[MarkerOnly]): # type: ignore[type-var] + pass + + +def test_model_declaring_the_mixin_by_hand_is_accepted(): + # the static bound asks for VersionedBase, but the version column and + # its mapper args are all the repository actually needs + assert HandVersionedRepository.model is HandVersioned + + +def test_abstract_subclass_needs_no_model(): + class Shared(VersionedRepository[VersionedArticle], abstract=True): + pass + + class Concrete(Shared): + pass + + assert Concrete.model is VersionedArticle + + +def test_model_can_be_passed_explicitly(): + class Explicit(VersionedRepository, model=VersionedArticle): + pass + + assert Explicit.model is VersionedArticle + + +# --- column & mapper --- + + +def test_uuid_version_column(): + column = VersionedArticle.__table__.c.version_id + assert isinstance(column.type, GUID) + assert not column.nullable + assert column.default.is_callable + assert isinstance(column.server_default.arg, GenerateUUID) + + +def test_mapper_version_id_col(): + assert ( + VersionedArticle.__mapper__.version_id_col + is VersionedArticle.__table__.c.version_id + ) + + +def test_mapper_version_id_generator(): + generator = VersionedArticle.__mapper__.version_id_generator + assert callable(generator) + # generator ignores the incoming version and returns a fresh UUID + assert isinstance(generator(None), UUID) + + +def test_version_excluded_from_default_conflict_set(): + assert "version_id" not in VersionedArticleRepository._get_default_set() + assert "version_id" in RawArticleRepository._get_default_set() + + +# --- auto-increment on update --- + + +@pytest.mark.usefixtures("tables") +async def test_create_sets_initial_version(repository): + obj = await repository.create(name="john") + + assert isinstance(obj.version_id, UUID) + + +@pytest.mark.usefixtures("tables") +async def test_update_one_increments_version(repository, article): + updated = await repository.update_one( + {"name": "jane"}, VersionedArticle.id == article.id + ) + + assert updated is not None + assert updated.version_id != article.version_id + + +@pytest.mark.usefixtures("tables") +async def test_update_many_increments_version(repository, raw): + await repository.bulk_create([{"name": "a"}, {"name": "b"}], return_results=False) + rows = await raw.all() + original = {row.id: row.version_id for row in rows} + + await repository.update_many( + {"name": "renamed"}, VersionedArticle.name.in_(["a", "b"]) + ) + + updated = {row.id: row.version_id for row in await raw.all()} + assert all(updated[rid] != original[rid] for rid in original) + + +@pytest.mark.usefixtures("tables") +async def test_bulk_update_increments_version(repository, raw): + await repository.bulk_create([{"name": "a"}, {"name": "b"}], return_results=False) + rows = await raw.all() + original = {row.id: row.version_id for row in rows} + + await repository.bulk_update([{"id": row.id, "name": "renamed"} for row in rows]) + + updated = {row.id: row.version_id for row in await raw.all()} + assert all(updated[rid] != original[rid] for rid in original) + + +# --- update_if_match --- + + +@pytest.mark.usefixtures("tables") +async def test_update_if_match_returns_row_on_matching_version(repository, article): + updated = await repository.update_if_match( + {"name": "jane"}, + VersionedArticle.id == article.id, + expected_version=article.version_id, + ) + + assert updated is not None + assert updated.name == "jane" + assert updated.version_id != article.version_id + + +@pytest.mark.usefixtures("tables") +async def test_update_if_match_returns_none_on_mismatched_version(repository, article): + from uuid_utils.compat import uuid4 as _uuid4 # noqa: PLC0415 + + wrong = _uuid4() + updated = await repository.update_if_match( + {"name": "jane"}, + VersionedArticle.id == article.id, + expected_version=wrong, + ) + + assert updated is None + + +@pytest.mark.usefixtures("tables") +async def test_update_if_match_raises_on_mismatch_when_configured(repository, article): + from uuid_utils.compat import uuid4 as _uuid4 # noqa: PLC0415 + + wrong = _uuid4() + with pytest.raises(ConcurrentModificationError, match="was modified or deleted"): + await repository.update_if_match( + {"name": "jane"}, + VersionedArticle.id == article.id, + expected_version=wrong, + raise_on_mismatch=True, + ) + + +@pytest.mark.usefixtures("tables") +async def test_update_if_match_accepts_kwargs(repository, article): + updated = await repository.update_if_match( + {"name": "jane"}, + expected_version=article.version_id, + id=article.id, + ) + + assert updated is not None + assert updated.name == "jane" + + +@pytest.mark.usefixtures("tables") +async def test_update_if_match_increments_version(repository, article): + updated = await repository.update_if_match( + {"name": "jane"}, + VersionedArticle.id == article.id, + expected_version=article.version_id, + ) + + assert updated is not None + assert updated.version_id != article.version_id + + +# --- delete_if_match --- + + +@pytest.mark.usefixtures("tables") +async def test_delete_if_match_returns_row_on_matching_version(repository, article): + deleted = await repository.delete_if_match( + VersionedArticle.id == article.id, + expected_version=article.version_id, + ) + + assert deleted is not None + assert deleted.name == "john" + assert await repository.count() == 0 + + +@pytest.mark.usefixtures("tables") +async def test_delete_if_match_returns_none_on_mismatched_version(repository, article): + from uuid_utils.compat import uuid4 as _uuid4 # noqa: PLC0415 + + wrong = _uuid4() + deleted = await repository.delete_if_match( + VersionedArticle.id == article.id, + expected_version=wrong, + ) + + assert deleted is None + assert await repository.count() == 1 + + +@pytest.mark.usefixtures("tables") +async def test_delete_if_match_raises_on_mismatch_when_configured(repository, article): + from uuid_utils.compat import uuid4 as _uuid4 # noqa: PLC0415 + + wrong = _uuid4() + with pytest.raises(ConcurrentModificationError, match="was modified or deleted"): + await repository.delete_if_match( + VersionedArticle.id == article.id, + expected_version=wrong, + raise_on_mismatch=True, + ) + + +@pytest.mark.usefixtures("tables") +async def test_delete_if_match_accepts_kwargs(repository, article): + deleted = await repository.delete_if_match( + expected_version=article.version_id, + id=article.id, + ) + + assert deleted is not None + assert deleted.name == "john" + + +# --- upsert / conflict --- + + +@pytest.mark.usefixtures("tables") +async def test_upsert_does_not_clobber_version(repository, raw, article): + original = article.version_id + + await repository.create_or_update(id=article.id, name="back?") + + stored = await raw.get(id=article.id) + assert stored is not None + assert stored.name == "back?" + assert stored.version_id == original + + +@pytest.mark.usefixtures("tables") +async def test_create_or_update_sets_version_on_insert(repository): + obj = await repository.create_or_update(name="fresh") + + assert obj is not None + assert isinstance(obj.version_id, UUID) + + +# --- regular methods with manual filters --- + + +@pytest.mark.usefixtures("tables") +async def test_update_one_with_manual_version_filter_returns_none_on_mismatch( + repository, article +): + from uuid_utils.compat import uuid4 as _uuid4 # noqa: PLC0415 + + wrong = _uuid4() + updated = await repository.update_one( + {"name": "jane"}, + VersionedArticle.id == article.id, + VersionedArticle.version_id == wrong, + ) + + assert updated is None + + +@pytest.mark.usefixtures("tables") +async def test_update_one_with_manual_version_filter_works_on_match( + repository, article +): + updated = await repository.update_one( + {"name": "jane"}, + VersionedArticle.id == article.id, + VersionedArticle.version_id == article.version_id, + ) + + assert updated is not None + assert updated.name == "jane" From adff554160a3875abe52d7682bf29e8af7b47491 Mon Sep 17 00:00:00 2001 From: RaRhAeu <37556570+RaRhAeu@users.noreply.github.com> Date: Mon, 24 Aug 2026 22:55:25 +0200 Subject: [PATCH 02/11] feat: outbox, tests optimization (#33) --- README.md | 114 ++- docs/auditable.md | 245 +++++ docs/cron.md | 7 + docs/index.md | 61 +- docs/outbox.md | 376 ++++++++ docs/reference/api.md | 30 +- docs/reference/dialects.md | 19 +- docs/reference/types.md | 83 +- docs/usage.md | 4 +- docs/vectors.md | 212 +++++ mkdocs.yaml | 3 + pyproject.toml | 23 +- sqlargon/__init__.py | 30 +- sqlargon/audit.py | 125 +++ sqlargon/cron/manager.py | 47 +- sqlargon/dialects/postgres.py | 135 ++- sqlargon/dialects/sqlite.py | 72 +- sqlargon/i18n/__init__.py | 40 + sqlargon/i18n/expression.py | 100 +++ sqlargon/i18n/mixin.py | 46 + sqlargon/i18n/repository.py | 39 + sqlargon/i18n/translatable.py | 243 +++++ sqlargon/i18n/translation.py | 216 +++++ sqlargon/integrations/eventiq.py | 81 ++ sqlargon/mixins.py | 148 +++ sqlargon/orm.py | 65 +- sqlargon/outbox/__init__.py | 23 + sqlargon/outbox/config.py | 76 ++ sqlargon/outbox/models.py | 45 + sqlargon/outbox/relay.py | 151 ++++ sqlargon/outbox/repository.py | 393 ++++++++ sqlargon/query_builder.py | 108 +++ sqlargon/repository/__init__.py | 14 + sqlargon/repository/auditable.py | 538 +++++++++++ .../{repository.py => repository/base.py} | 417 +-------- sqlargon/repository/soft_delete.py | 218 +++++ sqlargon/repository/versioned.py | 211 +++++ sqlargon/types/json.py | 617 ++++++++++++- sqlargon/types/vector.py | 230 +++++ sqlargon/vectors/__init__.py | 61 ++ sqlargon/vectors/loader.py | 84 ++ sqlargon/vectors/mixins.py | 158 ++++ sqlargon/vectors/models.py | 95 ++ sqlargon/vectors/repository.py | 288 ++++++ tests/conftest.py | 6 + tests/e2e/backends.py | 11 +- tests/e2e/conftest.py | 104 ++- tests/e2e/models.py | 156 +++- tests/e2e/test_auditable.py | 330 +++++++ tests/e2e/test_cron.py | 4 +- tests/e2e/test_outbox.py | 221 +++++ tests/e2e/test_types.py | 156 ++++ tests/e2e/test_vectors.py | 253 ++++++ tests/test_auditable.py | 846 ++++++++++++++++++ tests/test_cron.py | 41 + tests/test_database.py | 8 +- tests/test_dialects.py | 9 +- tests/test_eventiq.py | 223 +++++ tests/test_i18n.py | 190 ++++ tests/test_mixins.py | 17 + tests/test_outbox.py | 795 ++++++++++++++++ tests/test_types.py | 487 +++++++++- tests/test_vectors.py | 580 ++++++++++++ uv.lock | 154 +++- 64 files changed, 10377 insertions(+), 505 deletions(-) create mode 100644 docs/auditable.md create mode 100644 docs/outbox.md create mode 100644 docs/vectors.md create mode 100644 sqlargon/audit.py create mode 100644 sqlargon/i18n/__init__.py create mode 100644 sqlargon/i18n/expression.py create mode 100644 sqlargon/i18n/mixin.py create mode 100644 sqlargon/i18n/repository.py create mode 100644 sqlargon/i18n/translatable.py create mode 100644 sqlargon/i18n/translation.py create mode 100644 sqlargon/integrations/eventiq.py create mode 100644 sqlargon/outbox/__init__.py create mode 100644 sqlargon/outbox/config.py create mode 100644 sqlargon/outbox/models.py create mode 100644 sqlargon/outbox/relay.py create mode 100644 sqlargon/outbox/repository.py create mode 100644 sqlargon/repository/__init__.py create mode 100644 sqlargon/repository/auditable.py rename sqlargon/{repository.py => repository/base.py} (60%) create mode 100644 sqlargon/repository/soft_delete.py create mode 100644 sqlargon/repository/versioned.py create mode 100644 sqlargon/types/vector.py create mode 100644 sqlargon/vectors/__init__.py create mode 100644 sqlargon/vectors/loader.py create mode 100644 sqlargon/vectors/mixins.py create mode 100644 sqlargon/vectors/models.py create mode 100644 sqlargon/vectors/repository.py create mode 100644 tests/e2e/test_auditable.py create mode 100644 tests/e2e/test_outbox.py create mode 100644 tests/e2e/test_vectors.py create mode 100644 tests/test_auditable.py create mode 100644 tests/test_eventiq.py create mode 100644 tests/test_i18n.py create mode 100644 tests/test_outbox.py create mode 100644 tests/test_vectors.py diff --git a/README.md b/README.md index 34f63cb..c294002 100644 --- a/README.md +++ b/README.md @@ -21,6 +21,40 @@ Repository: https://github.com/asynq-io/sqlargon --- +## Features + +- **Repository pattern** — one object wraps async sessions, core queries and ORM models; + sessions are context-local and resolved at call time, so nothing gets passed around +- **High-level CRUD** — `create`, `get`, `get_or_create`, `create_or_update`, `all`, + `list`, `count`, `update_one`, `update_many`, `delete_one`, `delete_many` and `remove` + out of the box +- **Bulk operations** — `bulk_create`, `bulk_create_or_update` and `bulk_update` with + per-repository conflict handling +- **Query builder** — fluent, dialect-aware statements for upserts, `RETURNING`, advisory + locks and streaming, with terminal helpers that cast results to `.scalars()`, `.one()`, + `.mappings()`, ... +- **Multi-dialect** — PostgreSQL, SQLite, MySQL and MariaDB, with capability-gated SQL + generation per backend +- **Transactions** — `@atomic` and database-scoped `atomic()` blocks, plus named advisory + locks +- **Unit of work** — repositories declared as annotations on a unit of work share one + session and one transaction +- **Database routing** — clusters with read replicas, shards and vertical partitioning; + `using()`, `read_only` and per-request `use_context` +- **Pagination** — page-number, offset/limit and keyset cursor strategies +- **Outbox** — transactional outbox with a background relay and eventiq integration +- **Cron** — database-backed scheduler with namespaces and safe multi-instance claiming +- **Column types and mixins** — UUID (v4/v7), timestamp, orjson JSON and pydantic-validated + columns; mixins for UUID keys, created/updated timestamps and soft delete +- **Soft delete** — tombstone-based deletes via `SoftDeleteRepository` +- **Versioned models** — optimistic concurrency with UUID or PostgreSQL `xmin` versions +- **Auditable models** — append-only versioned history with point-in-time reads and restore +- **Vector search** — embeddings with cosine, L2, dot and L1 similarity, full-text and + hybrid reciprocal-rank-fusion search on PostgreSQL and SQLite +- **FastAPI-ready** — repositories and units of work work directly as dependencies +- **Alembic migrations** — async-first migration setup +- **OpenTelemetry** — optional SQLAlchemy instrumentation + ## About This library provides glue code to use sqlalchemy async sessions, core queries and orm models @@ -37,6 +71,7 @@ from one object which provides somewhat of repository pattern. This solution has - engines and routing policy are separate, so the same repository runs against one database, a primary with read replicas, or a set of shards + ## Installation ```shell @@ -50,7 +85,8 @@ uv add sqlargon ``` Optional extras: `postgres`, `sqlite`, `mysql`, `pagination` (cursor pagination), -`cron`, `opentelemetry`, or `standard` for all of them: +`cron`, `outbox`, `eventiq`, `opentelemetry`, or `standard` for all of them +except `eventiq`: ```shell pip install "sqlargon[standard]" @@ -292,14 +328,88 @@ Available strategies: `PageNumberPagination`, `TotalPageNumberPagination`, `LimitOffsetPagination`, `TotalLimitOffsetPagination` and `CursorPagination` (keyset, requires `sqlargon[pagination]`). +## Outbox + +`sqlargon.outbox` implements the transactional outbox pattern: a write through the repository +also appends a CloudEvent-shaped row to `outbox_events` **in the same transaction**, so an +event can neither be lost by a rollback nor published for a row that never committed. A +background relay then publishes them in write order (requires `sqlargon[outbox]`): + +```python +from sqlargon import Base +from sqlargon.outbox import OutboxConfig, OutboxRelay, OutboxRepository + + +class User(UUIDModelMixin, CreatedUpdatedMixin, Base): # is_new tells insert from update + name: Mapped[str] = mapped_column(sa.Unicode(255)) + password: Mapped[str] = mapped_column(sa.Unicode(255)) + + +class UserRepository(OutboxRepository[User]): + outbox = OutboxConfig(topic="users", exclude={"password"}) + + +await UserRepository().create(name="John", password=hashed) +# -> one outbox_events row, type "user.created", password left out + +async with OutboxRelay(publish).running(): # publish is any async callable + ... +``` + +The relay takes a plain publisher callable, so sqlargon depends on no broker client; +`sqlargon.integrations.eventiq` adapts the rows to `eventiq.CloudEvent`. Only writes that go +through the repository are recorded. See the +[outbox docs](https://asynq-io.github.io/sqlargon/outbox/) for retention, retries and +ordering. + ## Column types and mixins `sqlargon.types` provides dialect-aware column types: `GUID` with `GenerateUUID` / `GenerateUUIDV7` server defaults, `Timestamp` with a `now()` server default and `JSON` -(orjson-serialized). `sqlargon.types.pydantic` adds `Pydantic` and `ValidatedType` for +(orjson-serialized), whose comparator carries portable JSON operators — containment and +key tests, plus server-side mutation (`set_key`, `update`, `remove_key`) that rewrites a +document in the `UPDATE` itself. `sqlargon.types.pydantic` adds `Pydantic` and `ValidatedType` for pydantic-validated columns. `sqlargon.mixins` bundles them into `UUIDModelMixin`, `UUIDV7ModelMixin`, `CreatedUpdatedMixin` and `SoftDeleteMixin`. +## Auditable models + +`AuditableRepository` never updates a row: every write appends the next `version` of the +same entity, so the table *is* the audit log. Reads are scoped to the newest live version, +so the usual methods keep their usual meaning: + +```python +from sqlargon import AuditableBase, AuditableRepository +from sqlargon.mixins import UUIDModelMixin + + +class Article(UUIDModelMixin, AuditableBase): + title: Mapped[str] = mapped_column(sa.Unicode(255)) + + +class ArticleRepository(AuditableRepository[Article]): ... + + +articles = ArticleRepository() + +article = await articles.create(title="draft") # version 1 +await articles.update_one({"title": "final"}, Article.id == article.id) # version 2 + +await articles.get(id=article.id) # version 2 +await articles.history(id=article.id) # versions 1 and 2 +await articles.get_version(1, id=article.id) # version 1 +await articles.at(yesterday).list() # the state as of yesterday + +await articles.remove(Article.id == article.id) # appends a tombstoned version 3 +await articles.restore(Article.id == article.id) # and a live version 4 +``` + +The version joins the primary key, so concurrent appends collide there rather than one +silently winning, and `update_if_match` gives the cheaper check first. Versions are either +a human-readable counter (`AuditableBase`) or a sortable UUIDv7 (`UUIDAuditableBase`), and +`sqlargon.audit` relates other tables to one exact version or to whichever is newest. See +the [documentation](https://asynq-io.github.io/sqlargon/auditable/) for the full picture. + ## FastAPI Repository and unit-of-work `__init__` take no arguments, so subclasses work directly as diff --git a/docs/auditable.md b/docs/auditable.md new file mode 100644 index 0000000..b1c3d2e --- /dev/null +++ b/docs/auditable.md @@ -0,0 +1,245 @@ +# Auditable models + +An auditable repository never updates a row. Every change appends a new row carrying the +next `version` of the same entity, so the table *is* the audit log — the history is the +data, not a copy of it in a side table that can drift. + +Reads are scoped to the newest live version, so the usual repository methods keep their +usual meaning and the history stays underneath, one method away. + +## Declaring the model + +Inherit `AuditableBase` and combine it with whatever identifies the entity — usually +`UUIDModelMixin`: + +```python +import sqlalchemy as sa +from sqlalchemy.orm import Mapped, mapped_column + +from sqlargon import AuditableBase, AuditableRepository +from sqlargon.mixins import UUIDModelMixin + + +class Article(UUIDModelMixin, AuditableBase): + title: Mapped[str] = mapped_column(sa.Unicode(255)) + body: Mapped[str | None] = mapped_column(sa.Text, nullable=True) + + +class ArticleRepository(AuditableRepository[Article]): ... +``` + +The version column joins the primary key, so `Article` is keyed by `(id, version)` and its +**entity key** is derived as everything in the primary key except the version — here +`(id,)`. Nothing else to configure. + +`AuditableBase` also brings `created_at` / `updated_at` and the `tombstone` column, so the +model carries when each version was written and which one records a deletion. Because a row +is never updated, `updated_at` always equals `created_at`. + +!!! tip "Put the entity key in a mixin" + Columns are ordered by when they were declared, so an entity key declared in the model + body lands *after* `version` in the primary key and its index. Taking the key from a + mixin — `UUIDModelMixin` above — keeps the index `(id, version)`, which is the order + every latest-version read wants. + +## Writing + +Every write appends. `create` starts an entity at version 1, and each subsequent write adds +the next version: + +```python +articles = ArticleRepository() + +article = await articles.create(title="draft") # version 1 +await articles.update_one({"title": "revised"}, Article.id == article.id) # version 2 +await articles.update_one({"title": "final"}, Article.id == article.id) # version 3 +``` + +A column the write does not name is carried forward from the version it supersedes, so an +append is a change to the entity, not a replacement of it. + +`update()` and `delete()` still build a statement, so the fluent form works unchanged — it +compiles to an `INSERT ... SELECT` reading the current heads and writing their successors, +which makes an append one statement rather than a read followed by a write: + +```python +await articles.update({"title": "final"}).filter(Article.id == article.id) +``` + +`create_or_update` appends to an entity that exists and creates one that does not; +`bulk_create_or_update` is its bulk form, and `bulk_update` appends a different set of +values per entity: + +```python +await articles.create_or_update(id=article.id, title="fourth") +await articles.bulk_create_or_update( + [{"id": article.id, "title": "fifth"}, {"id": uuid4(), "title": "brand new"}] +) +await articles.bulk_update([{"id": a_id, "title": "a"}, {"id": b_id, "title": "b"}]) +``` + +The one statement an append-only table cannot serve is `upsert()`, whose whole meaning is +"resolve a conflict by rewriting the conflicting row". It raises `AppendOnlyError` and +points at the two methods above. + +## Deleting + +Deletion appends a tombstoned version rather than removing anything, so `remove`, +`delete_one` and `delete_many` are all recoverable: + +```python +await articles.remove(Article.id == article.id) # appends version 4, tombstoned + +await articles.list() # the entity is gone from reads +await articles.versions().count() # but all four rows are still there + +await articles.restore(Article.id == article.id) # appends version 5, live again +``` + +## Reading + +| Call | Returns | +| --- | --- | +| `list()` / `get()` / `first()` | the newest live version of each entity | +| `count()` | how many **entities** there are | +| `versions()` | a copy covering every version of every entity | +| `versions().count()` | how many **rows** there are | +| `history(...)` | every version of the matched entities, oldest first | +| `get_version(n, ...)` | one exact version, tombstoned or not | +| `at(timestamp)` | a copy reading the state as it stood at that moment | +| `with_deleted()` | the newest version even when it is a tombstone | +| `only_deleted()` | entities whose newest version is a tombstone | + +```python +await articles.get(id=article.id) # version 3 +await articles.history(id=article.id) # versions 1, 2 and 3 +await articles.get_version(2, id=article.id) # version 2 + +yesterday = utc_now() - timedelta(days=1) +await articles.at(yesterday).list() # the state as of yesterday +``` + +`at()` resolves each entity to the newest version recorded up to that moment, and keeps an +entity tombstoned by then hidden — exactly as it would have been at the time. + +The latest-version scope is a correlated subquery, available on the model itself as +`Article.is_latest()`, so it composes into any query of your own: + +```sql +SELECT ... FROM article +WHERE article.version = (SELECT v.version FROM article v + WHERE v.id = article.id + ORDER BY v.version DESC LIMIT 1) + AND NOT article.tombstone +``` + +## Concurrency + +Because the version is part of the primary key, two writers deriving the same successor +collide there rather than one of them silently winning — the loser gets an `IntegrityError` +instead of losing its append. + +`update_if_match` and `delete_if_match` are inherited from +[`VersionedRepository`](usage.md#versioned-models) and land on the append path, giving the +cheaper check first: + +```python +from sqlargon import ConcurrentModificationError + +article = await articles.get(id=article_id) +try: + await articles.update_if_match( + {"title": "final"}, + Article.id == article_id, + expected_version=article.version, + raise_on_mismatch=True, + ) +except ConcurrentModificationError: + ... # someone appended a version first +``` + +## Versioning strategies + +| Base | Version column | Successor | +| --- | --- | --- | +| `AuditableBase` | `Integer`, starting at 1 | `version + 1`, in SQL | +| `UUIDAuditableBase` | `GUID`, UUIDv7 | a fresh UUIDv7, minted without reading the current one | + +Pick `AuditableBase` for a human-readable document version — 1, 2, 3 — and for strict +chronological ordering. Pick `UUIDAuditableBase` when writers cannot coordinate on a +counter: a UUIDv7 is time-sortable, so the newest version is still the greatest one. + +!!! warning "UUIDv7 ordering across processes" + `uuid7()` is monotonic within a process, but two processes appending in the same + millisecond can produce an inverted pair, which would make the older of the two look + newest. Use `AuditableBase` where that matters. + +## Relationships + +A row of an auditable model is one *version* of an entity, so a reference to it has to say +which version it means. `sqlargon.audit` covers both answers. + +### Pinned to an exact version + +The child stores the entity key and the version, under a real composite foreign key. That +makes it an ordinary many-to-one — writable, and joined without a `primaryjoin`: + +```python +from sqlargon.audit import version_foreign_key, version_mapped_column + + +class Comment(UUIDModelMixin, Base): + article_id: Mapped[UUID] = mapped_column(GUID()) + article_version = version_mapped_column(Article) + body: Mapped[str] = mapped_column(sa.Text) + + __table_args__ = ( + version_foreign_key(Article, "article_id", "article_version"), + ) + + article: Mapped[Article] = relationship() +``` + +`version_mapped_column` types itself from the parent, so the child never has to know which +versioning strategy it uses. The comment keeps pointing at the version it was written +against, however far the article moves on. + +### Following the latest version + +The child stores only the entity key. There can be no foreign key — the parent's primary +key holds a version this child deliberately does not pin — so the relationship carries the +latest predicate in its join and is necessarily `viewonly`: + +```python +from sqlargon.audit import latest_relationship + + +class Bookmark(UUIDModelMixin, Base): + article_id: Mapped[UUID] = mapped_column(GUID()) + + article: Mapped[Article] = latest_relationship(Article, "article_id") +``` + +Both work under `selectinload`, which the repository's `load()` uses: + +```python +await comments.load(Comment.article).all() +await bookmarks.load(Bookmark.article).all() +``` + +Pass `uselist=True` for a collection. + +## Retention + +`purge()` is the only method that destroys history — for retention, not for deletion, which +`remove` records instead. It physically deletes every superseded version and keeps the +newest: + +```python +await articles.purge(id=article.id) +await articles.purge(Article.created_at < cutoff) +``` + +A version a `version_foreign_key` still points at is protected by that key: purging it +raises rather than orphaning the child. Declare the key `ondelete="CASCADE"` if you would +rather the children went with it. diff --git a/docs/cron.md b/docs/cron.md index 49a7604..10cce7c 100644 --- a/docs/cron.md +++ b/docs/cron.md @@ -38,6 +38,13 @@ decorated function cleans up its row on the next start. async def send_report() -> None: ... ``` +`task()` also takes the function outright, for one that is not yours to +decorate — a method of an object built elsewhere, say: + +```python +cron.task("0 3 * * *", "purge_outbox", relay.purge) +``` + Sync only ever deletes rows it created (they are marked as declarative), so imperative tasks in the same namespace are never deleted. Declaring a name that an imperative row already holds takes that row over — the declaration is diff --git a/docs/index.md b/docs/index.md index 8e88fdd..fefc688 100644 --- a/docs/index.md +++ b/docs/index.md @@ -13,7 +13,6 @@ *SQLAlchemy repository pattern and utilities* --- -Version: 1.0.0b1 Docs: [https://asynq-io.github.io/sqlargon/](https://asynq-io.github.io/sqlargon/) @@ -21,23 +20,43 @@ Repository: [https://github.com/asynq-io/sqlargon](https://github.com/asynq-io/s --- -## About +## Features + +- **Repository pattern** — one object wraps async sessions, core queries and ORM models; + sessions are context-local and resolved at call time, so nothing gets passed around +- **High-level CRUD** — `create`, `get`, `get_or_create`, `create_or_update`, `all`, + `list`, `count`, `update_one`, `update_many`, `delete_one`, `delete_many` and `remove` + out of the box +- **Bulk operations** — `bulk_create`, `bulk_create_or_update` and `bulk_update` with + per-repository conflict handling +- **Query builder** — fluent, dialect-aware statements for upserts, `RETURNING`, advisory + locks and streaming, with terminal helpers that cast results to `.scalars()`, `.one()`, + `.mappings()`, ... +- **Multi-dialect** — PostgreSQL, SQLite, MySQL and MariaDB, with capability-gated SQL + generation per backend +- **Transactions** — `@atomic` and database-scoped `atomic()` blocks, plus named advisory + locks +- **Unit of work** — repositories declared as annotations on a unit of work share one + session and one transaction +- **Database routing** — [clusters](routing.md) with read replicas, shards and vertical + partitioning; `using()`, `read_only` and per-request `use_context` +- **Pagination** — [page-number, offset/limit and keyset cursor](pagination.md) strategies +- **Outbox** — [transactional outbox](outbox.md) with a background relay and eventiq + integration +- **Cron** — [database-backed scheduler](cron.md) with namespaces and safe multi-instance + claiming +- **Column types and mixins** — UUID (v4/v7), timestamp, orjson JSON and pydantic-validated + columns; mixins for UUID keys, created/updated timestamps and soft delete +- **Soft delete** — tombstone-based deletes via `SoftDeleteRepository` +- **Versioned models** — optimistic concurrency with UUID or PostgreSQL `xmin` versions +- **Auditable models** — [append-only versioned history](auditable.md) with point-in-time + reads and restore +- **Vector search** — [embeddings with similarity, full-text and hybrid + reciprocal-rank-fusion search](vectors.md) on PostgreSQL and SQLite +- **FastAPI-ready** — repositories and units of work work directly as dependencies +- **Alembic migrations** — async-first [migration setup](migrations.md) +- **OpenTelemetry** — optional SQLAlchemy instrumentation -SQLArgon provides glue code to use SQLAlchemy async sessions, core queries and ORM models -from one object which provides somewhat of a repository pattern. This solution has a few -advantages: - -- no need to pass a `session` object to every function/method — sessions are context-local - and resolved by the repository itself -- write data access queries in one place -- no need to import `insert`, `update`, `delete`, `select` from SQLAlchemy over and over again -- implicit cast of results to `.scalars().all()`, `.one()`, `.mappings()`, ... -- a dialect-aware query builder for upserts, `RETURNING` and advisory locks -- your view model (e.g. FastAPI routes) does not need to know about the underlying storage — - the repository class can be replaced at any moment with any object providing a similar - interface -- engines and routing policy are separate, so the same repository runs against one database, - a primary with read replicas, or a set of shards ## Installation @@ -60,8 +79,10 @@ Drivers and optional features ship as extras: | `mysql` | `asyncmy` | | `pagination` | `sqlakeyset`, required for cursor pagination | | `cron` | `croniter`, `anyio` | +| `outbox` | `anyio` | +| `eventiq` | `eventiq`, required by the outbox integration layer | | `opentelemetry` | `opentelemetry-instrumentation-sqlalchemy` | -| `standard` | all of the above | +| `standard` | all of the above except `eventiq` | ```shell pip install "sqlargon[standard]" @@ -111,6 +132,10 @@ or from `DATABASE_*` environment variables. - **[Usage](usage.md)** — models, CRUD, query building, transactions and units of work. - **[Database Routing](routing.md)** — replicas, shards, routers and FastAPI wiring. - **[Pagination](pagination.md)** — page-number, offset/limit and cursor strategies. +- **[Cron](cron.md)** — database-backed scheduling with namespaces and multi-instance safety. +- **[Outbox](outbox.md)** — the transactional outbox pattern and its relay. +- **[Vector Search](vectors.md)** — embeddings, similarity and hybrid search. +- **[Auditable Models](auditable.md)** — append-only versioned history. - **[Examples](examples.md)** — end-to-end recipes: a FastAPI service, batch workers, multi-tenant sharding, testing. - **Reference** — [types and mixins](reference/types.md), [dialects](reference/dialects.md), diff --git a/docs/outbox.md b/docs/outbox.md new file mode 100644 index 0000000..8cbcf11 --- /dev/null +++ b/docs/outbox.md @@ -0,0 +1,376 @@ +# Outbox + +`sqlargon.outbox` implements the transactional outbox pattern: a write through +a repository also appends a [CloudEvent](https://cloudevents.io)-shaped row to +the `outbox_events` table **in the same transaction**, and a background relay +publishes those rows to a broker. Because the row and its event commit or roll +back together, an event can neither be lost by a rollback nor published for a +row that never committed. + +It requires the `sqlargon[outbox]` extra (`anyio`, for the relay). + +```python +import sqlalchemy as sa +from sqlalchemy.orm import Mapped, mapped_column + +from sqlargon import Base +from sqlargon.mixins import CreatedUpdatedMixin, UUIDModelMixin +from sqlargon.outbox import OutboxConfig, OutboxRelay, OutboxRepository + + +class User(UUIDModelMixin, CreatedUpdatedMixin, Base): + name: Mapped[str] = mapped_column(sa.Unicode(255)) + password: Mapped[str] = mapped_column(sa.Unicode(255)) + + +class UserRepository(OutboxRepository[User]): + outbox = OutboxConfig(topic="users", exclude={"password"}) + + +await UserRepository().create(name="John", password=hashed) +# -> one outbox_events row, type "user.created", password left out +``` + +sqlargon never talks to a broker itself. The relay takes a *publisher* — any +`async def publish(event: OutboxEvent) -> None` — so any client will do; see +[eventiq](#eventiq) for the adapter that ships with it. + +## Configuring a repository + +The `OutboxConfig` assigned to the repository's `outbox` attribute decides what +its writes look like as events. The model must carry `CreatedUpdatedMixin` — +its `is_new` is what tells an insert from an update — but nothing beyond that: +no declarative base of its own, and whether a write is recorded is decided by +the repository it goes through, the same way `on_conflict` is. So two +repositories over the same model can publish differently — useful when one of +them serves a less trusted consumer: + +```python +class PublicUserRepository(OutboxRepository[User]): + outbox = OutboxConfig(topic="users.public", include=frozenset({"id", "name"})) +``` + +The default is a bare `OutboxConfig()`, so a repository that configures nothing +records every write with the table name as topic and type prefix. + +| Field | Meaning | +| --- | --- | +| `topic` | The CloudEvents `subject`. Defaults to the table name. | +| `type_prefix` | Prefix of the CloudEvents `type`. Defaults to the table name. | +| `source` | The CloudEvents `source`. Defaults to `None`, leaving it to the relay. | +| `exclude` | Columns kept out of the payload. | +| `include` | The payload columns outright; wins over `exclude`. | +| `operations` | Which writes are recorded at all. | +| `attributes` | Extra CloudEvents attributes, next to the payload rather than inside it. | + +The event `type` is `f"{type_prefix}.{operation}"`, so the default gives +`user.created`, `user.updated` and `user.deleted`. + +Restricting the operations recorded, for a table whose updates nobody cares +about: + +```python +from sqlargon.outbox import Operation + +OutboxConfig(operations=frozenset({Operation.CREATED, Operation.DELETED})) +``` + +The topic, the event type of each recorded operation and the payload columns +are derived from the config, and readable for inspection as `repository.topic`, +`repository.event_types` and `repository.payload_columns`. + +### Topic templating + +A `topic` may carry `{placeholders}` naming attributes of the written row, for +a topic that identifies its subject rather than the table that holds it — say +one topic per tenant or per entity. Each placeholder is filled from the row the +event was written from, at write time, using `str.format`: + +```python +class UserRepository(OutboxRepository[User]): + outbox = OutboxConfig( + topic="events.organizations.{organization_id}.deleted", + exclude={"password"}, + ) + + +await UserRepository().create( + name="John", password=hashed, organization_id=21 +) +# -> topic "events.organizations.21.deleted" +``` + +A topic without placeholders is used verbatim, so nothing changes for topics +that do not use them. A placeholder names a column — or any other attribute of +the row — and its value is stringified as it is: a `UUID` becomes its +canonical string form. `{id}` reads the row's `id`, `{organization_id}` its +`organization_id`, and so on. + +The value is read the same way the payload and the extra attributes are: from +the row the write produced (before it, for a delete). It is read **when the +write happens**, not when the relay publishes the event, so a templated topic +always reflects the state at write time. A repository serves one event per +written row, so a bulk write of rows from different organizations lands on +their own topics. + +`format_topic(topic, row)` does the substitution on its own and is exported +from `sqlargon.outbox`. + +### Extra attributes + +Some values belong *next to* the payload rather than inside it — a `tenant_id` +a broker routes on, the trace of the request that caused the write. `attributes` +names them, and they are stored in the event's `attributes` column and +published as top-level CloudEvents attributes: + +```python +from contextvars import ContextVar + +traceparent: ContextVar[str] = ContextVar("traceparent") + + +class UserRepository(OutboxRepository[User]): + outbox = OutboxConfig( + topic="users", + exclude={"password", "tenant_id"}, + attributes={ + "tenant_id": "tenant_id", # an attribute of the written row + "traceparent": lambda _: traceparent.get(), # anything else + }, + ) +``` + +A `str` names an attribute of the written row — a column, or any other +attribute of the model — and a callable is handed the row and returns the +value, which is how a `ContextVar` is read. Either way the value is read +**when the write happens**, not when the event is published: the relay runs in +a background task, long after the request context the write ran in is gone. + +A column can be excluded from the payload and promoted to an attribute at +once, as `tenant_id` is above. Values are reduced to JSON primitives the way +the payload is, so a `UUID` is stored as a string and parsed back by the +CloudEvent class that declares it. + +Names of CloudEvents core attributes (`id`, `type`, `source`, `subject`, +`time`, `data`, …) are rejected: every event carries those of its own. + +## What is recorded + +Every write through the repository is recorded, one event per affected row: + +| Method | Operation | +| --- | --- | +| `create`, `bulk_create`, `create_many`, `get_or_create` | `created` | +| `create_or_update`, `bulk_create_or_update` | `created` or `updated`, per row | +| `update_one`, `update_many`, `bulk_update` | `updated` | +| `delete_one`, `delete_many`, `remove` | `deleted` | + +The payload is the row after the write (before it, for a delete), reduced to +JSON primitives — UUIDs and timestamps become strings — and keyed by column +name. + +An upsert is not reported as a kind of its own: each written row is recorded +as what it turned out to be. `CreatedUpdatedMixin` gives `created_at` and +`updated_at` one shared timestamp per statement, so a row still carrying it — +`is_new` — was inserted, and one whose `updated_at` has moved on was updated. +A single `bulk_create_or_update()` that inserts some rows and updates others +therefore records `created` for the first and `updated` for the second. This +is why the mixin is required, and why an upsert through *any* repository +leaves `created_at` alone: a row was created when it was created, and +rewriting it would make an updated row look new. + +Building the events needs the written rows back, so `remove()` and +`delete_many()` use a RETURNING statement here instead of a bare `DELETE`, and +`bulk_update()` reads its rows back by the keys it matched on. On MySQL and +MariaDB, which have no `RETURNING`, the repository's usual fallback (select the +identities, write, re-fetch) supplies them, so nothing dialect-specific is +involved — the outbox works on every backend sqlargon supports. + +Only writes that go *through the repository* are recorded. Raw SQL, a +migration or another service writing to the table leaves no event behind. + +!!! note "Combining with soft delete" + `SoftDeleteRepository` rewrites `delete()` into an update raising the + tombstone. Combining the two repositories records a `deleted` event for a + soft delete, which is what the caller asked for, but the statement behind + it is an `UPDATE` — a consumer replaying the raw payload sees the row is + still there, with `tombstone` set. + +## Running the relay + +`OutboxRelay` polls for unpublished events and hands them to the publisher. +Use it as a long-running entrypoint, or start it in the background: + +```python +from contextlib import asynccontextmanager + +from fastapi import FastAPI + +from sqlargon.outbox import OutboxEvent, OutboxRelay + + +async def publish(event: OutboxEvent) -> None: + await broker.send(event.topic, event.data) + + +relay = OutboxRelay(publish) + + +@asynccontextmanager +async def lifespan(app: FastAPI): + async with relay.running(): # yields once the relay is live + yield + + +app = FastAPI(lifespan=lifespan) +``` + +`run()` also supports anyio's initialization barrier, so it can be embedded in +an existing task group with `await tg.start(relay.run)`. + +Delivery semantics: + +- Events are published **one at a time, in the order they were written** + (`created_at`, then `id`). An outbox that reorders its events is not much of + an outbox. +- A publisher that raises stops the batch, so a broker outage delays the + events behind the failed one rather than letting them overtake it. The + failure is logged (logger `sqlargon.outbox.relay`), recorded in + `last_error`, and retried with an exponential backoff. +- Delivery is **at least once**: an event published just before the process + dies is published again, because it had not been marked yet. Consumers must + be idempotent — the CloudEvents `id` is stable across redeliveries. +- After `max_attempts` failures an event stops being claimed and keeps its + `last_error` for inspection. +- A failing poll (a dropped connection, say) is logged and retried on the next + tick; the relay keeps running. + +`dispatch_once()` runs a single claim-and-publish pass and returns how many +events it published — handy in tests, or to drain the table from a script. + +### Multiple relays + +A claim locks the due rows with `SELECT ... FOR UPDATE SKIP LOCKED` and leases +them — `available_at` is pushed `lease` into the future and `attempts` bumped — +within the same transaction. Relays polling concurrently skip locked rows, and +a relay that dies mid-batch leaves nothing wedged: its lease expires and +another picks the events up. + +## Retention + +Published events are marked, not deleted, so the table doubles as an audit +trail. Deleting the ones older than `retention` (7 days by default) is +`purge()`, which the relay never runs on its own — polling and pruning keep +different schedules, and nothing is dropped unless you ask for it: + +```python +relay = OutboxRelay(publish, retention=timedelta(days=30)) +cron.task("0 3 * * *", "purge_outbox", relay.purge) +``` + +`purge()` takes no arguments and returns how many rows it deleted, so it works +as a cron task as it stands. It only ever deletes rows that have been +published. Skip registering it to keep every event. + +## Configuration + +`OutboxRelay` takes: + +- `publisher` — the coroutine function each claimed event is handed to. +- `repository` — the `OutboxEventRepository` to read events with; defaults to + a fresh one bound to the process-wide default database. Pass + `OutboxEventRepository().using(db=other_database)` to relay another + database. +- `poll_interval` — seconds between polls (default `1.0`). +- `batch_size` — maximum events claimed per poll (default `100`). +- `lease` — how long a claimed event stays invisible to other relays + (default 5 minutes). +- `max_attempts` — attempts after which an event is left alone (default `10`). +- `retry_backoff` / `max_retry_delay` — base and ceiling, in seconds, of the + exponential retry delay (defaults `2.0` and `300.0`). +- `retention` — how much history `purge()` keeps (default 7 days). + +## eventiq + +`sqlargon.integrations.eventiq` converts an event row into an +`eventiq.CloudEvent` and publishes it through a `Service`. It is the only +module that imports `eventiq`, so the dependency stays optional — install the +`sqlargon[eventiq]` extra to use it. + +```python +from sqlargon.integrations.eventiq import eventiq_publisher +from sqlargon.outbox import OutboxRelay + +relay = OutboxRelay(eventiq_publisher(service, source="orders")) +``` + +The mapping is direct: + +| `OutboxEvent` | `CloudEvent` | +| --- | --- | +| `id` | `id` | +| `created_at` | `time` | +| `topic` | `topic` (serialized as `subject`) | +| `type` | `type` | +| `source`, or the fallback given to the publisher | `source` | +| `data` | `data` | +| `attributes` | top-level attributes, one per key | +| `headers` | the broker's transport headers | + +`to_cloud_event(event, source=...)` does the conversion on its own, for code +that publishes through something other than a relay. + +### Publishing your own CloudEvent class + +Services usually publish a base class of their own rather than the plain +`CloudEvent` — one that declares the [extra attributes](#extra-attributes) +every event of theirs carries. Name it with `event_class` and the relay builds +and validates that class instead: + +```python +from uuid import UUID + +from eventiq import CloudEvent + + +class TenantEvent(CloudEvent): + tenant_id: UUID + + +relay = OutboxRelay(eventiq_publisher(service, event_class=TenantEvent)) +``` + +The row's `attributes` are passed to it as top-level attributes, so the +`tenant_id` stored as a string comes back parsed as a `UUID`, and an event +missing an attribute the class requires fails validation rather than being +published half-formed. The core CloudEvents attributes always come from the +row's own columns, so nothing stored in `attributes` can shadow them. +`to_cloud_event(event, event_class=TenantEvent)` returns the same class, +typed as such. + +## The outbox_events table + +Events are stored in the `outbox_events` table, registered on the shared +`Base` metadata when `sqlargon.outbox` is imported — `db.create_all()` creates +it, and an Alembic autogenerate pass picks it up (see +[Alembic Migrations](migrations.md)). + +| Column | Purpose | +| --- | --- | +| `id` | The CloudEvents `id`. | +| `created_at` | The CloudEvents `time`, and the dispatch order. | +| `topic`, `type`, `source` | The CloudEvents routing attributes. | +| `data` | The row snapshot, as JSON. | +| `attributes` | The extra CloudEvents attributes, as JSON. | +| `headers` | The broker's transport headers, as JSON. | +| `published_at` | `NULL` until published; set on success. | +| `available_at` | Not before which the event may be claimed. | +| `attempts` | Claims so far, used by `max_attempts`. | +| `last_error` | Why the last attempt failed. | + +Every column carries a server default as well as a client one, so a row +inserted without naming them is still complete. + +`OutboxEventRepository` exposes the table directly, for monitoring or for a +dispatcher of your own: `pending_count()`, `claim_pending()`, +`mark_published()`, `mark_failed()` and `purge()`. diff --git a/docs/reference/api.md b/docs/reference/api.md index ceacfd9..bd25dbc 100644 --- a/docs/reference/api.md +++ b/docs/reference/api.md @@ -10,10 +10,14 @@ Generated from the source. See [Usage](../usage.md) for a narrative introduction ::: sqlargon.repository.VersionedRepository +::: sqlargon.repository.AuditableRepository + ::: sqlargon.repository.DeletedRowExistsError ::: sqlargon.repository.ConcurrentModificationError +::: sqlargon.repository.AppendOnlyError + ::: sqlargon.functools.atomic ## Unit of work @@ -104,9 +108,29 @@ Generated from the source. See [Usage](../usage.md) for a narrative introduction ::: sqlargon.cron.CronTaskRepository -::: sqlargon.cron.validate_schedule +::: sqlargon.cron.utils.validate_schedule + +::: sqlargon.cron.utils.next_run_time + +## Outbox + +::: sqlargon.outbox.OutboxRepository + +::: sqlargon.outbox.OutboxEventRepository + +::: sqlargon.outbox.OutboxRelay -::: sqlargon.cron.next_run_time +::: sqlargon.outbox.OutboxEvent + +::: sqlargon.outbox.OutboxConfig + +::: sqlargon.outbox.Operation + +::: sqlargon.outbox.format_topic + +::: sqlargon.integrations.eventiq.to_cloud_event + +::: sqlargon.integrations.eventiq.eventiq_publisher ## ORM and types @@ -114,6 +138,8 @@ Generated from the source. See [Usage](../usage.md) for a narrative introduction ::: sqlargon.mixins +::: sqlargon.audit + ::: sqlargon.types.uuid ::: sqlargon.types.datetime diff --git a/docs/reference/dialects.md b/docs/reference/dialects.md index daf81c3..d60ec2b 100644 --- a/docs/reference/dialects.md +++ b/docs/reference/dialects.md @@ -11,12 +11,12 @@ is isolated. Builders declare what they support as an `Option` flag, checked with `db.query_builder.supports(...)`: -| Dialect | `RETURNING` | `CONFLICTS` | `LOCKS` | -| --- | --- | --- | --- | -| `postgresql` | ✅ | ✅ | ✅ `pg_advisory_lock` | -| `sqlite` | ✅ (SQLite ≥ 3.35) | ✅ | ❌ | -| `mysql` | ❌ | ✅ | ✅ `GET_LOCK` | -| anything else | ❌ | ❌ | ❌ | +| Dialect | `RETURNING` | `CONFLICTS` | `LOCKS` | `VECTORS` | `FULL_TEXT` | +| --- | --- | --- | --- | --- | --- | +| `postgresql` | ✅ | ✅ | ✅ `pg_advisory_lock` | ✅ pgvector | ✅ `ts_rank` | +| `sqlite` | ✅ (SQLite ≥ 3.35) | ✅ | ❌ | ✅ sqlite-vector | ❌ | +| `mysql` | ❌ | ✅ | ✅ `GET_LOCK` | ❌ | ❌ | +| anything else | ❌ | ❌ | ❌ | ❌ | ❌ | ```python from sqlargon.query_builder import Option @@ -106,6 +106,13 @@ Methods: `select`, `insert`, `update`, `delete`, `filter`, `count`, `page`, `loc and `get_lock_pair`. `insert`, `update` and `delete` take `return_results=True` to append a `RETURNING` clause for the whole table. +The search hooks — `vector_search`, `vector_distance`, `vector_init`, `text_search` and +`rrf_search` — are the same idea for [vector search](../vectors.md): the base class refuses +them with `UnsupportedDialectError`, and the two backends that can search express it in +shapes with nothing in common. PostgreSQL orders by a pgvector operator; SQLite joins the +table valued scan sqlite-vector exposes, because it has no scalar distance function at all. +Keeping both behind one hook is what lets `VectorRepository.search()` be portable. + ## Adding a dialect Subclass `QueryBuilder`, declare the supported options and override what differs — the base diff --git a/docs/reference/types.md b/docs/reference/types.md index 96b385b..8073d1a 100644 --- a/docs/reference/types.md +++ b/docs/reference/types.md @@ -91,18 +91,69 @@ await repo.list(Document.meta.json_value("owner") == "john") | `has_any_key([...])` | `?\|` | `JSON_CONTAINS_PATH(..., 'one', ...)` | `EXISTS` over `json_each` | | `has_all_keys([...])` | `?&` | `JSON_CONTAINS_PATH(..., 'all', ...)` | `json_each` self-join | | `json_value(key)` | `->>` | `JSON_EXTRACT` | `JSON_EXTRACT` | +| `get(key)` | `->` | `JSON_EXTRACT` | `JSON_EXTRACT` | +| `has_key(key)` | `?` | `JSON_CONTAINS_PATH(..., 'one', ...)` | `JSON_TYPE(...) IS NOT NULL` | +| `array_length()` | `JSONB_ARRAY_LENGTH` | `JSON_LENGTH` | `JSON_ARRAY_LENGTH` | +| `keys()` | `JSONB_OBJECT_KEYS` + `JSONB_AGG` | `JSON_KEYS` | `JSON_GROUP_ARRAY` over `json_each` | !!! warning "`has_any_key` / `has_all_keys` are portable over arrays, not objects" Use them to test membership in a JSON **array** — that is the one meaning all three dialects agree on. Against a JSON **object** they diverge: PostgreSQL and MySQL test the object's *keys*, while the SQLite fallback tests the *values* produced by `json_each`. - To query a key of an object portably, use `json_value(key)` instead. + To test a key of an object portably, use `has_key(key)`, which addresses object keys on + every dialect. -The underlying function elements — `json_contains`, `json_has_any_key`, `json_has_all_keys` -and `json_value` — are importable from `sqlargon.types.json` for use outside a `JSON` -column. `has_any_key` and `has_all_keys` require string keys and raise `ValueError` -otherwise. +### Mutating a document server-side + +The mutation operators rewrite a document in the `UPDATE` itself, so a single key can be +changed without reading the row into Python and writing it back — no lost update, one +round trip: + +```python +await repo.update({Document.meta: Document.meta.set_key("owner", "john")}).execute() +await repo.update({Document.meta: Document.meta.update({"owner": "john", "hits": 0})}).execute() +await repo.update({Document.meta: Document.meta.remove_key("owner")}).execute() +``` + +| Operator | PostgreSQL | MySQL | SQLite | +| --- | --- | --- | --- | +| `set_key(key, value)` | `\|\|` | `JSON_SET` | `JSON_SET` | +| `update({...})` | `\|\|` | `JSON_SET` | `JSON_SET` | +| `remove_key(*keys)` | `-` over `text[]` | `JSON_REMOVE` | `JSON_REMOVE` | +| `insert_key(key, value)` | `\|\|`, patch on the left | `JSON_INSERT` | `JSON_INSERT` | +| `replace_key(key, value)` | `JSONB_SET(..., false)` | `JSON_REPLACE` | `JSON_REPLACE` | +| `array_append(value)` | `\|\|` + `JSONB_BUILD_ARRAY` | `JSON_ARRAY_APPEND` | `JSON_INSERT(..., '$[#]', ...)` | + +`insert_key` only writes a key that is **absent**; `replace_key` only one already +**present**. Every mutation returns a JSON expression, so they nest: + +```python +Document.meta.update({"c": 3}).remove_key("a") +``` + +!!! warning "What the mutation operators do not smooth over" + + - **`NULL` in, `NULL` out.** `JSONB_SET` and `JSON_SET` both return `NULL` for a `NULL` + document, and these operators match that rather than coalescing to `{}`. Give the + column a `server_default` of `'{}'` if you need a document to always be there. + - **Objects only.** The `JSON_SET` family addresses `$."key"`, so `set_key`, + `update`, `remove_key`, `insert_key` and `replace_key` assume the document is an + object. Use `array_append` for arrays. + - **Top-level keys only.** There are no nested paths or array indices; a key is always + one level down. + - **`update` is a shallow merge.** A top-level key is replaced wholesale, not merged + into recursively — the semantics of PostgreSQL's `||`. Deep merge-patch + (`JSON_MERGE_PATCH`, `json_patch`) is deliberately absent: PostgreSQL has no builtin + for it. + - **`array_length` is portable over arrays only.** Given an object PostgreSQL raises, + SQLite answers 0 and MySQL answers 1. + +The underlying function elements — `json_contains`, `json_has_any_key`, `json_has_all_keys`, +`json_value`, `json_get`, `json_has_key`, `json_array_length`, `json_keys`, `json_update`, +`json_set_key`, `json_remove_key`, `json_insert_key`, `json_replace_key` and +`json_array_append` — are importable from `sqlargon.types.json` for use outside a `JSON` +column. The key operators require string keys and raise `ValueError` otherwise. ## Pydantic-validated columns @@ -144,6 +195,9 @@ Both accept `sa_column_type=` to store in something other than `JSON` (e.g. `sa. | `VersionedMixin` | *(abstract marker — no columns)* | | | `UUIDVersionedMixin` | `version_id` — `GUID`, `uuid4` default, `GenerateUUID()` server default | `__mapper_args__` with `version_id_col` + UUID generator | | `XminVersionedMixin` | `xmin` — PostgreSQL system column, `String`, `system=True`, `FetchedValue()` | `__mapper_args__` with `version_id_col` + `version_id_generator=False` | +| `AuditableMixin` | *(abstract marker — no columns; extends `SoftDeleteMixin` and `VersionedMixin`)* | `audit_key()`, `latest_version()`, `is_latest()` | +| `IntegerAuditableMixin` | `version` — `Integer` primary key, starting at 1 | successor is `version + 1` | +| `UUIDAuditableMixin` | `version` — `GUID` primary key, `uuid7` default, `GenerateUUIDV7()` server default | successor is a fresh UUIDv7 | ```python from sqlargon.mixins import CreatedUpdatedMixin, SoftDeleteMixin, UUIDV7ModelMixin @@ -213,3 +267,22 @@ Both set `__mapper_args__` with `version_id_col`, enabling SQLAlchemy's ORM-leve versioning when using `AsyncSession` directly. `VersionedModel` is the matching type variable, bound to `VersionedBase`. See [Versioned models](../usage.md#versioned-models) for the repository API. + +`AuditableMixin` is an abstract marker too — use `IntegerAuditableMixin` or +`UUIDAuditableMixin`, or the bases combining them with `Base`: + +```python +from sqlargon import AuditableBase, UUIDAuditableBase + + +class Article(UUIDModelMixin, AuditableBase): # versions 1, 2, 3 ... + title: Mapped[str] = mapped_column(sa.Unicode(255)) + + +class Draft(UUIDModelMixin, UUIDAuditableBase): # UUIDv7 versions + title: Mapped[str] = mapped_column(sa.Unicode(255)) +``` + +`AnyAuditableBase` is the abstract base both share and `AuditableModel` the matching type +variable. Pair either with [`AuditableRepository`](../auditable.md), which appends a new +version instead of updating a row. diff --git a/docs/usage.md b/docs/usage.md index 288ea2b..b9f2504 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -182,7 +182,9 @@ row on the `on_` columns — one round trip for the whole batch. `insert(..., ignore_conflicts=True)`, `upsert(...)` and the bulk helpers derive their `ON CONFLICT` clause from the `on_conflict` property, which defaults to the model's primary -key as `index_elements` and every remaining column in `set_`. Override it per repository: +key as `index_elements` and every remaining column in `set_` — bar `created_at`, which a +`CreatedUpdatedMixin` row keeps from the insert that created it, so `is_new` stays +truthful after an upsert. Override it per repository: ```python from sqlargon.typing import OnConflictOptions diff --git a/docs/vectors.md b/docs/vectors.md new file mode 100644 index 0000000..fa4e941 --- /dev/null +++ b/docs/vectors.md @@ -0,0 +1,212 @@ +# Vector Search + +`sqlargon.vectors` stores embeddings and searches them by similarity. The +column type, the model mixins and the repositories are all separate, so a model +takes only the pieces it needs — an embedding alone, or an embedding beside +text, JSON attributes and a collection. + +It requires the `sqlargon[vectors]` extra (`pgvector`), plus +`sqlargon[vectors-sqlite]` (`sqliteai-vector`) to search on SQLite. + +```python +from sqlalchemy.orm import declared_attr + +from sqlargon import Database +from sqlargon.mixins import UUIDV7ModelMixin +from sqlargon.vectors import EmbeddingBase, VectorRepository, init_vectors + + +class Note(UUIDV7ModelMixin, EmbeddingBase): + __vector_dim__ = 384 + + @declared_attr.directive + def __table_args__(cls): + return (cls.embedding_index(),) + + +class NoteRepository(VectorRepository[Note]): + pass + + +db = Database.from_env() +await init_vectors(db) # before create_all() +await db.create_all() + +await NoteRepository().create(embedding=[0.1] * 384) +nearest = await NoteRepository().search([0.1] * 384, limit=5) +``` + +## Choosing columns + +Each mixin adds one column, its index helper and the expressions that column +needs. Combine only the ones the model wants. + +| Mixin | Column | Configure with | Index helper | +| --- | --- | --- | --- | +| `EmbeddingMixin` | `embedding` | `__vector_dim__`, `__vector_distance__` | `embedding_index()` | +| `TextMixin` | `text` | `__text_regconfig__` | `text_index()` | +| `AttributesMixin` | `attributes` | — | `attributes_index()` | +| `VectorCollectionMixin` | `collection_id` | `__collection_table__` | — | + +Four abstract bases pre-compose them: `EmbeddingBase` (embedding only), +`TextBase` (text only), `TextEmbeddingBase` (both) and `VectorDocument` +(everything, plus `UUIDV7ModelMixin` and `CreatedUpdatedMixin`). Anything else +is composed directly: + +```python +class Chunk(UUIDV7ModelMixin, AttributesMixin, EmbeddingBase): + """An embedding and JSON attributes -- no text, no collection.""" + + __vector_dim__ = 1536 +``` + +Index DDL only runs on PostgreSQL, so the same model is portable: SQLite needs +no index, and on MySQL the embedding degrades to JSON storage with no search. + +`VectorDocument` points `collection_id` at `VectorCollection`, a concrete model +this package declares — importing it registers the `vector_collection` table +with the shared metadata, so `create_all()` creates it. Point +`__collection_table__` at a table of your own to group documents differently. + +## Choosing a repository + +Each repository requires the mixins its search needs and raises `TypeError` at +subclass time when the model lacks one. + +| Repository | Model needs | Adds | +| --- | --- | --- | +| `VectorRepository` | `EmbeddingMixin` | `search()` | +| `TextSearchRepository` | `TextMixin` | `text_search()` | +| `HybridVectorRepository` | both | both, plus `rrf_search()` | + +## Similarity and hybrid search + +`search()` returns the models nearest to a vector, most similar first. Extra +positional expressions and keyword equalities narrow it exactly as `where()` +does, which is all hybrid search is — an ordinary `WHERE` beside the ordering: + +```python +found = await documents.search( + embedding, + Document.attributes_contain({"lang": "en"}), + Document.text.like("%apple%"), + collection_id=collection.id, + limit=10, +) +``` + +The filter is applied before the limit, so filtering never returns fewer rows +than it should. `with_distance=True` returns `(model, distance)` pairs: + +```python +for document, distance in await documents.search(embedding, with_distance=True): + print(document.text, distance) +``` + +Attributes need no repository support — `attributes_contain()` is a predicate, +and the `JSON` column type also offers `contains()`, `has_any_key()` and +`json_value()` through its comparator. + +## Distance metrics + +`__vector_distance__` sets the metric a model's index is built for and its +searches default to. Distances always sort ascending, so the nearest row comes +first whichever metric is chosen. + +| Metric | pgvector operator | Index opclass | +| --- | --- | --- | +| `DistanceMetric.COSINE` (default) | `<=>` | `vector_cosine_ops` | +| `DistanceMetric.L2` | `<->` | `vector_l2_ops` | +| `DistanceMetric.DOT` | `<#>` | `vector_ip_ops` | +| `DistanceMetric.L1` | `<+>` | `vector_l1_ops` | + +`DOT` is the *negative* inner product, which is what keeps ascending order +meaningful for it. + +On PostgreSQL a single query can override the metric with `search(..., +metric=DistanceMetric.L2)`, though it will not use an index built for another +one. SQLite fixes the metric per column, so overriding it there raises +`UnsupportedDialectError`. + +Outside a repository the comparator gives the same expressions directly, for +instance `Note.embedding.cosine_distance(vector)`. They compile on PostgreSQL +only; on other dialects they raise `UnsupportedDialectError` rather than +emitting SQL the server would reject. + +## Full text and reciprocal rank fusion + +`TextSearchRepository.text_search()` ranks rows by `ts_rank` over +`to_tsvector(__text_regconfig__, text)` — the same expression `text_index()` +builds, so the index applies. `HybridVectorRepository.rrf_search()` fuses that +ranking with the vector one by reciprocal rank fusion: it ranks the `candidates` +nearest rows and the `candidates` best text matches, then scores each row +`sum(1 / (k + rank))` over the rankings it appears in, so a row both agree on +outranks one that only either found. + +```python +for document, score in await documents.rrf_search(embedding, "red apple", limit=10): + print(score, document.text) +``` + +Both are PostgreSQL only and raise `UnsupportedDialectError` elsewhere. + +## Setting a backend up + +`init_vectors(db)` prepares a database and has to run before `create_all()` — +a `VECTOR` column cannot be declared before the type exists. On PostgreSQL it +issues `CREATE EXTENSION IF NOT EXISTS vector`; on SQLite it registers the +loadable extension on the engine's pool, so register it at startup, before any +query, because connections checked out earlier never get it. + +Applications that manage their schema with alembic create the extension in a +migration instead, ahead of the table: + +```python +def upgrade() -> None: + op.execute("CREATE EXTENSION IF NOT EXISTS vector") +``` + +Index DDL belongs in the migration too, since `create_all()` is not what builds +the schema there. + +## Backend support + +| | PostgreSQL | SQLite | MySQL / MariaDB | +| --- | --- | --- | --- | +| Storage | `VECTOR(n)` | float32 `BLOB` | JSON | +| `search()` | distance operator | `vector_full_scan` join | not supported | +| `text_search()` / `rrf_search()` | yes | no | no | +| Indexes | HNSW, GIN | none needed | none | + +SQLite goes through sqlite-vector, which differs enough to be worth knowing: +it has no scalar distance function, so searches join its table valued scan +rather than ordering by an expression; the metric is fixed per column when the +column is declared to it; and vectors are plain float32 blobs. The declaration +happens on first search per connection and is remembered for that connection. + +Values read back are `list[float]` on every backend. + +Two things are deliberately left out for now: sqlite-vector's quantized scan +(`vector_quantize_scan`), which needs a quantization lifecycle of its own, and +fusing SQLite FTS5 with vector ranking. + +## Where the statements come from + +The repositories build no SQL of their own. Each statement comes from the +dialect's [`QueryBuilder`](reference/dialects.md), which is what makes the two +backends' very different shapes interchangeable behind one `search()`: + +| Hook | Answers | +| --- | --- | +| `vector_search()` | the whole similarity query, filters included | +| `vector_distance()` | the ordering expression, where the backend has one | +| `vector_init()` | the per-connection declaration, or `None` when unneeded | +| `text_search()` | the ranked full text query | +| `rrf_search()` | the fused query | + +A repository asks `supports(Option.VECTORS)` or `Option.FULL_TEXT` before +building anything, so an unsupported backend raises +`UnsupportedDialectError` naming the dialect instead of emitting SQL the +server would reject. Teaching sqlargon another vector backend therefore means +overriding these hooks on that dialect's builder and claiming the options — +no repository change at all. diff --git a/mkdocs.yaml b/mkdocs.yaml index 62ca89e..f5d93fd 100644 --- a/mkdocs.yaml +++ b/mkdocs.yaml @@ -24,6 +24,9 @@ nav: - "Database Routing": routing.md - "Pagination": pagination.md - "Cron": cron.md + - "Outbox": outbox.md + - "Vector Search": vectors.md + - "Auditable Models": auditable.md - "Examples": examples.md - "Reference": - "Column Types & Mixins": reference/types.md diff --git a/pyproject.toml b/pyproject.toml index 30cc370..41816d6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -21,9 +21,15 @@ sqlite = ["aiosqlite>=0.19.0,<1"] mysql = ["asyncmy>=0.2.11"] pagination = ["sqlakeyset>=2.0.1716332987,<3"] cron = ["croniter>=2.0,<7", "anyio>=4.0,<5"] +outbox = ["anyio>=4.0,<5"] +eventiq = ["eventiq>=1.1.14,<2"] opentelemetry = ["opentelemetry-instrumentation-sqlalchemy"] standard = ["asyncpg<1.0", "aiosqlite>=0.19.0,<1", "sqlakeyset>=2.0.1716332987,<3", "croniter>=2.0,<7", "anyio>=4.0,<5", "opentelemetry-instrumentation-sqlalchemy"] +vectors = [ + "pgvector>=0.5.0", +] +vectors-sqlite = ["sqliteai-vector>=1.0.0,<2"] [dependency-groups] dev = [ @@ -31,6 +37,7 @@ dev = [ { include-group = "test" }, { include-group = "e2e" }, { include-group = "docs" }, + "eventiq>=1.1.14", "fastapi>=0.138.0", "greenlet", ] @@ -39,14 +46,17 @@ e2e = [ "testcontainers>=4.15.0", # asyncmy needs it for the caching_sha2_password auth of MySQL 8 "cryptography", + # the loadable extension the SQLite vector search runs on + "sqliteai-vector>=1.0.0,<2", ] test = [ + "anyio>=4.14.1", "pytest", "pytest-cov", "pytest-sugar", "pytest-repeat", - "pytest-asyncio", + "pytest-xdist>=3.6", "httpx>=0.28.1", "pytest-timeout>=2.4.0", ] @@ -76,9 +86,14 @@ build-backend = "uv_build" [tool.pytest.ini_options] -addopts = "--cov=./sqlargon --count=3" +# loadgroup keeps every test of an e2e backend on the worker that started its +# container, so the backends run in parallel but each one boots only once +addopts = "--cov=./sqlargon --count=3 --numprocesses=auto --dist=loadgroup" testpaths = ["./tests"] -asyncio_mode = "auto" +anyio_mode = "auto" +# the anyio plugin only runs an async fixture for a test that pulls the backend +# in, which a sync test requesting an autouse async fixture would not do +usefixtures = ["anyio_backend"] timeout = 20 markers = [ "e2e: runs against a real database backend; enabled with --e2e", @@ -144,6 +159,8 @@ classmethod-decorators = [ [tool.ruff.lint.per-file-ignores] "sqlargon/*" = ["PLC0415"] "sqlargon/types/*" = ["ARG002"] +# the search hooks name the arguments their dialect overrides act on +"sqlargon/query_builder.py" = ["ARG002"] "tests/*" = ["S101", "ANN001", "ANN002", "ANN003", "ANN201", "ANN202", "SLF001", "PLR2004", "ARG002"] # a fixture requested for its side effect only, and testcontainers imported # lazily so an ordinary run never pulls it in diff --git a/sqlargon/__init__.py b/sqlargon/__init__.py index bab5718..5bd6da3 100644 --- a/sqlargon/__init__.py +++ b/sqlargon/__init__.py @@ -1,21 +1,36 @@ from importlib.metadata import version +from .audit import latest_relationship, version_foreign_key, version_mapped_column from .cluster import AnyDatabase, DatabaseCluster from .database import BaseDatabase, Database, ReadOnlyDatabase, ReadOnlyError from .functools import atomic -from .mixins import UUIDVersionedMixin, VersionedMixin, XminVersionedMixin +from .mixins import ( + AuditableMixin, + IntegerAuditableMixin, + UUIDAuditableMixin, + UUIDVersionedMixin, + VersionedMixin, + XminVersionedMixin, +) from .orm import ( + AnyAuditableBase, + AnyVersionedBase, + AuditableBase, + AuditableModel, Base, Model, ORMModel, SoftDeleteBase, SoftDeleteModel, + UUIDAuditableBase, VersionedBase, VersionedModel, XminVersionedBase, ) from .registry import get_default_database, set_default_database from .repository import ( + AppendOnlyError, + AuditableRepository, ConcurrentModificationError, DeletedRowExistsError, SoftDeleteRepository, @@ -41,7 +56,14 @@ __all__ = [ "AbstractUnitOfWork", + "AnyAuditableBase", "AnyDatabase", + "AnyVersionedBase", + "AppendOnlyError", + "AuditableBase", + "AuditableMixin", + "AuditableModel", + "AuditableRepository", "Base", "BaseDatabase", "ConcurrentModificationError", @@ -49,6 +71,7 @@ "DatabaseCluster", "DefaultRouter", "DeletedRowExistsError", + "IntegerAuditableMixin", "Model", "ModelRouter", "ORMModel", @@ -65,6 +88,8 @@ "SoftDeleteBase", "SoftDeleteModel", "SoftDeleteRepository", + "UUIDAuditableBase", + "UUIDAuditableMixin", "UUIDVersionedMixin", "VersionedBase", "VersionedMixin", @@ -75,8 +100,11 @@ "__version__", "atomic", "get_default_database", + "latest_relationship", "read_only", "set_default_database", "use_context", "using", + "version_foreign_key", + "version_mapped_column", ] diff --git a/sqlargon/audit.py b/sqlargon/audit.py new file mode 100644 index 0000000..12e0da0 --- /dev/null +++ b/sqlargon/audit.py @@ -0,0 +1,125 @@ +"""Relating other tables to an append-only, versioned model. + +A row of an :class:`~sqlargon.orm.AnyAuditableBase` model is one version of +an entity, so a reference to it has to say *which* version it means. The +helpers here cover the two answers: + +* :func:`version_mapped_column` and :func:`version_foreign_key` pin a child + to one exact version, under a real composite foreign key; +* :func:`latest_relationship` follows the entity forward, resolving to + whichever version is newest at read time. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +import sqlalchemy as sa +from sqlalchemy.orm import declared_attr, mapped_column, relationship + +if TYPE_CHECKING: + from sqlalchemy.orm import MappedColumn, Relationship + + from sqlargon.mixins import AuditableMixin + +__all__ = ["latest_relationship", "version_foreign_key", "version_mapped_column"] + + +def _remote_columns( + target: type[AuditableMixin], columns: tuple[str, ...] +) -> tuple[str, ...]: + remote = target.audit_key() + if len(columns) != len(remote): + msg = ( + f"{target.__name__} is identified by {remote}, so {len(remote)} " + f"local column(s) are needed, got {len(columns)}: {columns}" + ) + raise ValueError(msg) + return remote + + +def version_mapped_column( + target: type[AuditableMixin], **kwargs: Any +) -> MappedColumn[Any]: + """A column holding a version of ``target``, typed to match it. + + The child never has to know which versioning strategy the parent uses:: + + class Comment(UUIDModelMixin, Base): + article_id: Mapped[UUID] = mapped_column(GUID()) + article_version = version_mapped_column(Article) + """ + return mapped_column(target.__table__.c.version.type, **kwargs) + + +def version_foreign_key( + target: type[AuditableMixin], *columns: str, **kwargs: Any +) -> sa.ForeignKeyConstraint: + """A composite foreign key pinning ``columns`` to one version of ``target``. + + ``columns`` names the local columns mirroring the entity key and the + version, in that order:: + + class Comment(UUIDModelMixin, Base): + article_id: Mapped[UUID] = mapped_column(GUID()) + article_version = version_mapped_column(Article) + + __table_args__ = ( + version_foreign_key(Article, "article_id", "article_version"), + ) + + article: Mapped[Article] = relationship() + + Because the key is real, the relationship needs no ``primaryjoin`` and is + writable: assigning an ``Article`` fills both columns. It also stops + :meth:`~sqlargon.repository.AuditableRepository.purge` from removing a + version something still points at, unless declared ``ondelete="CASCADE"``. + """ + remote = _remote_columns(target, columns[:-1]) + table = target.__table__ + return sa.ForeignKeyConstraint( + list(columns), [table.c[name] for name in (*remote, "version")], **kwargs + ) + + +def latest_relationship( + target: type[AuditableMixin], *columns: str, uselist: bool = False, **kwargs: Any +) -> Any: + """A relationship resolving to the newest version of ``target``. + + ``columns`` names the local columns mirroring the entity key -- no + version column, and no foreign key, since the primary key of ``target`` + holds a version this child deliberately does not pin:: + + class Tag(UUIDModelMixin, Base): + article_id: Mapped[UUID] = mapped_column(GUID()) + + article = latest_relationship(Article, "article_id") + + The join carries :meth:`~sqlargon.mixins.AuditableMixin.is_latest`, so the + child follows the entity forward as versions are appended. Having no + foreign key to write back through, it is necessarily ``viewonly``. + """ + remote = _remote_columns(target, columns) + + # declared_attr only to reach the class being declared: the join needs its + # columns, and a helper called from a class body has no other handle on it + @declared_attr + def _latest(cls: Any) -> Relationship[Any]: + local = [getattr(cls, name) for name in columns] + return relationship( + target, + primaryjoin=sa.and_( + *( + column == getattr(target, name) + for column, name in zip(local, remote, strict=True) + ), + target.is_latest(), + ), + foreign_keys=local, + viewonly=True, + uselist=uselist, + **kwargs, + ) + + return _latest diff --git a/sqlargon/cron/manager.py b/sqlargon/cron/manager.py index 45e8795..36b60b5 100644 --- a/sqlargon/cron/manager.py +++ b/sqlargon/cron/manager.py @@ -6,7 +6,7 @@ from dataclasses import dataclass, replace from functools import partial from inspect import iscoroutinefunction -from typing import TYPE_CHECKING, Any, TypeGuard, TypeVar +from typing import TYPE_CHECKING, Any, TypeGuard, TypeVar, overload import anyio from anyio import TASK_STATUS_IGNORED, CapacityLimiter, to_thread @@ -71,10 +71,11 @@ class _Registration: class Cron: """Database-backed cron scheduler for a single namespace. - Declarative mode: decorate functions with :meth:`task` and call - :meth:`run` (or :meth:`sync`); declared tasks are created or updated and - tasks no longer declared are deleted. Imperative mode: manage schedules - at runtime with :meth:`schedule` and :meth:`unschedule`. + Declarative mode: declare functions with :meth:`task`, as a decorator or + a plain call, and call :meth:`run` (or :meth:`sync`); declared tasks are + created or updated and tasks no longer declared are deleted. Imperative + mode: manage schedules at runtime with :meth:`schedule` and + :meth:`unschedule`. Multiple instances may run the same namespace concurrently; due tasks are claimed with ``FOR UPDATE SKIP LOCKED`` so each run executes once. @@ -113,21 +114,41 @@ def limiter(self) -> CapacityLimiter: self._limiter = CapacityLimiter(self.max_concurrency) return self._limiter + @overload def task( - self, schedule: str | None = None, *, name: str | None = None - ) -> Callable[[F], F]: - """Register the decorated function; with ``schedule`` it becomes a - declarative task reconciled by :meth:`sync`.""" + self, schedule: str | None = ..., name: str | None = ... + ) -> Callable[[F], F]: ... - def decorator(func: F) -> F: - task_name = self.register(func, name=name) + @overload + def task(self, schedule: str | None, name: str | None, func: F) -> F: ... + + @overload + def task(self, schedule: str | None = ..., *, func: F) -> F: ... + + def task( + self, + schedule: str | None = None, + name: str | None = None, + func: TaskFunc | None = None, + ) -> Any: + """Register a function; with ``schedule`` it becomes a declarative + task reconciled by :meth:`sync`. + + Used as a decorator when ``func`` is left out, and called outright + otherwise -- for a function that is not yours to decorate:: + + cron.task("0 3 * * *", "purge_outbox", relay.purge) + """ + + def decorator(target: F) -> F: + task_name = self.register(target, name=name) if schedule is not None: self._registry[task_name] = replace( self._registry[task_name], schedule=validate_schedule(schedule) ) - return func + return target - return decorator + return decorator if func is None else decorator(func) def register(self, func: TaskFunc, *, name: str | None = None) -> str: """Make ``func`` executable by this scheduler and return its name. diff --git a/sqlargon/dialects/postgres.py b/sqlargon/dialects/postgres.py index 336d914..8e24068 100644 --- a/sqlargon/dialects/postgres.py +++ b/sqlargon/dialects/postgres.py @@ -10,8 +10,11 @@ from sqlargon.query_builder import Option, QueryBuilder if TYPE_CHECKING: - from sqlalchemy.sql._typing import _DMLTableArgument + from collections.abc import Sequence + from sqlalchemy.sql._typing import _ColumnExpressionArgument, _DMLTableArgument + + from sqlargon.types.vector import DistanceMetric from sqlargon.typing import OnConflict, Values INT64_SIZE = 2**63 - 1 @@ -25,7 +28,13 @@ def _key_to_int(key: str) -> int: class PostgresqlQueryBuilder(QueryBuilder): - supported_options = Option.RETURNING | Option.CONFLICTS | Option.LOCKS + supported_options = ( + Option.RETURNING + | Option.CONFLICTS + | Option.LOCKS + | Option.VECTORS + | Option.FULL_TEXT + ) _lock_query = sa.text("SELECT pg_advisory_lock(:key)") _unlock_query = sa.text("SELECT pg_advisory_unlock(:key)") @@ -58,6 +67,128 @@ def _insert( assert_never(on_conflict.do) return query + def vector_distance( + self, + model: Any, + embedding: Sequence[float], + metric: DistanceMetric | None = None, + ) -> sa.ColumnElement[float]: + """A pgvector distance operator between the column and ``embedding``.""" + from sqlargon.types.vector import distance_for + + element = distance_for(metric or model.__vector_distance__) + return element(model.embedding, self.query_vector(model, embedding)) + + def vector_search( + self, + model: Any, + embedding: Sequence[float], + *filters: _ColumnExpressionArgument[bool], + limit: int, + metric: DistanceMetric | None = None, + ) -> sa.Select[Any]: + distance = self.vector_distance(model, embedding, metric) + return ( + sa.select(model, distance.label("distance")) + .where(*filters) + .order_by(distance) + .limit(limit) + ) + + def text_score(self, model: Any, query: str) -> sa.ColumnElement[float]: + """How well the model's text matches ``query``; higher is better.""" + return sa.func.ts_rank(model.text_document(), model.text_query(query)) + + def text_match(self, model: Any, query: str) -> sa.ColumnElement[bool]: + """Whether the model's text matches ``query`` at all.""" + return model.text_document().op("@@")(model.text_query(query)) + + def text_search( + self, + model: Any, + query: str, + *filters: _ColumnExpressionArgument[bool], + limit: int, + ) -> sa.Select[Any]: + score = self.text_score(model, query) + return ( + sa.select(model, score.label("score")) + .where(self.text_match(model, query), *filters) + .order_by(score.desc()) + .limit(limit) + ) + + def _rank_cte( + self, + model: Any, + order_by: sa.ColumnElement[Any], + filters: Sequence[_ColumnExpressionArgument[bool]], + *, + candidates: int, + name: str, + ) -> sa.CTE: + """The ``candidates`` best rows by ``order_by``, numbered from one.""" + return ( + sa.select( + self.identity_column(model).label("id"), + sa.func.row_number().over(order_by=order_by).label("rank"), + ) + .where(*filters) + .order_by(order_by) + .limit(candidates) + .cte(name) + ) + + def rrf_search( # noqa: PLR0913 -- keyword-only tuning knobs, each with a default + self, + model: Any, + embedding: Sequence[float], + query: str, + *filters: _ColumnExpressionArgument[bool], + k: int = 60, + limit: int = 10, + candidates: int = 50, + ) -> sa.Select[Any]: + vector_rank = self._rank_cte( + model, + self.vector_distance(model, embedding), + filters, + candidates=candidates, + name="vector_candidates", + ) + text_rank = self._rank_cte( + model, + self.text_score(model, query).desc(), + (self.text_match(model, query), *filters), + candidates=candidates, + name="text_candidates", + ) + score = ( + sa.func.coalesce(1.0 / (k + vector_rank.c.rank), 0.0) + + sa.func.coalesce(1.0 / (k + text_rank.c.rank), 0.0) + ).label("score") + # a full outer join so a row either ranking alone found still scores + fused = ( + sa.select( + sa.func.coalesce(vector_rank.c.id, text_rank.c.id).label("id"), score + ) + .select_from( + sa.join( + vector_rank, + text_rank, + vector_rank.c.id == text_rank.c.id, + full=True, + ) + ) + .subquery("rrf") + ) + return ( + sa.select(model, fused.c.score) + .join(fused, self.identity_column(model) == fused.c.id) + .order_by(fused.c.score.desc()) + .limit(limit) + ) + def lock(self, key: str) -> sa.TextClause: int_key = _key_to_int(key) return self._lock_query.bindparams(key=int_key) diff --git a/sqlargon/dialects/sqlite.py b/sqlargon/dialects/sqlite.py index e606b5a..c84da86 100644 --- a/sqlargon/dialects/sqlite.py +++ b/sqlargon/dialects/sqlite.py @@ -3,24 +3,36 @@ import sqlite3 from typing import TYPE_CHECKING, Any +import sqlalchemy as sa from sqlalchemy.dialects.sqlite import Insert, insert from typing_extensions import assert_never -from sqlargon.query_builder import Option, QueryBuilder +from sqlargon.query_builder import Option, QueryBuilder, UnsupportedDialectError if TYPE_CHECKING: - from sqlalchemy.sql._typing import _DMLTableArgument + from collections.abc import Sequence + from sqlalchemy.sql._typing import _ColumnExpressionArgument, _DMLTableArgument + + from sqlargon.types.vector import DistanceMetric from sqlargon.typing import OnConflict, Values -_SQLITE_OPTIONS = Option.CONFLICTS +_SQLITE_OPTIONS = Option.CONFLICTS | Option.VECTORS if sqlite3.sqlite_version > "3.35": _SQLITE_OPTIONS |= Option.RETURNING class SQLiteQueryBuilder(QueryBuilder): + """Query builder for SQLite, searching vectors through sqlite-vector. + + That extension exposes no scalar distance function, so a search joins + the table valued scan it does expose rather than ordering by an + expression, and the column has to be declared to it per connection -- + see :meth:`vector_init`. + """ + supported_options = _SQLITE_OPTIONS def excluded(self, table: _DMLTableArgument) -> Any: @@ -53,3 +65,57 @@ def _insert( else: assert_never(on_conflict.do) return query + + def vector_search( + self, + model: Any, + embedding: Sequence[float], + *filters: _ColumnExpressionArgument[bool], + limit: int, + metric: DistanceMetric | None = None, + ) -> sa.Select[Any]: + """A join against the streaming scan, so filters cannot under-return. + + ``vector_full_scan`` is called without ``k``: a top-k scan would + pick its rows before the ``WHERE`` clause ran, and could then + return fewer than ``limit`` of them. + """ + if metric is not None and metric is not model.__vector_distance__: + msg = ( + "sqlite-vector fixes the distance metric per column; " + f"{model.__name__} uses {model.__vector_distance__.value!r}" + ) + raise UnsupportedDialectError(msg) + table = model.__table__ + scan = sa.func.vector_full_scan( + sa.literal(table.name, sa.String), + sa.literal("embedding", sa.String), + self.query_vector(model, embedding), + ).table_valued("rowid", "distance") + rowid = sa.literal_column(f'"{table.name}".rowid') + return ( + sa.select(model, scan.c.distance) + .select_from(sa.join(table, scan, rowid == scan.c.rowid)) + .where(*filters) + .order_by(scan.c.distance) + .limit(limit) + ) + + def vector_init(self, model: Any) -> sa.Executable: + """Declare the embedding column to sqlite-vector. + + Its dimension and metric are fixed here rather than in the schema, + which is why the statement has to run on every connection that + searches. + """ + options = ( + f"type=FLOAT32,dimension={model.__vector_dim__}," + f"distance={model.__vector_distance__.sqlite_option}" + ) + return sa.select( + sa.func.vector_init( + sa.literal(model.__table__.name, sa.String), + sa.literal("embedding", sa.String), + sa.literal(options, sa.String), + ) + ) diff --git a/sqlargon/i18n/__init__.py b/sqlargon/i18n/__init__.py new file mode 100644 index 0000000..816b540 --- /dev/null +++ b/sqlargon/i18n/__init__.py @@ -0,0 +1,40 @@ +from .expression import current_locale, get_locale, set_locale_getter, translated_value +from .mixin import TranslationMixin +from .repository import TranslatedRepository +from .translatable import ( + TranslatableMixin, + TranslationBase, + current_translation, + translation_class, + translation_table, +) +from .translation import ( + LocaleMap, + TranslatedString, + Translation, + as_translation, + fallback_chain, + select_current, + set_fallback_chain, +) + +__all__ = [ + "LocaleMap", + "TranslatableMixin", + "TranslatedRepository", + "TranslatedString", + "Translation", + "TranslationBase", + "TranslationMixin", + "as_translation", + "current_locale", + "current_translation", + "fallback_chain", + "get_locale", + "select_current", + "set_fallback_chain", + "set_locale_getter", + "translated_value", + "translation_class", + "translation_table", +] diff --git a/sqlargon/i18n/expression.py b/sqlargon/i18n/expression.py new file mode 100644 index 0000000..4032163 --- /dev/null +++ b/sqlargon/i18n/expression.py @@ -0,0 +1,100 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql +from sqlalchemy.ext.compiler import compiles + +if TYPE_CHECKING: + from collections.abc import Callable + + from sqlalchemy.sql.compiler import SQLCompiler + from sqlalchemy.sql.elements import BindParameter, ColumnElement + +_get_locale: Callable[[], str] | None = None + + +def set_locale_getter(fn: Callable[[], str]) -> None: + """Register the callable that returns the active locale per request. + + The callable is invoked at SQL execution time via a late-binding bind + parameter, so a single statement template is cached and reused across + requests of any locale. + """ + global _get_locale # noqa: PLW0603 + _get_locale = fn + + +def get_locale() -> str: + """Return the active locale for the current request. + + Falls through to the callable registered with :func:`set_locale_getter` + at startup. + """ + if _get_locale is None: + msg = ( + "No locale getter has been configured. " + "Call sqlargon.i18n.set_locale_getter() at startup." + ) + raise RuntimeError(msg) + return _get_locale() + + +def current_locale() -> BindParameter[str]: + """Bind parameter resolving to the active locale on every execution. + + Late binding keeps cached statements -- relationship join conditions in + particular, which are built once when mappers are configured -- aware of + the locale of the request being served. + """ + return sa.bindparam( + "current_locale", callable_=get_locale, type_=sa.String, unique=True + ) + + +def _locale_path() -> BindParameter[str]: + return sa.bindparam( + "locale_path", callable_=_current_locale_path, type_=sa.String, unique=True + ) + + +def _current_locale_path() -> str: + return f'$."{get_locale()}"' + + +class translated_value(sa.FunctionElement[str]): + """Text stored under the active locale key of a JSON translation column.""" + + name = "translated_value" + type = sa.String() + inherit_cache = True + + +def _operand(element: translated_value) -> ColumnElement[Any]: + """Read the wrapped column off the element itself. + + Clone and adapt machinery -- `ClauseAdapter`, `with_loader_criteria`, + `with_polymorphic` -- rewrites only the traversed ``clauses``, so anything + cached on the instance would still point at the pre-adaption column. + """ + return next(iter(element.clauses)) + + +@compiles(translated_value, "postgresql") +def _compile_postgresql( + element: translated_value, compiler: SQLCompiler, **kwargs: Any +) -> str: + column = sa.type_coerce(_operand(element), postgresql.JSONB) + return compiler.process( + column.op("->>")(sa.cast(current_locale(), sa.Text)), **kwargs + ) + + +@compiles(translated_value) +def _compile_default( + element: translated_value, compiler: SQLCompiler, **kwargs: Any +) -> str: + return compiler.process( + sa.func.json_extract(_operand(element), _locale_path()), **kwargs + ) diff --git a/sqlargon/i18n/mixin.py b/sqlargon/i18n/mixin.py new file mode 100644 index 0000000..fafab21 --- /dev/null +++ b/sqlargon/i18n/mixin.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +from .expression import get_locale +from .translation import LocaleMap, Translation, as_translation, select_current + + +class TranslationMixin: + """Multi-locale read and write helpers shared by both backends. + + Plain attribute access stays transparent: ``model.title`` reads as the text + of the active locale, while these helpers reach the other locales. Writing + merges -- no write drops a locale it does not name, except + `clear_translations`. + """ + + def get_translations(self, field: str) -> LocaleMap: + """Return every known translation of ``field``, keyed by locale.""" + translation = as_translation(getattr(self, field)) + return {} if translation is None else translation.data + + def get_translation(self, field: str, locale: str | None = None) -> str | None: + """Return the text of ``field`` for ``locale``, the active one by default. + + An explicit ``locale`` is looked up as given; only the active locale + walks its fallback chain. + """ + data = self.get_translations(field) + if locale is None: + return select_current(data) + return data.get(locale) + + def set_translation( + self, field: str, value: str, locale: str | None = None + ) -> None: + """Store ``value`` under ``locale``, keeping the other translations.""" + data = self.get_translations(field) + data[locale or get_locale()] = value + setattr(self, field, Translation(select_current(data) or value, data)) + + def clear_translations(self, field: str) -> None: + """Drop every translation of ``field``, leaving it empty rather than unset. + + The column keeps holding a translation -- an empty one -- so a model may + declare it non-nullable. + """ + setattr(self, field, Translation("", {})) diff --git a/sqlargon/i18n/repository.py b/sqlargon/i18n/repository.py new file mode 100644 index 0000000..79fec33 --- /dev/null +++ b/sqlargon/i18n/repository.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +from sqlargon.repository import SQLAlchemyRepository + +if TYPE_CHECKING: + from typing import Any + + from typing_extensions import Self + + +class TranslatedRepository(SQLAlchemyRepository, abstract=True): + """Repository for a model whose fields are backed by a translation table. + + Every ``select()`` outer-joins the active-locale translation row, so + filtering and ordering on translated columns -- the + :class:`~sqlalchemy.ext.hybrid.hybrid_property` class-level expressions + resolve to the translation table's columns -- works without an explicit + join in the calling code. + + The model must use :class:`TranslatableMixin`, whose + ``_current_translation`` relationship carries the join condition that + matches the model's primary key and the active locale. + """ + + def select( + self, + *args: Any, + **kwargs: Any, + ) -> Self: + return ( + super() + .select(*args, **kwargs) + .join( + self.model._current_translation, # noqa: SLF001 + isouter=True, # type: ignore[union-attr] + ) + ) diff --git a/sqlargon/i18n/translatable.py b/sqlargon/i18n/translatable.py new file mode 100644 index 0000000..b6ce34b --- /dev/null +++ b/sqlargon/i18n/translatable.py @@ -0,0 +1,243 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, ClassVar, cast +from weakref import WeakKeyDictionary + +import sqlalchemy as sa +from sqlalchemy.ext.hybrid import hybrid_property +from sqlalchemy.orm import ( + DeclarativeBase, + Mapped, + declared_attr, + mapped_column, + relationship, +) +from sqlalchemy.orm.collections import attribute_keyed_dict + +from .expression import current_locale +from .mixin import TranslationMixin +from .translation import LocaleMap, Translation, as_translation + +if TYPE_CHECKING: + from collections.abc import Callable, Mapping, MutableMapping + + from sqlalchemy import ColumnElement + +LOCALE_LENGTH = 10 + +_translation_classes: MutableMapping[type[Any], type[Any]] = WeakKeyDictionary() + + +def translation_class(parent: type[Any]) -> type[Any]: + """Return the translation model registered for ``parent``.""" + for klass in parent.__mro__: + translation = _translation_classes.get(klass) + if translation is not None: + return translation + msg = f"No translation table declared for {parent.__name__}" + raise LookupError(msg) + + +def current_translation(parent: Any) -> Any: + """Return the relationship joining a model to its active locale row. + + Given an instance instead of the class it reads the attribute, declared + ``lazy="raise"`` -- it is a join target, never a loader. + """ + return parent._current_translation # noqa: SLF001 + + +class TranslationBase: + """Marker base of every model built by `translation_table`.""" + + __translation_parent__: ClassVar[type[Any]] + + if TYPE_CHECKING: + locale: Mapped[str] + + def __init_subclass__(cls, **kwargs: Any) -> None: + super().__init_subclass__(**kwargs) + if not cls.__dict__.get("__abstract__"): + _translation_classes[cls.__translation_parent__] = cls + _allow_untranslated_fields(cls) + + +def translation_table(parent: type[Any]) -> type[TranslationBase]: + """Build the declarative base of ``parent``'s translation table. + + The returned class carries a copy of ``parent``'s primary key -- cascading + foreign keys back to it -- plus a ``locale`` column, all part of the + translation table's own primary key. The translated columns declared on the + concrete subclass are made nullable, whatever their annotation says -- see + `_allow_untranslated_fields`. + """ + table = parent.__table__ + namespace: dict[str, Any] = { + "__abstract__": True, + "__translation_parent__": parent, + } + for column in table.primary_key.columns: + namespace[column.key] = declared_attr(_parent_key_column(table, column)) + namespace["locale"] = declared_attr(_locale_column) + + base = _declarative_base(parent) + metaclass: Callable[..., type[Any]] = type(base) + return cast( + "type[TranslationBase]", + metaclass( + f"{parent.__name__}TranslationBase", (TranslationBase, base), namespace + ), + ) + + +def _allow_untranslated_fields(translation: type[Any]) -> None: + """Make the translated columns of ``translation`` nullable. + + A locale row carries the fields translated to that locale only: writing one + of them creates the row leaving the others out, and `clear_translations` + empties them again, so ``NULL`` is how an untranslated field is stored -- + which is what `TranslatableMixin.get_translations` skips over. + + A field the translation table does not define is a typo, raised here so it + fails where it is declared instead of at flush time, on a column the + developer never named. + """ + fields = getattr(translation.__translation_parent__, "__translated_fields__", ()) + for field in fields: + column = translation.__table__.columns.get(field) + if column is None: + msg = f"{translation.__name__} has no column {field!r}" + raise TypeError(msg) + column.nullable = True + + +def _declarative_base(model: type[Any]) -> type[DeclarativeBase]: + for klass in model.__mro__: + if DeclarativeBase in klass.__bases__: + return cast("type[DeclarativeBase]", klass) + msg = f"{model.__name__} is not a declarative model" + raise TypeError(msg) + + +def _locale_column(cls: type[Any]) -> Mapped[str]: # noqa: ARG001 + return mapped_column( + "locale", sa.String(LOCALE_LENGTH), primary_key=True, nullable=False + ) + + +def _parent_key_column( + table: sa.Table, column: sa.Column[Any] +) -> Callable[[type[Any]], Mapped[Any]]: + def factory(cls: type[Any]) -> Mapped[Any]: # noqa: ARG001 + return mapped_column( + column.key, + column.type, + sa.ForeignKey(f"{table.fullname}.{column.key}", ondelete="CASCADE"), + primary_key=True, + autoincrement=False, + nullable=False, + ) + + return factory + + +def _locale_join(parent: type[Any], locale: ColumnElement[str]) -> ColumnElement[bool]: + target = translation_class(parent) + clauses = [ + getattr(parent, column.key) == getattr(target, column.key) + for column in parent.__table__.primary_key.columns + ] + clauses.append(target.locale == locale) + return sa.and_(*clauses) + + +class TranslatableMixin(TranslationMixin): + """Keeps the ``__translated_fields__`` of a model in a translation table. + + Each field becomes a hybrid property reading the text of the active locale + (walking its fallback chain) and writing to it, creating the locale row on + demand; assigning ``None`` clears the field in every locale. At class level + the field resolves to the translation table column, so queries must outer + join `current_translation` -- see `TranslatableFilterResolver`. + + Reads go through `_translations`, eagerly loaded with the row itself, while + `_current_translation` is a join target only: it never loads on its own, so + a joined query costs no extra statement and an async session cannot trip + over it. + """ + + __translated_fields__: ClassVar[tuple[str, ...]] = () + + def __init_subclass__(cls, **kwargs: Any) -> None: + super().__init_subclass__(**kwargs) + for field in cls.__translated_fields__: + setattr(cls, field, _translated_property(field)) + + @declared_attr + @classmethod + def _translations(cls) -> Mapped[dict[str, Any]]: + return relationship( + lambda: translation_class(cls), + collection_class=attribute_keyed_dict("locale"), + cascade="all, delete-orphan", + lazy="selectin", + ) + + @declared_attr + @classmethod + def _current_translation(cls) -> Mapped[Any]: + return relationship( + lambda: translation_class(cls), + primaryjoin=lambda: _locale_join(cls, current_locale()), + uselist=False, + viewonly=True, + lazy="raise", + ) + + def get_translations(self, field: str) -> LocaleMap: + if field not in self.__translated_fields__: + return super().get_translations(field) + return { + locale: value + for locale, row in self._translations.items() + if isinstance(value := getattr(row, field, None), str) + } + + def write_translations(self, field: str, data: Mapping[str, str]) -> None: + """Write ``field`` for every locale in ``data``, adding missing rows. + + Locales absent from ``data`` keep their text: writes merge, so assigning + a plain string only touches the active locale. A row created here holds + ``field`` alone, the other translated columns staying ``NULL`` until + they are written. + """ + target = translation_class(type(self)) + for locale, value in data.items(): + row = self._translations.get(locale) + if row is None: + row = target(locale=locale) + self._translations[locale] = row + setattr(row, field, value) + + def clear_translations(self, field: str) -> None: + """Drop ``field`` from every locale row, leaving the other fields.""" + for row in self._translations.values(): + setattr(row, field, None) + + +def _translated_property(field: str) -> hybrid_property[Translation | None]: + def getter(self: TranslatableMixin) -> Translation | None: + data = self.get_translations(field) + return as_translation(data) if data else None + + def setter(self: TranslatableMixin, value: Any) -> None: + translation = as_translation(value) + if translation is None: + self.clear_translations(field) + else: + self.write_translations(field, translation.data) + + def expression(cls: type[Any]) -> ColumnElement[Any]: + return cast("ColumnElement[Any]", getattr(translation_class(cls), field)) + + return hybrid_property(getter).setter(setter).expression(expression) diff --git a/sqlargon/i18n/translation.py b/sqlargon/i18n/translation.py new file mode 100644 index 0000000..88d2a8a --- /dev/null +++ b/sqlargon/i18n/translation.py @@ -0,0 +1,216 @@ +from __future__ import annotations + +from collections.abc import Callable, Mapping +from typing import TYPE_CHECKING, Any + +from pydantic_core import core_schema +from sqlalchemy import TypeDecorator + +from sqlargon.types import JSON + +from .expression import get_locale, translated_value + +if TYPE_CHECKING: + from pydantic import GetCoreSchemaHandler, GetJsonSchemaHandler + from pydantic.json_schema import JsonSchemaValue + from sqlalchemy import ColumnElement, Dialect + from sqlalchemy.sql.operators import OperatorType + from typing_extensions import Self + +LocaleMap = dict[str, str] + +_get_fallback: Callable[[str | None], tuple[str, ...]] | None = None + + +def set_fallback_chain(fn: Callable[[str | None], tuple[str, ...]]) -> None: + """Register the callable that builds the fallback chain for ``locale``. + + The callable receives an explicit locale or ``None`` (meaning "the + active one") and is expected to return the ordered locales to try. + """ + global _get_fallback # noqa: PLW0603 + _get_fallback = fn + + +def fallback_chain(locale: str | None = None) -> tuple[str, ...]: + """Return the locales to look up, best match first.""" + if _get_fallback is None: + msg = ( + "No fallback chain has been configured. " + "Call sqlargon.i18n.set_fallback_chain() at startup." + ) + raise RuntimeError(msg) + return _get_fallback(locale) + + +def select_current(data: Mapping[str, str], locale: str | None = None) -> str | None: + """Pick the value for ``locale`` walking down its fallback chain.""" + for candidate in fallback_chain(locale): + value = data.get(candidate) + if value is not None: + return value + return next(iter(data.values()), None) + + +def _input_schema() -> core_schema.CoreSchema: + """The shapes `Translation._validate` accepts: text, or ``{locale: text}``.""" + return core_schema.union_schema( + [ + core_schema.str_schema(), + core_schema.dict_schema(core_schema.str_schema(), core_schema.str_schema()), + ] + ) + + +class Translation(str): # noqa: SLOT000 + """Current locale text that also carries every other known translation.""" + + _data: LocaleMap + + def __new__(cls, current: str, data: Mapping[str, str] | None = None) -> Self: + translation = super().__new__(cls, current) + translation._data = dict(data) if data is not None else {get_locale(): current} + return translation + + @property + def data(self) -> LocaleMap: + """All known translations, keyed by locale.""" + return dict(self._data) + + def get(self, locale: str) -> str | None: + """Return the text for ``locale``, or ``None`` when it is missing.""" + return self._data.get(locale) + + def update(self, value: str, locale: str | None = None) -> Translation: + """Return a copy with ``value`` stored under ``locale``.""" + data = dict(self._data) + data[locale or get_locale()] = value + return Translation(select_current(data) or value, data) + + @classmethod + def _validate(cls, value: Any) -> Translation: + translation = as_translation(value) + if translation is None: + msg = "Input should be a string or a mapping of locales to strings" + raise ValueError(msg) + return translation + + @classmethod + def __get_pydantic_core_schema__( + cls, source_type: Any, handler: GetCoreSchemaHandler + ) -> core_schema.CoreSchema: + return core_schema.no_info_plain_validator_function( + cls._validate, + serialization=core_schema.plain_serializer_function_ser_schema( + str, return_schema=core_schema.str_schema() + ), + ) + + @classmethod + def __get_pydantic_json_schema__( + cls, schema: core_schema.CoreSchema, handler: GetJsonSchemaHandler + ) -> JsonSchemaValue: + serialization = ( + schema.get("serialization") if handler.mode == "serialization" else None + ) + return_schema = serialization.get("return_schema") if serialization else None + return handler(return_schema or _input_schema()) + + +def as_translation(value: Any, locale: str | None = None) -> Translation | None: + """Normalize a raw attribute value into a `Translation`. + + Unusable input raises `ValueError`, the only error pydantic turns into a + validation error -- a `TypeError` would leak out of validation as a 500. + """ + if value is None: + return None + if isinstance(value, Translation): + return value + if isinstance(value, str): + return Translation(value, {locale or get_locale(): value}) + if isinstance(value, Mapping): + data = {str(key): _text(key, text) for key, text in value.items()} + return Translation(select_current(data, locale) or "", data) + msg = f"Cannot build a Translation from {type(value).__name__}" + raise ValueError(msg) + + +def _text(key: Any, value: Any) -> str: + """Reject non-string texts instead of stringifying them.""" + if not isinstance(value, str): + msg = ( + f"Translation of locale {key!r} must be a string, " + f"got {type(value).__name__}" + ) + raise ValueError(msg) # noqa: TRY004 + return value + + +def _locale_map(value: Any) -> LocaleMap | None: + translation = as_translation(value) + return None if translation is None else translation.data + + +class TranslatedString(TypeDecorator[Translation]): + """JSON column holding ``{locale: text}`` and reading as a `Translation`. + + Every column operator is rewritten to act on the text of the locale active + when the statement runs, so plain ``select(...).where(Model.field == value)`` + needs no join. The JSON methods inherited from `JSON.ComparatorFactory` -- + the reads ``contains``, ``has_any_key``, ``has_all_keys``, ``has_key``, + ``json_value``, ``get``, ``keys``, ``array_length``, the mutations + ``update``, ``set_key``, ``remove_key``, ``insert_key``, ``replace_key``, + ``array_append``, and indexing -- all still address the whole locale map. + A mutation therefore rewrites one locale's entry, keyed by locale name, + rather than the active locale's text. + """ + + impl = JSON + cache_ok = True + + class ComparatorFactory(JSON.ComparatorFactory): + """Redirects every column operator to the active locale's text.""" + + @property + def current(self) -> ColumnElement[str]: + return translated_value(self.expr) + + def operate( + self, op: OperatorType, *other: Any, **kwargs: Any + ) -> ColumnElement[Any]: + return self.current.operate(op, *other, **kwargs) + + def reverse_operate( + self, op: OperatorType, other: Any, **kwargs: Any + ) -> ColumnElement[Any]: + return self.current.reverse_operate(op, other, **kwargs) + + def asc(self) -> ColumnElement[str]: + return self.current.asc() + + def desc(self) -> ColumnElement[str]: + return self.current.desc() + + comparator_factory = ComparatorFactory # pyright: ignore[reportAssignmentType, reportIncompatibleMethodOverride] + + def compare_values(self, x: Any, y: Any) -> bool: + """Compare the whole locale maps, not just the active locale text.""" + return _locale_map(x) == _locale_map(y) + + def process_bind_param( + self, + value: Any, + dialect: Dialect, # noqa: ARG002 + ) -> LocaleMap | None: + return _locale_map(value) + + def process_result_value( + self, + value: Any, + dialect: Dialect, # noqa: ARG002 + ) -> Translation | None: + if value is None: + return None + data = {str(key): text for key, text in value.items() if isinstance(text, str)} + return Translation(select_current(data) or "", data) diff --git a/sqlargon/integrations/eventiq.py b/sqlargon/integrations/eventiq.py new file mode 100644 index 0000000..e258b9d --- /dev/null +++ b/sqlargon/integrations/eventiq.py @@ -0,0 +1,81 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, TypeVar, overload + +from eventiq import CloudEvent + +if TYPE_CHECKING: + from eventiq import Service + + from sqlargon.outbox import OutboxEvent, Publisher + +CloudEventT = TypeVar("CloudEventT", bound="CloudEvent[Any]") + + +@overload +def to_cloud_event( + event: OutboxEvent, *, source: str | None = ... +) -> CloudEvent[Any]: ... + + +@overload +def to_cloud_event( + event: OutboxEvent, *, event_class: type[CloudEventT], source: str | None = ... +) -> CloudEventT: ... + + +def to_cloud_event( + event: OutboxEvent, + *, + event_class: type[CloudEvent[Any]] = CloudEvent, + source: str | None = None, +) -> Any: + """Build the CloudEvent an outbox row stands for. + + ``event_class`` is the class the message is validated against, so a + service with a base of its own -- one carrying a ``tenant_id``, say -- + gets back that class rather than the plain + :class:`~eventiq.CloudEvent`. The row's ``attributes`` are passed as + top-level attributes of it, and the core CloudEvents ones always come + from the row's own columns:: + + message = to_cloud_event(event, event_class=TenantEvent) + """ + attributes: dict[str, Any] = dict(event.attributes or {}) + attributes.update( + id=event.id, + time=event.created_at, + subject=event.topic, + type=event.type, + source=event.source or source, + ) + headers = {name: str(value) for name, value in (event.headers or {}).items()} + return event_class.new(event.data, headers=headers or None, **attributes) + + +def eventiq_publisher( + service: Service, + *, + event_class: type[CloudEvent[Any]] = CloudEvent, + source: str | None = None, +) -> Publisher: + """Build a relay publisher that hands outbox rows to an eventiq service. + + Every claimed event is built as ``event_class``, the base class the + service publishes, and ``source`` is the one events carrying none of + their own are given:: + + relay = OutboxRelay(eventiq_publisher(service, event_class=TenantEvent)) + + async with relay.running(): + ... + """ + + async def publish(event: OutboxEvent) -> None: + message = to_cloud_event(event, event_class=event_class, source=source) + await service.publish(message) + + return publish + + +__all__ = ["eventiq_publisher", "to_cloud_event"] diff --git a/sqlargon/mixins.py b/sqlargon/mixins.py index d35de7d..e4d606e 100644 --- a/sqlargon/mixins.py +++ b/sqlargon/mixins.py @@ -173,3 +173,151 @@ def __mapper_args__(cls) -> dict[str, Any]: "version_id_col": cls.xmin, "version_id_generator": False, } + + +def _generate_version_uuid7(_current: Any) -> UUID: + """Generate a fresh, time sortable UUID for the version column.""" + return uuid7() + + +def _next_integer_version(current: Any) -> int: + """The successor of ``current``, starting the sequence at 1.""" + return 1 if current is None else current + 1 + + +class AuditableMixin(SoftDeleteMixin, VersionedMixin): + """Marker mixin for append-only, versioned models. + + Use :class:`IntegerAuditableMixin` (human readable 1, 2, 3 ...) or + :class:`UUIDAuditableMixin` (time sortable UUIDv7) -- this base declares + no version column of its own, it only carries the expressions every + strategy shares and lets + :class:`~sqlargon.repository.AuditableRepository` validate its model at + subclass time. + + A row is never updated: each change appends a row holding the next + ``version`` of the same entity, so the table *is* the audit log. The + entity is identified by :meth:`audit_key` -- the primary key minus + ``version`` -- and the tombstone inherited from :class:`SoftDeleteMixin` + marks the version that records a deletion. + """ + + if TYPE_CHECKING: + __table__: sa.Table + version: Mapped[Any] + + @classmethod + def audit_key(cls) -> tuple[str, ...]: + """The columns identifying the entity: the primary key minus ``version``.""" + return tuple( + c.name for c in cls.__table__.primary_key.columns if c.name != "version" + ) + + @classmethod + def latest_version(cls, *, before: datetime | None = None) -> Any: + """A scalar subquery holding the newest version of each entity. + + ``ORDER BY version DESC LIMIT 1`` rather than ``MAX(version)``: it is + one expression for both strategies, and it does not rely on a ``max`` + aggregate existing for the UUID type of every backend. Either way the + ordering is total -- native ``uuid`` byte order on PostgreSQL and + lowercase hex ``CHAR(36)`` order elsewhere, both of which put a + UUIDv7's timestamp first. + + ``before`` restricts the subquery to versions recorded up to that + moment, which is what makes an as-of read possible. + """ + table = cls.__table__ + alias = table.alias(f"{table.name}_latest") + conditions = [alias.c[name] == table.c[name] for name in cls.audit_key()] + if before is not None: + conditions.append(alias.c.created_at <= before) + return ( + sa.select(alias.c.version) + .where(*conditions) + .order_by(alias.c.version.desc()) + .limit(1) + # the outer table appears only in the WHERE clause, so say what + # this subquery correlates to rather than leaving it to be guessed + .correlate(table) + .scalar_subquery() + ) + + @classmethod + def is_latest(cls, *, before: datetime | None = None) -> sa.ColumnElement[bool]: + """Whether a row is the newest version of its entity.""" + return cls.version == cls.latest_version(before=before) + + @classmethod + def next_version_expression(cls) -> Any: + """The version superseding the one of the row being read, in SQL. + + This is what lets an append be a single ``INSERT ... SELECT`` rather + than a read followed by a write. + """ + raise NotImplementedError + + +class IntegerAuditableMixin(AuditableMixin): + """Append-only versioning with a human readable counter: 1, 2, 3 ... + + The version is part of the primary key, so two writers deriving the same + successor collide on it rather than one silently overwriting the other. + """ + + version: Mapped[int] = mapped_column( + sa.Integer(), + primary_key=True, + nullable=False, + default=1, + server_default=sa.text("1"), + ) + + @declared_attr.directive + def __mapper_args__(cls) -> dict[str, Any]: + return { + "eager_defaults": True, + "version_id_col": cls.version, + "version_id_generator": _next_integer_version, + } + + @classmethod + def next_version_expression(cls) -> Any: + return cls.__table__.c.version + 1 + + +class UUIDAuditableMixin(AuditableMixin): + """Append-only versioning with a time sortable UUIDv7. + + The successor of a version can be minted without reading the current one, + which suits writers that cannot coordinate. UUIDv7 is monotonic within a + process, but two processes appending in the same millisecond can produce + an inverted pair -- prefer :class:`IntegerAuditableMixin` when strict + chronological ordering across writers has to hold. + """ + + version: Mapped[UUID] = mapped_column( + GUID(), + primary_key=True, + nullable=False, + default=uuid7, + server_default=GenerateUUIDV7(), + ) + + @declared_attr.directive + def __mapper_args__(cls) -> dict[str, Any]: + return { + "eager_defaults": True, + "version_id_col": cls.version, + "version_id_generator": _generate_version_uuid7, + } + + @classmethod + def next_version_expression(cls) -> Any: + """One fresh UUIDv7, shared by every row a single statement appends. + + Sharing it is harmless: rows of one statement belong to different + entities, so the primary key stays unique, and a value minted now + sorts above every version already recorded. + """ + return sa.literal(uuid7(), GUID()) diff --git a/sqlargon/orm.py b/sqlargon/orm.py index 9492182..e82064b 100644 --- a/sqlargon/orm.py +++ b/sqlargon/orm.py @@ -5,8 +5,13 @@ from sqlalchemy.orm import DeclarativeBase, declared_attr from .mixins import ( + AuditableMixin, + CreatedUpdatedMixin, + IntegerAuditableMixin, SoftDeleteMixin, + UUIDAuditableMixin, UUIDVersionedMixin, + VersionedMixin, XminVersionedMixin, ) @@ -70,7 +75,18 @@ class User(UUIDModelMixin, SoftDeleteBase): SoftDeleteModel = TypeVar("SoftDeleteModel", bound=SoftDeleteBase) -class VersionedBase(UUIDVersionedMixin, Base): +class AnyVersionedBase(VersionedMixin, Base): + """Declarative base shared by every versioning strategy. + + It declares no version column of its own; it exists so + :class:`~sqlargon.repository.VersionedRepository` can type its model + against any strategy rather than against the UUID one alone. + """ + + __abstract__ = True + + +class VersionedBase(UUIDVersionedMixin, AnyVersionedBase): """Declarative base for models versioned with a UUID column. Inherit it instead of combining :class:`UUIDVersionedMixin` with @@ -84,7 +100,7 @@ class User(UUIDModelMixin, VersionedBase): __abstract__ = True -class XminVersionedBase(XminVersionedMixin, Base): +class XminVersionedBase(XminVersionedMixin, AnyVersionedBase): """Declarative base for PostgreSQL models versioned via ``xmin``. Only works on PostgreSQL — the ``xmin`` system column does not exist @@ -94,4 +110,47 @@ class XminVersionedBase(XminVersionedMixin, Base): __abstract__ = True -VersionedModel = TypeVar("VersionedModel", bound=VersionedBase) +VersionedModel = TypeVar("VersionedModel", bound=AnyVersionedBase) + + +class AnyAuditableBase( + AuditableMixin, CreatedUpdatedMixin, SoftDeleteBase, AnyVersionedBase +): + """Declarative base shared by every append-only versioning strategy. + + It declares no version column of its own -- inherit + :class:`AuditableBase` or :class:`UUIDAuditableBase`. + + ``created_at`` timestamps the version rather than the entity, and + ``updated_at`` always equals it, because a row of an append-only table + is never updated. + """ + + __abstract__ = True + + +class AuditableBase(IntegerAuditableMixin, AnyAuditableBase): + """Declarative base for append-only models with counted versions. + + The version column joins the primary key, so a model combining this with + :class:`~sqlargon.mixins.UUIDModelMixin` is keyed by ``(id, version)`` + and its entity key is derived as ``(id,)``:: + + class Article(UUIDModelMixin, AuditableBase): + title: Mapped[str] = mapped_column(sa.Unicode(255)) + """ + + __abstract__ = True + + +class UUIDAuditableBase(UUIDAuditableMixin, AnyAuditableBase): + """Declarative base for append-only models versioned by UUIDv7. + + The counterpart of :class:`AuditableBase` for writers that cannot + coordinate on a counter. + """ + + __abstract__ = True + + +AuditableModel = TypeVar("AuditableModel", bound=AnyAuditableBase) diff --git a/sqlargon/outbox/__init__.py b/sqlargon/outbox/__init__.py new file mode 100644 index 0000000..2738f92 --- /dev/null +++ b/sqlargon/outbox/__init__.py @@ -0,0 +1,23 @@ +from .config import ( + ALL_OPERATIONS, + AttributeSource, + Operation, + OutboxConfig, + format_topic, +) +from .models import OutboxEvent +from .relay import OutboxRelay, Publisher +from .repository import OutboxEventRepository, OutboxRepository + +__all__ = [ + "ALL_OPERATIONS", + "AttributeSource", + "Operation", + "OutboxConfig", + "OutboxEvent", + "OutboxEventRepository", + "OutboxRelay", + "OutboxRepository", + "Publisher", + "format_topic", +] diff --git a/sqlargon/outbox/config.py b/sqlargon/outbox/config.py new file mode 100644 index 0000000..790185c --- /dev/null +++ b/sqlargon/outbox/config.py @@ -0,0 +1,76 @@ +from collections.abc import Callable, Mapping +from collections.abc import Set as AbstractSet +from dataclasses import dataclass +from enum import Enum +from types import MappingProxyType +from typing import Any + +__all__ = [ + "ALL_OPERATIONS", + "AttributeSource", + "Operation", + "OutboxConfig", + "format_topic", +] + +# Where an extra CloudEvent attribute comes from: a name of an attribute of +# the written row, or a callable handed that row -- how a ContextVar is read +AttributeSource = str | Callable[[Any], Any] + +CORE_ATTRIBUTES = frozenset( + { + "data", + "datacontenttype", + "dataschema", + "id", + "source", + "specversion", + "subject", + "time", + "topic", + "type", + } +) + + +class Operation(str, Enum): + """The kind of write an outbox event was produced by.""" + + CREATED = "created" + UPDATED = "updated" + DELETED = "deleted" + + +ALL_OPERATIONS = frozenset(Operation) + + +def format_topic(topic: str, row: Any) -> str: + if "{" not in topic: + return topic + return topic.format(**vars(row)) + + +@dataclass(frozen=True, slots=True) +class OutboxConfig: + topic: str | None = None + type_prefix: str | None = None + source: str | None = None + exclude: AbstractSet[str] = frozenset() + include: AbstractSet[str] | None = None + operations: AbstractSet[Operation] = ALL_OPERATIONS + attributes: Mapping[str, AttributeSource] = MappingProxyType({}) + + def __post_init__(self) -> None: + shadowed = CORE_ATTRIBUTES.intersection(self.attributes) + if shadowed: + msg = ( + f"{', '.join(sorted(shadowed))} name CloudEvents core attributes, " + f"which every event carries of its own" + ) + raise ValueError(msg) + + def includes_column(self, name: str) -> bool: + """Whether the column is copied into the payload.""" + if self.include is not None: + return name in self.include + return name not in self.exclude diff --git a/sqlargon/outbox/models.py b/sqlargon/outbox/models.py new file mode 100644 index 0000000..5ee2442 --- /dev/null +++ b/sqlargon/outbox/models.py @@ -0,0 +1,45 @@ +from datetime import datetime +from typing import Any + +import sqlalchemy as sa +from sqlalchemy.orm import Mapped, mapped_column + +from sqlargon.mixins import CreatedUpdatedMixin, UUIDModelMixin +from sqlargon.orm import Base +from sqlargon.types import JSON, Timestamp, now +from sqlargon.utils import utc_now + + +class OutboxEvent(UUIDModelMixin, CreatedUpdatedMixin, Base): + """A CloudEvent awaiting publication. + + ``id`` is the CloudEvent ``id`` and ``created_at`` its ``time``; the + remaining CloudEvent attributes map one to one, and ``specversion`` is + left to the publisher as the constant it is. ``attributes`` holds the + extension attributes -- the ones the CloudEvents spec leaves to the + application -- while ``headers`` holds the broker's transport headers. + Every column carries a server default as well as a client one, so a row + inserted without naming them is still complete. + """ + + __tablename__ = "outbox_events" + __table_args__ = ( + sa.Index( + "idx_outbox_events_pending", "published_at", "available_at", "created_at" + ), + ) + + topic: Mapped[str] = mapped_column(sa.String(255), nullable=False) + type: Mapped[str] = mapped_column(sa.String(255), nullable=False) + source: Mapped[str | None] = mapped_column(sa.String(255), nullable=True) + data: Mapped[dict[str, Any]] = mapped_column(JSON(), nullable=False) + attributes: Mapped[dict[str, Any] | None] = mapped_column(JSON(), nullable=True) + headers: Mapped[dict[str, Any] | None] = mapped_column(JSON(), nullable=True) + published_at: Mapped[datetime | None] = mapped_column(Timestamp(), nullable=True) + available_at: Mapped[datetime] = mapped_column( + Timestamp(), nullable=False, default=utc_now, server_default=now() + ) + attempts: Mapped[int] = mapped_column( + sa.Integer(), nullable=False, default=0, server_default=sa.text("0") + ) + last_error: Mapped[str | None] = mapped_column(sa.Text(), nullable=True) diff --git a/sqlargon/outbox/relay.py b/sqlargon/outbox/relay.py new file mode 100644 index 0000000..1781159 --- /dev/null +++ b/sqlargon/outbox/relay.py @@ -0,0 +1,151 @@ +from __future__ import annotations + +import logging +from collections.abc import Awaitable, Callable +from contextlib import asynccontextmanager +from datetime import timedelta +from typing import TYPE_CHECKING + +import anyio +from anyio import TASK_STATUS_IGNORED + +from sqlargon.utils import utc_now + +from .models import OutboxEvent +from .repository import OutboxEventRepository + +if TYPE_CHECKING: + from collections.abc import AsyncGenerator + from uuid import UUID + + from anyio.abc import TaskStatus + + +logger = logging.getLogger(__name__) + +# What a claimed event is handed to +Publisher = Callable[[OutboxEvent], Awaitable[None]] + +ERROR_LENGTH = 1000 + + +class OutboxRelay: + """Publishes the events an :class:`.OutboxRepository` recorded. + + Claimed events are published one at a time, in the order they were + written; an outbox that reorders its events is not much of an outbox. A + publisher that raises stops the batch, so a transient broker outage + delays the events behind it rather than overtaking them, and the failed + event is retried with an exponential backoff until ``max_attempts``. + + Several relays may run against the same table: events are claimed with + ``FOR UPDATE SKIP LOCKED`` and leased, so no two relays publish the same + one:: + + relay = OutboxRelay(publish) + + async with relay.running(): + ... + + Published events are kept as an audit trail; deleting the ones older than + ``retention`` is :meth:`purge`, which the relay never runs on its own. + """ + + def __init__( # noqa: PLR0913 -- keyword-only tuning knobs, each with a default + self, + publisher: Publisher, + *, + repository: OutboxEventRepository | None = None, + poll_interval: float = 1.0, + batch_size: int = 100, + lease: timedelta = timedelta(minutes=5), + max_attempts: int = 10, + retry_backoff: float = 2.0, + max_retry_delay: float = 300.0, + retention: timedelta = timedelta(days=7), + ) -> None: + self.publisher = publisher + # a repository carries the statement it is building, so each relay + # needs its own; pass one bound elsewhere with + # ``OutboxEventRepository().using(db=...)`` to use another database + self.repository = ( + repository if repository is not None else OutboxEventRepository() + ) + self.poll_interval = poll_interval + self.batch_size = batch_size + self.lease = lease + self.max_attempts = max_attempts + self.retry_backoff = retry_backoff + self.max_retry_delay = max_retry_delay + self.retention = retention + + def retry_delay(self, attempts: int) -> timedelta: + """How long to wait before retrying an event that failed ``attempts`` times.""" + delay = min(self.retry_backoff ** max(attempts - 1, 0), self.max_retry_delay) + return timedelta(seconds=delay) + + async def dispatch_once(self) -> int: + """Publish one batch of due events and return how many were published.""" + events = await self.repository.claim_pending( + utc_now(), + lease=self.lease, + limit=self.batch_size, + max_attempts=self.max_attempts, + ) + published: list[UUID] = [] + for event in events: + if not await self._publish(event): + break + published.append(event.id) + await self.repository.mark_published(published, utc_now()) + return len(published) + + async def purge(self) -> int: + """Delete the events published longer ago than ``retention``. + + The relay never purges on its own -- running it is up to the caller, + and it takes no arguments so it can be registered as a cron task:: + + cron.task("0 3 * * *", "purge_outbox", relay.purge) + """ + return await self.repository.purge(utc_now() - self.retention) + + async def run(self, *, task_status: TaskStatus[None] = TASK_STATUS_IGNORED) -> None: + """Poll for due events and publish them until cancelled. + + Polling errors are logged and retried on the next tick, so a database + blip never takes the relay down. + """ + task_status.started() + while True: + await self._tick() + await anyio.sleep(self.poll_interval) + + @asynccontextmanager + async def running(self) -> AsyncGenerator[OutboxRelay]: + """Run the relay in the background, e.g. in an ASGI lifespan.""" + async with anyio.create_task_group() as tg: + await tg.start(self.run) + try: + yield self + finally: + tg.cancel_scope.cancel() + + async def _publish(self, event: OutboxEvent) -> bool: + try: + await self.publisher(event) + except Exception as exc: + logger.exception("Publishing outbox event %s failed", event.id) + await self.repository.mark_failed( + event.id, + f"{type(exc).__name__}: {exc}"[:ERROR_LENGTH], + utc_now() + self.retry_delay(event.attempts), + ) + return False + return True + + async def _tick(self) -> None: + try: + await self.dispatch_once() + except Exception: + logger.exception("Dispatching outbox events failed") diff --git a/sqlargon/outbox/repository.py b/sqlargon/outbox/repository.py new file mode 100644 index 0000000..86a9731 --- /dev/null +++ b/sqlargon/outbox/repository.py @@ -0,0 +1,393 @@ +from __future__ import annotations + +from operator import attrgetter +from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast, overload + +import sqlalchemy as sa +from pydantic_core import to_jsonable_python + +from sqlargon.mixins import CreatedUpdatedMixin +from sqlargon.orm import Model +from sqlargon.repository import SQLAlchemyRepository +from sqlargon.repository.base import _as_result, _as_scalars + +from .config import Operation, OutboxConfig, format_topic +from .models import OutboxEvent + +if TYPE_CHECKING: + from collections.abc import Callable, Collection, Sequence + from datetime import datetime, timedelta + from uuid import UUID + + from sqlalchemy import CursorResult, Result, ScalarResult + from sqlalchemy.orm import Mapper + from sqlalchemy.sql._typing import _ColumnExpressionArgument + from typing_extensions import Unpack + + from sqlargon.typing import MultipleValues, OnConflictOptions, SingleValue, Values + + +class OutboxEventRepository(SQLAlchemyRepository[OutboxEvent]): + """Repository for :class:`OutboxEvent` rows, used by the relay.""" + + async def claim_pending( + self, + now: datetime, + *, + lease: timedelta, + limit: int = 100, + max_attempts: int | None = None, + ) -> Sequence[OutboxEvent]: + """Claim the events due at ``now`` and return them in write order. + + Due rows are locked with ``FOR UPDATE SKIP LOCKED`` and leased -- their + ``available_at`` pushed to ``now + lease`` and ``attempts`` bumped -- + within the same transaction, so concurrent relays never claim the same + event and one that dies mid-batch leaves nothing wedged. + + Events that already failed ``max_attempts`` times are left alone; they + keep their ``last_error`` for inspection. + """ + filters: list[_ColumnExpressionArgument[bool]] = [ + OutboxEvent.published_at.is_(None), + OutboxEvent.available_at <= now, + ] + if max_attempts is not None: + filters.append(OutboxEvent.attempts < max_attempts) + async with self.session(): + events = await ( + self.select(with_for_update={"skip_locked": True}) + .filter(*filters) + .order_by(OutboxEvent.created_at, OutboxEvent.id) + .limit(limit) + .all() + ) + for event in events: + event.attempts += 1 + event.available_at = now + lease + return events + + async def mark_published( + self, ids: Collection[UUID], published_at: datetime + ) -> None: + """Mark the events as published, clearing any recorded error.""" + if not ids: + return + await ( + self.update({"published_at": published_at, "last_error": None}) + .filter(OutboxEvent.id.in_(list(ids))) + .execute() + ) + + async def mark_failed(self, event_id: UUID, error: str, retry_at: datetime) -> None: + """Record why publishing failed and when to try again.""" + await ( + self.update({"last_error": error, "available_at": retry_at}) + .filter(OutboxEvent.id == event_id) + .execute() + ) + + async def purge(self, published_before: datetime) -> int: + """Delete the events published before the cutoff and count them.""" + result = await ( + self.delete() + .filter( + OutboxEvent.published_at.is_not(None), + OutboxEvent.published_at < published_before, + ) + .execute() + ) + return cast("CursorResult[Any]", result).rowcount + + async def pending_count(self) -> int: + """How many events are still waiting to be published.""" + return await self.count(OutboxEvent.published_at.is_(None)) + + +class OutboxRepository(SQLAlchemyRepository[Model], abstract=True): + """Repository that records every write as an event, in the same transaction. + + Each statement it writes is followed by an insert into ``outbox_events`` + on the same session, so an event can neither be lost by a rollback nor + published for a row that never committed:: + + class User(UUIDModelMixin, CreatedUpdatedMixin, Base): ... + + + class UserRepository(OutboxRepository[User]): + outbox = OutboxConfig(topic="users", exclude={"password"}) + + + await UserRepository().create(name="John", password=hashed) + # -> one outbox_events row, type "user.created", password left out + + The model must carry + :class:`~sqlargon.mixins.CreatedUpdatedMixin`, whose ``is_new`` is what + tells an insert from an update: an upsert writes some rows and updates + others in one statement, and each is recorded as what it turned out to + be. What those writes look like as events is otherwise settled by the + repository rather than by the model, the way + :attr:`~sqlargon.repository.SQLAlchemyRepository.on_conflict` is, so two + repositories over the same model can publish differently and a model + reached through an ordinary repository records nothing. The payload is the + full row, minus whatever :attr:`outbox` excludes. + + Reads back are needed to build the events, so ``remove`` and + ``delete_many`` go through a RETURNING statement here rather than a bare + DELETE. On backends without RETURNING the repository's existing fallback + (select the identities, write, re-fetch) supplies the rows, so nothing + dialect-specific is involved. + """ + + event_repository_class: ClassVar[type[OutboxEventRepository]] = ( + OutboxEventRepository + ) + outbox: ClassVar[OutboxConfig] = OutboxConfig() + + def __init_subclass__(cls, *, abstract: bool = False, **kwargs: Any) -> None: + super().__init_subclass__(abstract=abstract, **kwargs) + if not abstract and not issubclass(cls.model, CreatedUpdatedMixin): + msg = ( + f"{cls.model.__name__} must inherit from CreatedUpdatedMixin " + f"to be used with {cls.__name__}" + ) + raise TypeError(msg) + + @property + def events(self) -> OutboxEventRepository: + """A repository for the event table, bound to the same database.""" + return self.event_repository_class().using(db=self.db) + + @property + def topic(self) -> str: + """The topic events of this repository are published on. + + The configured topic, or the table name when none is configured. A + topic with ``{placeholders}`` is a template: each event's topic is + filled from the row it was written from. + """ + return self.outbox.topic or self.model.__tablename__ + + @property + def event_types(self) -> dict[Operation, str]: + """The event type of each recorded operation, and only of those.""" + prefix = self.outbox.type_prefix or self.model.__tablename__ + return { + operation: f"{prefix}.{operation.value}" + for operation in self.outbox.operations + } + + @property + def payload_columns(self) -> tuple[tuple[str, str], ...]: + """``(column name, attribute key)`` pairs making up the payload. + + Keyed by column name rather than by attribute key, so the payload + speaks the database's names rather than the mapper's. + """ + mapper: Mapper[Model] = sa.inspect(self.model) + includes_column = self.outbox.includes_column + return tuple( + (column.name, key) + for key, column in mapper.columns.items() + if includes_column(column.name) + ) + + @property + def attribute_sources(self) -> tuple[tuple[str, Callable[[Any], Any]], ...]: + """``(attribute name, getter)`` pairs for the extra CloudEvent attributes. + + A source named as a string is resolved once, here, into the getter + reading it off the written row. + """ + return tuple( + (name, source if callable(source) else attrgetter(source)) + for name, source in self.outbox.attributes.items() + ) + + def _operation_of(self, row: Model) -> Operation: + """Whether the write inserted the row or updated one already there. + + An upsert does both in one statement, and only the row itself knows + which it was: ``created_at`` is the value it was inserted with, so a + row whose ``updated_at`` still matches it has just been created. + """ + is_new = cast("CreatedUpdatedMixin", row).is_new + return Operation.CREATED if is_new else Operation.UPDATED + + def _build_events( + self, written: Sequence[tuple[Model, Operation]] + ) -> list[dict[str, Any]]: + """One event per written row, with a payload of JSON primitives. + + The payload is reduced to primitives here rather than left to the + engine's ``json_serializer``: a :class:`~sqlargon.Database` built + straight from a URL carries the standard library's, which knows + neither UUIDs nor datetimes. The extra attributes are read here too, + and a templated topic is filled from the row, while the context the + write ran in is still the current one. + """ + columns = self.payload_columns + sources = self.attribute_sources + event_types = self.event_types + topic = self.topic + return [ + { + "topic": format_topic(topic, row), + "type": event_type, + "source": self.outbox.source, + "data": to_jsonable_python( + {name: getattr(row, key) for name, key in columns} + ), + "attributes": to_jsonable_python( + {name: source(row) for name, source in sources} + ) + if sources + else None, + } + for row, operation in written + if (event_type := event_types.get(operation)) is not None + ] + + async def _record( + self, rows: Sequence[Model], operation: Operation | None = None + ) -> Sequence[Model]: + """Record one event per row, on the session the write ran on. + + Without an ``operation`` every row reports its own, which is what an + insert that may have updated instead needs. + """ + events = self._build_events( + [ + (row, operation if operation is not None else self._operation_of(row)) + for row in rows + ] + ) + if events: + await self.events.bulk_create(events, ignore_conflicts=False) + return rows + + async def _insert_returning( + self, + values: MultipleValues, + *, + do: Literal["ignore", "update"] | None = None, + **options: Unpack[OnConflictOptions], + ) -> ScalarResult[Model]: + async with self.session(): + result = await super()._insert_returning(values, do=do, **options) + return _as_scalars(await self._record(result.all())) + + async def _update_returning( + self, values: Values, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> ScalarResult[Model]: + async with self.session(): + result = await super()._update_returning(values, *args, **kwargs) + return _as_scalars(await self._record(result.all(), Operation.UPDATED)) + + async def _delete_returning( + self, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> ScalarResult[Model]: + async with self.session(): + result = await super()._delete_returning(*args, **kwargs) + return _as_scalars(await self._record(result.all(), Operation.DELETED)) + + async def remove( + self, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> None: + """Delete the matched rows, recording one event per row.""" + await self._delete_returning(*args, **kwargs) + + async def delete_many( + self, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> None: + """Delete the matched rows, recording one event per row.""" + await self._delete_returning(*args, **kwargs) + + @overload + async def bulk_create( + self, + values: MultipleValues, + *, + ignore_conflicts: bool = ..., + return_results: Literal[False] = ..., + **options: Unpack[OnConflictOptions], + ) -> None: ... + + @overload + async def bulk_create( + self, + values: MultipleValues, + *, + ignore_conflicts: bool = ..., + return_results: Literal[True], + **options: Unpack[OnConflictOptions], + ) -> Sequence[Model]: ... + + async def bulk_create( + self, + values: MultipleValues, + *, + ignore_conflicts: bool = True, + return_results: bool = False, + **options: Unpack[OnConflictOptions], + ) -> Sequence[Model] | None: + """Insert the rows, recording one event per written row.""" + result = await self._insert_returning( + values, do="ignore" if ignore_conflicts else None, **options + ) + return result.all() if return_results else None + + @overload + async def bulk_create_or_update( + self, + values: MultipleValues, + *, + return_results: Literal[False] = ..., + **options: Unpack[OnConflictOptions], + ) -> Result: ... + + @overload + async def bulk_create_or_update( + self, + values: MultipleValues, + *, + return_results: Literal[True], + **options: Unpack[OnConflictOptions], + ) -> Sequence[Model]: ... + + async def bulk_create_or_update( + self, + values: MultipleValues, + *, + return_results: bool = False, + **options: Unpack[OnConflictOptions], + ) -> Sequence[Model] | Result: + """Upsert the rows, recording one event per written row. + + Without ``return_results`` the written rows are still handed back, in a + result built from them rather than the cursor's -- so unlike the base + repository's, it carries no ``rowcount``. + """ + rows = (await self._insert_returning(values, do="update", **options)).all() + return rows if return_results else _as_result(rows) + + async def bulk_update( + self, + values: MultipleValues, + *args: Any, + on_: set[str] | None = None, + **kwargs: Any, + ) -> None: + """Update many rows in one statement, recording one event per row. + + An executemany cannot return rows, so the updated rows are read back + by the keys the statement matched on. + """ + on_ = on_ or self._get_default_index_elements() + elements = tuple(on_) + async with self.session(): + await super().bulk_update(values, *args, on_=on_, **kwargs) + rows = await self._fetch( + self._identity_filter(cast("Sequence[SingleValue]", values), elements) + ) + await self._record(rows.all(), Operation.UPDATED) diff --git a/sqlargon/query_builder.py b/sqlargon/query_builder.py index 8e56a0c..1055efd 100644 --- a/sqlargon/query_builder.py +++ b/sqlargon/query_builder.py @@ -7,6 +7,8 @@ import sqlalchemy as sa if TYPE_CHECKING: + from collections.abc import Sequence + from sqlalchemy.sql._typing import ( _ColumnExpressionArgument, _DMLTableArgument, @@ -18,6 +20,7 @@ ReturningUpdate, ) + from .types.vector import DistanceMetric from .typing import OnConflict, OnConflictOptions, Values, WithForUpdate from enum import Flag, auto @@ -28,6 +31,10 @@ class Option(Flag): RETURNING = auto() CONFLICTS = auto() LOCKS = auto() + #: similarity search over an embedding column + VECTORS = auto() + #: ranked full text search + FULL_TEXT = auto() class QueryBuilderError(Exception): @@ -38,6 +45,10 @@ class UnsupportedOption(QueryBuilderError): pass +class UnsupportedDialectError(QueryBuilderError): + """The dialect cannot express the query that was asked of it.""" + + class QueryBuilder: supported_options: Option = Option.NONE @@ -261,6 +272,103 @@ def page( total_query = self.count(query.subquery()) if include_total else None return page_query, total_query + def identity_column(self, model: Any) -> sa.Column[Any]: + """The single column identifying a row of ``model``. + + Rank fusion joins its candidate sets on it, so a model keyed by + more than one column cannot take part. + """ + primary_key = model.__table__.primary_key.columns + if len(primary_key) != 1: + msg = f"{model.__name__} must have a single-column primary key" + raise TypeError(msg) + return next(iter(primary_key)) + + def query_vector( + self, model: Any, embedding: Sequence[float] + ) -> sa.BindParameter[Any]: + """``embedding``, bound as the type of the embedding column. + + Reading the type off the column rather than naming it keeps the + backend's own encoding -- a pgvector literal on PostgreSQL, a + packed float32 blob on SQLite. + """ + return sa.bindparam( + "search_embedding", + list(embedding), + type_=model.embedding.type, + unique=True, + ) + + def _unsupported(self, feature: str) -> UnsupportedDialectError: + msg = f"{feature} is not supported by {type(self).__name__}" + return UnsupportedDialectError(msg) + + def vector_distance( + self, + model: Any, + embedding: Sequence[float], + metric: DistanceMetric | None = None, + ) -> sa.ColumnElement[float]: + """How far the model's embedding is from ``embedding``.""" + feature = "a vector distance expression" + raise self._unsupported(feature) + + def vector_search( + self, + model: Any, + embedding: Sequence[float], + *filters: _ColumnExpressionArgument[bool], + limit: int, + metric: DistanceMetric | None = None, + ) -> sa.Select[Any]: + """Rows of ``model`` nearest ``embedding``, with their distance. + + The distance is selected as a second column, and ``filters`` are + applied before the limit so narrowing the search cannot return + fewer rows than it should. + """ + feature = "vector search" + raise self._unsupported(feature) + + def vector_init(self, model: Any) -> sa.Executable | None: + """A statement declaring the embedding column to the backend. + + Only sqlite-vector needs one, and it has to run on the connection + the search will run on -- which is why this is a statement rather + than part of the schema. ``None`` means no declaration is needed. + """ + return None + + def text_search( + self, + model: Any, + query: str, + *filters: _ColumnExpressionArgument[bool], + limit: int, + ) -> sa.Select[Any]: + """Rows of ``model`` matching ``query``, with their score.""" + feature = "full text search" + raise self._unsupported(feature) + + def rrf_search( # noqa: PLR0913 -- keyword-only tuning knobs, each with a default + self, + model: Any, + embedding: Sequence[float], + query: str, + *filters: _ColumnExpressionArgument[bool], + k: int = 60, + limit: int = 10, + candidates: int = 50, + ) -> sa.Select[Any]: + """Rows of ``model`` ranked by fusing similarity with full text. + + Scores each row ``sum(1 / (k + rank))`` over the two rankings it + appears in, each cut to ``candidates`` rows. + """ + feature = "reciprocal rank fusion" + raise self._unsupported(feature) + def lock(self, key: str) -> sa.TextClause: msg = f"Cannot obtain lock for key {key}" raise NotImplementedError(msg) diff --git a/sqlargon/repository/__init__.py b/sqlargon/repository/__init__.py new file mode 100644 index 0000000..fcc6a81 --- /dev/null +++ b/sqlargon/repository/__init__.py @@ -0,0 +1,14 @@ +from .auditable import AppendOnlyError, AuditableRepository +from .base import SQLAlchemyRepository +from .soft_delete import DeletedRowExistsError, SoftDeleteRepository +from .versioned import ConcurrentModificationError, VersionedRepository + +__all__ = [ + "AppendOnlyError", + "AuditableRepository", + "ConcurrentModificationError", + "DeletedRowExistsError", + "SQLAlchemyRepository", + "SoftDeleteRepository", + "VersionedRepository", +] diff --git a/sqlargon/repository/auditable.py b/sqlargon/repository/auditable.py new file mode 100644 index 0000000..25262f0 --- /dev/null +++ b/sqlargon/repository/auditable.py @@ -0,0 +1,538 @@ +from __future__ import annotations + +from collections.abc import Mapping as MappingABC +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any + +import sqlalchemy as sa + +from sqlargon.mixins import AuditableMixin +from sqlargon.orm import AuditableModel +from sqlargon.query_builder import Option + +from .base import _as_result +from .soft_delete import SoftDeleteRepository +from .versioned import VersionedRepository + +if TYPE_CHECKING: + from collections.abc import Sequence + from datetime import datetime + + from sqlalchemy import Result, ScalarResult + from sqlalchemy.sql._typing import _ColumnExpressionArgument + from typing_extensions import Self, Unpack + + from sqlargon.typing import ( + MultipleValues, + OnConflictOptions, + Params, + SingleValue, + Values, + ) + +__all__ = ["AppendOnlyError", "AuditableRepository"] + + +class AppendOnlyError(RuntimeError): + """A statement would have rewritten a row of an append-only table.""" + + +@dataclass(slots=True, frozen=True) +class _Append: + """An append a builder has staged but not yet turned into a statement. + + It is held rather than compiled straight away so ``.filter(...)`` keeps + narrowing the select the append reads from, exactly as it would narrow an + ordinary ``UPDATE``. + """ + + values: SingleValue + tombstone: bool + return_results: bool + + +class AuditableRepository( + SoftDeleteRepository[AuditableModel], + VersionedRepository[AuditableModel], + abstract=True, +): + """Repository that appends a new version instead of updating a row. + + The table *is* the audit log: every write appends a row carrying the next + ``version`` of the same entity, and nothing is ever rewritten. Reads are + scoped to the newest live version, so the usual methods keep their usual + meaning while the history stays underneath:: + + class Article(UUIDModelMixin, AuditableBase): ... + + + class ArticleRepository(AuditableRepository[Article]): ... + + + articles = ArticleRepository() + + article = await articles.create(title="draft") # version 1 + await articles.update_one({"title": "final"}, Article.id == article.id) + + await articles.get(id=article.id) # version 2 + await articles.history(id=article.id) # versions 1 and 2 + await articles.remove(Article.id == article.id) # appends version 3, + await articles.list() # tombstoned, so the entity is gone from reads + await articles.versions().count() # but all three rows are still there + + The entity is identified by :meth:`~sqlargon.mixins.AuditableMixin.audit_key` + -- the primary key minus ``version`` -- so ``count()`` counts entities + while ``versions().count()`` counts rows. Deletion appends a tombstoned + version rather than removing anything, which makes ``remove``, + ``delete_one`` and ``delete_many`` all recoverable through + :meth:`~sqlargon.repository.SoftDeleteRepository.restore`. + + :meth:`update` and :meth:`delete` still build a statement, so the fluent + form works unchanged -- it is an ``INSERT ... SELECT`` reading the current + heads and writing their successors, which makes an append one statement + rather than a read followed by a write:: + + await articles.update({"title": "final"}).filter(Article.id == aid) + + Because the version is part of the primary key, two writers deriving the + same successor collide there rather than one of them silently winning. + :meth:`~sqlargon.repository.VersionedRepository.update_if_match` and + ``delete_if_match`` are inherited and land on the append path, giving the + cheaper check first:: + + await articles.update_if_match( + {"title": "final"}, + Article.id == article.id, + expected_version=1, + raise_on_mismatch=True, + ) + + Only :meth:`upsert` is refused: resolving a conflict by rewriting the + conflicting row is the one thing an append-only table cannot do, and + :meth:`create_or_update` and :meth:`bulk_create_or_update` express the + intent behind it. + + The model type variable is bound to + :class:`~sqlargon.orm.AnyAuditableBase`, so a type checker rejects a model + that cannot be audited. At runtime the looser + :class:`~sqlargon.mixins.AuditableMixin` is enough; anything else raises + ``TypeError`` on subclassing. + """ + + __slots__ = ("all_versions", "as_of", "pending") + + def __init__(self) -> None: + super().__init__() + self.all_versions = False + self.as_of: datetime | None = None + self.pending: _Append | None = None + + def __init_subclass__(cls, *, abstract: bool = False, **kwargs: Any) -> None: + super().__init_subclass__(abstract=abstract, **kwargs) + if abstract: + return + if not cls.model.audit_key(): + msg = ( + f"{cls.model.__name__} is keyed by its version alone, so " + f"{cls.__name__} could not tell one entity from another; put " + "the columns identifying the entity in its primary key" + ) + raise TypeError(msg) + + @classmethod + def _required_mixin(cls) -> tuple[type, str]: + return ( + AuditableMixin, + "AuditableMixin (IntegerAuditableMixin or UUIDAuditableMixin)", + ) + + # --- scoping --- + + @property + def _scope(self) -> _ColumnExpressionArgument[bool] | None: + tombstone = super()._scope + if self.all_versions: + return tombstone + latest = self.model.is_latest(before=self.as_of) + if tombstone is None: + return latest + return sa.and_(latest, tombstone) + + def copy(self, query: Any) -> Self: + clone = super().copy(query) + clone.all_versions = self.all_versions + clone.as_of = self.as_of + clone.pending = self.pending + return clone + + def versions(self) -> Self: + """Return a copy covering every version of every entity. + + The scope holds for every statement the copy builds, so + ``versions().count()`` counts rows rather than entities and + ``versions().filter(...)`` searches the whole history. + """ + clone = self.copy(self._query) + clone.all_versions = True + clone.include_deleted = True + clone.deleted_only = False + return clone + + def at(self, timestamp: datetime) -> Self: + """Return a copy reading the state as it stood at ``timestamp``. + + Each entity resolves to the newest version recorded up to that moment, + and one already tombstoned by then stays hidden, exactly as it would + have been at the time. + """ + clone = self.copy(self._query) + clone.as_of = timestamp + return clone + + # --- history --- + + async def history( + self, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> Sequence[AuditableModel]: + """Every version of the matched entities, oldest first.""" + model = self.model + order_by = ( + *(getattr(model, name) for name in model.audit_key()), + model.version, + ) + return ( + await self.versions() + .select() + .filter(*args, **kwargs) + .order_by(*order_by) + .all() + ) + + async def get_version( + self, version: Any, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> AuditableModel | None: + """One exact version of the matched entity, tombstoned or not.""" + return ( + await self.versions() + .select() + .filter(self.model.version == version, *args, **kwargs) + .one_or_none() + ) + + # --- staging an append --- + + def update(self, values: Values, *, return_results: bool = False) -> Self: + """Stage the next version of every entity the statement goes on to match. + + Nothing is rewritten: what this builds is an ``INSERT ... SELECT`` + over the current heads, so ``.filter(...)`` narrows the *source* + select exactly as it would narrow an ``UPDATE``:: + + await repo.update({"title": "final"}).filter(Article.id == aid) + """ + if not isinstance(values, MappingABC): + name = type(self).__name__ + msg = ( + f"{name} appends one version per matched entity and so takes " + "a single mapping; use bulk_update to give each entity its " + "own columns" + ) + raise AppendOnlyError(msg) + clone = self.select() + clone.pending = _Append(values, tombstone=False, return_results=return_results) + return clone + + def delete(self, *, return_results: bool = False) -> Self: + """Stage a tombstoned version of every entity the statement matches.""" + clone = self.select() + clone.pending = _Append({}, tombstone=True, return_results=return_results) + return clone + + def _append_statement(self, pending: _Append) -> Any: + """Compile the staged append against the select built so far.""" + table = self.model.__table__ + carried = self._carried_columns() + columns: list[str] = [] + selected: list[Any] = [] + for column in table.columns: + name = column.name + if name in {"version", "tombstone"}: + continue + if name in pending.values: + value = pending.values[name] + columns.append(name) + selected.append( + value + # a mapped attribute is not a ClauseElement but resolves + # to one, and either may stand in for a literal + if isinstance(value, sa.ClauseElement) + or hasattr(value, "__clause_element__") + else sa.literal(value, column.type) + ) + elif name in carried: + columns.append(name) + selected.append(column) + columns += ["version", "tombstone"] + selected += [ + self.model.next_version_expression(), + sa.literal(pending.tombstone, table.c.tombstone.type), + ] + statement = sa.insert(self.model).from_select( + columns, self.query.with_only_columns(*selected) + ) + if pending.return_results and self.qb.supports(Option.RETURNING): + return statement.returning(self.model) + return statement + + async def execute( + self, + params: Params | None = None, + *, + read_only: bool | None = None, + **kwargs: Any, + ) -> Result: + pending = self.pending + if pending is None: + return await super().execute(params, read_only=read_only, **kwargs) + statement = self._append_statement(pending) + if not pending.return_results or self.qb.supports(Option.RETURNING): + return await self.execute_query( + statement, params, read_only=read_only, **kwargs + ) + # a backend without RETURNING has to find the appended rows again, + # which it can: an append leaves its entity's newest version behind + elements = self.model.audit_key() + columns = tuple(self._column(name) for name in elements) + async with self.session(): + identities = [ + dict(zip(elements, tuple(row), strict=True)) + for row in ( + await self.execute_query(self.query.with_only_columns(*columns)) + ).all() + ] + await self.execute_query(statement, params, **kwargs) + if not identities: + return _as_result([]) + appended = await self._fetch( + sa.and_( + self._identity_filter(identities, elements), + self.model.is_latest(), + ) + ) + return _as_result(appended.all()) + + def stream( + self, + params: Params | None = None, + *, + read_only: bool | None = None, + **kwargs: Any, + ) -> Any: + pending = self.pending + query = self.query if pending is None else self._append_statement(pending) + return self.stream_query(query, params, read_only=read_only, **kwargs) + + async def _update_returning( + self, values: Values, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> ScalarResult[AuditableModel]: + return ( + await self.update(values, return_results=True) + .filter(*args, **kwargs) + .scalars() + ) + + async def _delete_returning( + self, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> ScalarResult[AuditableModel]: + return await self.delete(return_results=True).filter(*args, **kwargs).scalars() + + # --- building a version client side, where one statement cannot --- + + @classmethod + def _next_version(cls, current: Any) -> Any: + """The version superseding ``current``, per the model's strategy. + + The counterpart of + :meth:`~sqlargon.mixins.AuditableMixin.next_version_expression` for + the paths that carry a different set of values per entity and so + cannot be one ``INSERT ... SELECT``. + """ + return cls._version_generator()(current) + + @classmethod + def _carried_columns(cls) -> set[str]: + """The columns an appended version inherits from the one it supersedes. + + Everything the append derives itself is left out, so a value the + caller did not name is carried forward while the new row still gets + its own version and timestamps. + """ + derived = {"version", "tombstone", "created_at", "updated_at"} + return {c.name for c in cls.model.__table__.columns if c.name not in derived} + + def _successor( + self, row: AuditableModel, values: SingleValue, *, tombstone: bool = False + ) -> dict[str, Any]: + return { + **{name: getattr(row, name) for name in self._carried_columns()}, + **values, + "version": self._next_version(row.version), + "tombstone": tombstone, + } + + def _keyed(self, rows: MultipleValues) -> tuple[str, ...]: + elements = self.model.audit_key() + missing = [name for name in elements if any(name not in row for row in rows)] + if missing: + msg = ( + f"every row has to name the entity key of {self.model.__name__} " + f"{elements}, but {missing} is missing from at least one" + ) + raise AppendOnlyError(msg) + return elements + + # --- writes --- + + async def create_or_update(self, **kwargs: Any) -> AuditableModel: + """Append the next version of an entity, creating it if it has none. + + An entity whose newest version is a tombstone is revived rather than + refused: unlike a soft delete, the append leaves the deletion in the + history for anyone to read. + """ + key = self.model.audit_key() + identity = {name: kwargs[name] for name in key if name in kwargs} + async with self.session(): + if len(identity) == len(key): + appended = await self.with_deleted().update_many(kwargs, **identity) + if appended: + return appended[0] + return (await self._insert_returning([kwargs])).one() + + async def bulk_create_or_update( + self, + values: MultipleValues, + *, + return_results: bool = False, + **options: Unpack[OnConflictOptions], # noqa: ARG002 + ) -> Any: + """Append a version per known entity and create the rest at version 1. + + The bulk form of :meth:`create_or_update`. Conflict options are + ignored -- an append has no conflicting row to resolve against. + """ + rows = list(values) + if not rows: + return [] if return_results else _as_result([]) + elements = self._keyed(rows) + async with self.session(): + heads = ( + await self.with_deleted() + .select() + .filter(self._identity_filter(rows, elements)) + .all() + ) + by_identity = { + tuple(getattr(head, name) for name in elements): head for head in heads + } + appended, created = [], [] + for row in rows: + head = by_identity.get(self._identity(row, elements)) + if head is None: + created.append(dict(row)) + else: + appended.append(self._successor(head, row)) + # an appended row names every carried column while a created one + # names only what the caller gave, so the two cannot share an + # executemany -- a column missing from one row of a multi-values + # INSERT has no bound parameter to render + if return_results: + return [ + row + for batch in (appended, created) + if batch + for row in (await self._insert_returning(batch)).all() + ] + result: Result = _as_result([]) + for batch in (appended, created): + if batch: + result = await self.insert(batch).execute() + return result + + async def bulk_update( + self, + values: MultipleValues, + *args: Any, + on_: set[str] | None = None, + **kwargs: Any, + ) -> None: + """Append one version per row, matching entities on ``on_``. + + ``on_`` defaults to the entity key rather than the primary key, since + naming the version would pin every row to the one it already has. + """ + rows = list(values) + if not rows: + return + elements = tuple(on_) if on_ else self._keyed(rows) + async with self.session(): + current = await ( + self.select() + .filter(self._identity_filter(rows, elements), *args, **kwargs) + .all() + ) + by_identity = { + tuple(getattr(row, name) for name in elements): row for row in current + } + appended = [ + self._successor(head, row) + for row in rows + if (head := by_identity.get(self._identity(row, elements))) is not None + ] + if appended: + await self.insert(appended).execute() + + async def purge( + self, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> None: + """Physically delete every superseded version, keeping the newest. + + The only method here that destroys history -- for retention, not for + deletion, which :meth:`remove` records instead. A version something + still points at through + :func:`~sqlargon.audit.version_foreign_key` is protected by that key. + """ + elements = (*self.model.audit_key(), "version") + query = self.qb.filter( + self.qb.select(*(self._column(name) for name in elements)).where( + sa.not_(self.model.is_latest()) + ), + *args, + **kwargs, + ) + async with self.session(): + superseded = [ + dict(zip(elements, tuple(row), strict=True)) + for row in (await self.execute_query(query)).all() + ] + if not superseded: + return + await self.versions().hard_delete( + self._identity_filter(superseded, elements) + ) + + # --- the one statement an append-only table cannot serve --- + + def upsert( + self, + values: Values, # noqa: ARG002 + *, + return_results: bool = False, # noqa: ARG002 + **options: Unpack[OnConflictOptions], # noqa: ARG002 + ) -> Self: + msg = ( + f"{type(self).__name__} cannot resolve a conflict by rewriting the " + "conflicting row; use create_or_update or bulk_create_or_update to " + "append the next version instead" + ) + raise AppendOnlyError(msg) diff --git a/sqlargon/repository.py b/sqlargon/repository/base.py similarity index 60% rename from sqlargon/repository.py rename to sqlargon/repository/base.py index aa374f8..75cccef 100644 --- a/sqlargon/repository.py +++ b/sqlargon/repository/base.py @@ -1,8 +1,7 @@ from __future__ import annotations -from collections.abc import Mapping as MappingABC from contextlib import asynccontextmanager -from typing import TYPE_CHECKING, Any, ClassVar, Generic, Literal, overload +from typing import TYPE_CHECKING, Any, ClassVar, Generic, Literal, cast, overload from sqlalchemy import ( Delete, @@ -12,21 +11,20 @@ Row, ScalarResult, SQLColumnExpression, - Text, and_, bindparam, - cast, false, or_, ) +from sqlalchemy.engine.result import IteratorResult, SimpleResultMetaData from sqlalchemy.orm import QueryableAttribute, selectinload -from .mixins import SoftDeleteMixin, VersionedMixin -from .orm import Base, Model, SoftDeleteModel, VersionedModel -from .query_builder import Option, UnsupportedOption -from .registry import get_default_database -from .routing import RoutingContext, RoutingOptions -from .typing import ( +from sqlargon.mixins import CreatedUpdatedMixin +from sqlargon.orm import Base, Model +from sqlargon.query_builder import Option, UnsupportedOption +from sqlargon.registry import get_default_database +from sqlargon.routing import RoutingContext, RoutingOptions +from sqlargon.typing import ( MultipleValues, OnConflict, OnConflictOptions, @@ -54,10 +52,26 @@ from sqlalchemy.sql.selectable import TypedReturnsRows from typing_extensions import Self, Unpack + from sqlargon.cluster import AnyDatabase from sqlargon.database import Database + from sqlargon.query_builder import QueryBuilder - from .cluster import AnyDatabase - from .query_builder import QueryBuilder +__all__ = ["SQLAlchemyRepository"] + + +def _as_result(rows: Sequence[object]) -> IteratorResult[Any]: + """Re-wrap already fetched rows as a single-column result. + + A hook that has to consume the result of the write it wraps -- to capture + the written rows, or because it built them itself -- hands them back in a + fresh result rather than the exhausted one. + """ + metadata = SimpleResultMetaData(("value",)) + return IteratorResult(metadata, iter([(row,) for row in rows])) + + +def _as_scalars(rows: Sequence[Model]) -> ScalarResult[Model]: + return cast("ScalarResult[Model]", _as_result(rows).scalars()) class SQLAlchemyRepository(Generic[Model]): @@ -145,10 +159,16 @@ def _get_default_index_elements(cls) -> set[str]: @classmethod def _get_default_set(cls) -> set[str]: + # a conflicting row was created when it was created, so ``created_at`` + # keeps the value it was inserted with -- rewriting it would also make + # ``is_new`` report an upserted row as freshly inserted + untouched = cls._get_default_index_elements() + if issubclass(cls.model, CreatedUpdatedMixin): + untouched = untouched | {"created_at"} return { c.name for c in cls.model.__table__.columns # type: ignore[attr-defined] - if c.name not in cls._get_default_index_elements() + if c.name not in untouched } @property @@ -688,374 +708,3 @@ async def get_chunk_for_update( ) else: await self.remove(pk.in_([getattr(r, on_) for r in results])) - - -class DeletedRowExistsError(RuntimeError): - """A tombstoned row holds the unique key a new row was to be created with.""" - - -class SoftDeleteRepository(SQLAlchemyRepository[SoftDeleteModel], abstract=True): - """Repository that tombstones rows instead of deleting them. - - Every statement is scoped to live rows: selects and :meth:`count` skip - tombstoned rows, :meth:`update` refuses to touch them, and :meth:`delete` - is rewritten into an update raising the flag -- so ``remove``, - ``delete_one`` and ``delete_many`` all soft delete:: - - class User(UUIDModelMixin, SoftDeleteBase): ... - - - class UserRepository(SoftDeleteRepository[User]): ... - - - users = UserRepository() - - await users.remove(User.id == user_id) # UPDATE ... SET tombstone = true - await users.list() # the row is gone from reads - await users.restore(User.id == user_id) # and back again - - The flag belongs to :meth:`delete` and :meth:`restore` alone: it is left - out of the default ``ON CONFLICT DO UPDATE`` set, so an upsert cannot - silently resurrect a deleted row. Reach past the scope with - :meth:`with_deleted`, :meth:`only_deleted` and :meth:`hard_delete`. - - The model type variable is bound to - :class:`~sqlargon.mixins.SoftDeleteBase`, so a type checker rejects a - model that cannot be soft deleted. At runtime the looser - :class:`~sqlargon.mixins.SoftDeleteMixin` is enough; anything else raises - ``TypeError`` on subclassing. - """ - - __slots__ = ("deleted_only", "include_deleted") - - def __init__(self) -> None: - super().__init__() - self.include_deleted = False - self.deleted_only = False - - def __init_subclass__(cls, *, abstract: bool = False, **kwargs: Any) -> None: - super().__init_subclass__(abstract=abstract, **kwargs) - if not abstract and not issubclass(cls.model, SoftDeleteMixin): - msg = ( - f"{cls.model.__name__} must inherit from SoftDeleteMixin " - f"to be used with {cls.__name__}" - ) - raise TypeError(msg) - - @classmethod - def _get_default_set(cls) -> set[str]: - return super()._get_default_set() - {"tombstone"} - - @property - def _not_deleted(self) -> _ColumnExpressionArgument[bool]: - return self.model.not_deleted - - @property - def _is_deleted(self) -> _ColumnExpressionArgument[bool]: - return self.model.is_deleted - - @property - def _scope(self) -> _ColumnExpressionArgument[bool] | None: - """The predicate every statement is narrowed to, if any.""" - if self.deleted_only: - return self._is_deleted - if self.include_deleted: - return None - return self._not_deleted - - def _scoped(self, query: Any) -> Any: - """Narrow ``query`` to the rows this repository is scoped to.""" - scope = self._scope - if scope is None: - return query - return query.where(scope) - - def copy(self, query: Any) -> Self: - clone = super().copy(query) - clone.include_deleted = self.include_deleted - clone.deleted_only = self.deleted_only - return clone - - @property - def query(self) -> Any: - if self._query is None: - self.use_query(self._scoped(super().query)) - return self._query - - def select( - self, - *args: Any, - with_for_update: bool | WithForUpdate | None = None, - options: tuple[Any, ...] | None = None, - ) -> Self: - clone = super().select(*args, with_for_update=with_for_update, options=options) - return clone.use_query(self._scoped(clone.query)) - - def update(self, values: Values, *, return_results: bool = False) -> Self: - clone = super().update(values, return_results=return_results) - return clone.use_query(self._scoped(clone.query)) - - def delete(self, *, return_results: bool = False) -> Self: - """Raise the tombstone on the matched rows instead of removing them.""" - return self.update({"tombstone": True}, return_results=return_results) - - async def count(self, *args: _ColumnExpressionArgument[bool], **kwargs: Any) -> int: - scope = self._scope - if scope is not None: - args = (*args, scope) - return await super().count(*args, **kwargs) - - async def bulk_update( - self, - values: MultipleValues, - *args: Any, - on_: set[str] | None = None, - **kwargs: Any, - ) -> None: - scope = self._scope - if scope is not None: - args = (*args, scope) - await super().bulk_update(values, *args, on_=on_, **kwargs) - - async def get_or_create( - self, defaults: SingleValue | None = None, **kwargs: Any - ) -> SoftDeleteModel: - """Get the live row matching ``kwargs`` or create it. - - Raises :class:`DeletedRowExistsError` when the row exists but is - tombstoned: creating it would violate the unique key, and returning it - would resurrect a deleted row behind the caller's back. - """ - async with self.session(): - values = {**(defaults or {}), **kwargs} - result = await self._insert_returning([values], do="ignore") - created = result.one_or_none() - if created is not None: - return created - obj = await self.select().filter(**kwargs).one_or_none() - if obj is None: - msg = ( - f"A deleted {self.model.__name__} row already matches " - f"{kwargs}; restore or hard delete it first" - ) - raise DeletedRowExistsError(msg) - return obj - - def with_deleted(self) -> Self: - """Return a copy whose statements cover tombstoned rows as well.""" - clone = self.copy(self._query) - clone.include_deleted = True - clone.deleted_only = False - return clone - - def only_deleted(self) -> Self: - """Return a copy scoped to tombstoned rows. - - The scope holds for every statement the copy builds, so - ``only_deleted().count()`` counts the trash and - ``only_deleted().hard_delete()`` empties it. - """ - clone = self.copy(self._query) - clone.include_deleted = True - clone.deleted_only = True - return clone - - async def restore( - self, *args: _ColumnExpressionArgument[bool], **kwargs: Any - ) -> Sequence[SoftDeleteModel]: - """Clear the tombstone on the matched rows and return them.""" - return await self.only_deleted().update_many( - {"tombstone": False}, *args, **kwargs - ) - - async def hard_delete( - self, *args: _ColumnExpressionArgument[bool], **kwargs: Any - ) -> None: - """Physically delete the matched rows, bypassing the tombstone.""" - if self.deleted_only: - args = (*args, self._is_deleted) - await super().delete().filter(*args, **kwargs).execute() - - -class ConcurrentModificationError(RuntimeError): - """An update or delete matched zero rows — the version was stale.""" - - -class VersionedRepository(SQLAlchemyRepository[VersionedModel], abstract=True): - """Repository that auto-increments the version column on every update - and offers optimistic concurrency checks via - :meth:`update_if_match` / :meth:`delete_if_match`. - - Works with both :class:`~sqlargon.mixins.UUIDVersionedMixin` (UUID - column, client-managed) and - :class:`~sqlargon.mixins.XminVersionedMixin` (PostgreSQL ``xmin``, - server-managed) — the strategy is read from the SQLAlchemy mapper at - class creation time:: - - class User(UUIDModelMixin, VersionedBase): ... - - - class UserRepository(VersionedRepository[User]): ... - - - users = UserRepository() - - user = await users.get(id=user_id) - updated = await users.update_if_match( - {"name": "jane"}, User.id == user_id, expected_version=user.version_id - ) - if updated is None: - # someone else modified the row first - - The version column is left out of the default ``ON CONFLICT DO - UPDATE`` set, so an upsert cannot silently clobber a version. Regular - ``update_one`` / ``update_many`` / ``bulk_update`` all auto-increment - the version but do **not** check it — use ``update_if_match`` / - ``delete_if_match`` for the optimistic guard, or add a manual - ``Model.version_id == expected`` filter to any method. - - The model type variable is bound to - :class:`~sqlargon.orm.VersionedBase`, so a type checker rejects a - model that cannot be versioned. At runtime the looser - :class:`~sqlargon.mixins.VersionedMixin` is enough; anything else - raises ``TypeError`` on subclassing. - """ - - def __init_subclass__(cls, *, abstract: bool = False, **kwargs: Any) -> None: - super().__init_subclass__(abstract=abstract, **kwargs) - if abstract: - return - if not issubclass(cls.model, VersionedMixin): - msg = ( - f"{cls.model.__name__} must inherit from VersionedMixin " - f"(UUIDVersionedMixin or XminVersionedMixin) to be used " - f"with {cls.__name__}" - ) - raise TypeError(msg) - if cls._version_col() is None: - msg = ( - f"{cls.model.__name__} maps no version column: its " - f"__mapper_args__ carry no version_id_col, so " - f"{cls.__name__} could not tell a stale row from a fresh one" - ) - raise TypeError(msg) - - @classmethod - def _version_col(cls) -> Any: - """The version column from the mapper, or ``None``.""" - return getattr(cls.model.__mapper__, "version_id_col", None) - - @classmethod - def _version_generator(cls) -> Any: - """The version generator: ``False`` (server-side), ``True`` - (integer increment), or a callable.""" - return getattr(cls.model.__mapper__, "version_id_generator", True) - - @classmethod - def _is_server_versioned(cls) -> bool: - return cls._version_generator() is False - - @classmethod - def _get_default_set(cls) -> set[str]: - col = cls._version_col() - excluded = {col.name} if col is not None else set() - return super()._get_default_set() - excluded - - def _with_version_increment(self, values: Values) -> Values: - """Return ``values`` with the version column bumped.""" - col = self._version_col() - generator = self._version_generator() - if col is None or generator is False: - return values - name = col.name - if callable(generator): - if isinstance(values, MappingABC): - return {**values, name: generator(None)} - return [{**row, name: generator(None)} for row in values] - # generator is True — integer increment via SQL expression - if isinstance(values, MappingABC): - return {**values, name: col + 1} - # integer increment is not supported for executemany - return values - - def _version_filter(self, expected: Any) -> Any: - """The guard matching ``expected`` against the version column. - - A server managed version is PostgreSQL's ``xmin``, an ``xid`` no - driver binds natively: asyncpg decodes it into an ``int`` while a - version round tripped through a client comes back as a ``str``, and - ``xid = varchar`` is not an operator PostgreSQL has. Comparing the - column as text accepts either. - """ - col = self._version_col() - if self._is_server_versioned(): - return cast(col, Text) == str(expected) - return col == expected - - def update(self, values: Values, *, return_results: bool = False) -> Self: - values = self._with_version_increment(values) - return super().update(values, return_results=return_results) - - async def bulk_update( - self, - values: MultipleValues, - *args: Any, - on_: set[str] | None = None, - **kwargs: Any, - ) -> None: - col = self._version_col() - generator = self._version_generator() - if col is not None and callable(generator): - name = col.name - values = [{**row, name: generator(None)} for row in values] - await super().bulk_update(values, *args, on_=on_, **kwargs) - - async def update_if_match( - self, - values: SingleValue, - *args: _ColumnExpressionArgument[bool], - expected_version: Any, - raise_on_mismatch: bool = False, - **kwargs: Any, - ) -> VersionedModel | None: - """Update with ``WHERE version_col = expected_version``. - - Returns the updated model, or ``None`` if no row matched (the - version was stale or the row is gone). Raises - :class:`ConcurrentModificationError` when ``raise_on_mismatch`` is - ``True`` and no row matched. - """ - filters = (*args, self._version_filter(expected_version)) - result = await self._update_returning(values, *filters, **kwargs) - row = result.one_or_none() - if row is None and raise_on_mismatch: - msg = ( - f"{self.model.__name__} with version {expected_version!r} " - "was modified or deleted by a concurrent transaction" - ) - raise ConcurrentModificationError(msg) - return row - - async def delete_if_match( - self, - *args: _ColumnExpressionArgument[bool], - expected_version: Any, - raise_on_mismatch: bool = False, - **kwargs: Any, - ) -> VersionedModel | None: - """Delete with ``WHERE version_col = expected_version``. - - Returns the deleted model, or ``None`` if no row matched. Raises - :class:`ConcurrentModificationError` when ``raise_on_mismatch`` is - ``True`` and no row matched. - """ - filters = (*args, self._version_filter(expected_version)) - result = await self._delete_returning(*filters, **kwargs) - row = result.one_or_none() - if row is None and raise_on_mismatch: - msg = ( - f"{self.model.__name__} with version {expected_version!r} " - "was modified or deleted by a concurrent transaction" - ) - raise ConcurrentModificationError(msg) - return row diff --git a/sqlargon/repository/soft_delete.py b/sqlargon/repository/soft_delete.py new file mode 100644 index 0000000..7d4170c --- /dev/null +++ b/sqlargon/repository/soft_delete.py @@ -0,0 +1,218 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from sqlargon.mixins import SoftDeleteMixin +from sqlargon.orm import SoftDeleteModel + +from .base import SQLAlchemyRepository + +if TYPE_CHECKING: + from collections.abc import Sequence + + from sqlalchemy.sql._typing import _ColumnExpressionArgument + from typing_extensions import Self + + from sqlargon.typing import MultipleValues, SingleValue, Values, WithForUpdate + +__all__ = ["DeletedRowExistsError", "SoftDeleteRepository"] + + +class DeletedRowExistsError(RuntimeError): + """A tombstoned row holds the unique key a new row was to be created with.""" + + +class SoftDeleteRepository(SQLAlchemyRepository[SoftDeleteModel], abstract=True): + """Repository that tombstones rows instead of deleting them. + + Every statement is scoped to live rows: selects and :meth:`count` skip + tombstoned rows, :meth:`update` refuses to touch them, and :meth:`delete` + is rewritten into an update raising the flag -- so ``remove``, + ``delete_one`` and ``delete_many`` all soft delete:: + + class User(UUIDModelMixin, SoftDeleteBase): ... + + + class UserRepository(SoftDeleteRepository[User]): ... + + + users = UserRepository() + + await users.remove(User.id == user_id) # UPDATE ... SET tombstone = true + await users.list() # the row is gone from reads + await users.restore(User.id == user_id) # and back again + + The flag belongs to :meth:`delete` and :meth:`restore` alone: it is left + out of the default ``ON CONFLICT DO UPDATE`` set, so an upsert cannot + silently resurrect a deleted row. Reach past the scope with + :meth:`with_deleted`, :meth:`only_deleted` and :meth:`hard_delete`. + + The model type variable is bound to + :class:`~sqlargon.mixins.SoftDeleteBase`, so a type checker rejects a + model that cannot be soft deleted. At runtime the looser + :class:`~sqlargon.mixins.SoftDeleteMixin` is enough; anything else raises + ``TypeError`` on subclassing. + """ + + __slots__ = ("deleted_only", "include_deleted") + + def __init__(self) -> None: + super().__init__() + self.include_deleted = False + self.deleted_only = False + + def __init_subclass__(cls, *, abstract: bool = False, **kwargs: Any) -> None: + super().__init_subclass__(abstract=abstract, **kwargs) + if abstract: + return + mixin, name = cls._required_mixin() + if not issubclass(cls.model, mixin): + msg = ( + f"{cls.model.__name__} must inherit from {name} " + f"to be used with {cls.__name__}" + ) + raise TypeError(msg) + + @classmethod + def _required_mixin(cls) -> tuple[type, str]: + """The mixin a model must carry, and how to name it when it does not. + + A subclass narrowing the requirement -- as + :class:`~sqlargon.repository.AuditableRepository` does -- overrides + this so the error names the mixin its own models need. + """ + return SoftDeleteMixin, "SoftDeleteMixin" + + @classmethod + def _get_default_set(cls) -> set[str]: + return super()._get_default_set() - {"tombstone"} + + @property + def _not_deleted(self) -> _ColumnExpressionArgument[bool]: + return self.model.not_deleted + + @property + def _is_deleted(self) -> _ColumnExpressionArgument[bool]: + return self.model.is_deleted + + @property + def _scope(self) -> _ColumnExpressionArgument[bool] | None: + """The predicate every statement is narrowed to, if any.""" + if self.deleted_only: + return self._is_deleted + if self.include_deleted: + return None + return self._not_deleted + + def _scoped(self, query: Any) -> Any: + """Narrow ``query`` to the rows this repository is scoped to.""" + scope = self._scope + if scope is None: + return query + return query.where(scope) + + def copy(self, query: Any) -> Self: + clone = super().copy(query) + clone.include_deleted = self.include_deleted + clone.deleted_only = self.deleted_only + return clone + + @property + def query(self) -> Any: + if self._query is None: + self.use_query(self._scoped(super().query)) + return self._query + + def select( + self, + *args: Any, + with_for_update: bool | WithForUpdate | None = None, + options: tuple[Any, ...] | None = None, + ) -> Self: + clone = super().select(*args, with_for_update=with_for_update, options=options) + return clone.use_query(self._scoped(clone.query)) + + def update(self, values: Values, *, return_results: bool = False) -> Self: + clone = super().update(values, return_results=return_results) + return clone.use_query(self._scoped(clone.query)) + + def delete(self, *, return_results: bool = False) -> Self: + """Raise the tombstone on the matched rows instead of removing them.""" + return self.update({"tombstone": True}, return_results=return_results) + + async def count(self, *args: _ColumnExpressionArgument[bool], **kwargs: Any) -> int: + scope = self._scope + if scope is not None: + args = (*args, scope) + return await super().count(*args, **kwargs) + + async def bulk_update( + self, + values: MultipleValues, + *args: Any, + on_: set[str] | None = None, + **kwargs: Any, + ) -> None: + scope = self._scope + if scope is not None: + args = (*args, scope) + await super().bulk_update(values, *args, on_=on_, **kwargs) + + async def get_or_create( + self, defaults: SingleValue | None = None, **kwargs: Any + ) -> SoftDeleteModel: + """Get the live row matching ``kwargs`` or create it. + + Raises :class:`DeletedRowExistsError` when the row exists but is + tombstoned: creating it would violate the unique key, and returning it + would resurrect a deleted row behind the caller's back. + """ + async with self.session(): + values = {**(defaults or {}), **kwargs} + result = await self._insert_returning([values], do="ignore") + created = result.one_or_none() + if created is not None: + return created + obj = await self.select().filter(**kwargs).one_or_none() + if obj is None: + msg = ( + f"A deleted {self.model.__name__} row already matches " + f"{kwargs}; restore or hard delete it first" + ) + raise DeletedRowExistsError(msg) + return obj + + def with_deleted(self) -> Self: + """Return a copy whose statements cover tombstoned rows as well.""" + clone = self.copy(self._query) + clone.include_deleted = True + clone.deleted_only = False + return clone + + def only_deleted(self) -> Self: + """Return a copy scoped to tombstoned rows. + + The scope holds for every statement the copy builds, so + ``only_deleted().count()`` counts the trash and + ``only_deleted().hard_delete()`` empties it. + """ + clone = self.copy(self._query) + clone.include_deleted = True + clone.deleted_only = True + return clone + + async def restore( + self, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> Sequence[SoftDeleteModel]: + """Clear the tombstone on the matched rows and return them.""" + return await self.only_deleted().update_many( + {"tombstone": False}, *args, **kwargs + ) + + async def hard_delete( + self, *args: _ColumnExpressionArgument[bool], **kwargs: Any + ) -> None: + """Physically delete the matched rows, bypassing the tombstone.""" + if self.deleted_only: + args = (*args, self._is_deleted) + await super().delete().filter(*args, **kwargs).execute() diff --git a/sqlargon/repository/versioned.py b/sqlargon/repository/versioned.py new file mode 100644 index 0000000..e216a2a --- /dev/null +++ b/sqlargon/repository/versioned.py @@ -0,0 +1,211 @@ +from __future__ import annotations + +from collections.abc import Mapping as MappingABC +from typing import TYPE_CHECKING, Any + +from sqlalchemy import Text, cast + +from sqlargon.mixins import VersionedMixin +from sqlargon.orm import VersionedModel + +from .base import SQLAlchemyRepository + +if TYPE_CHECKING: + from sqlalchemy.sql._typing import _ColumnExpressionArgument + from typing_extensions import Self + + from sqlargon.typing import MultipleValues, SingleValue, Values + +__all__ = ["ConcurrentModificationError", "VersionedRepository"] + + +class ConcurrentModificationError(RuntimeError): + """An update or delete matched zero rows — the version was stale.""" + + +class VersionedRepository(SQLAlchemyRepository[VersionedModel], abstract=True): + """Repository that auto-increments the version column on every update + and offers optimistic concurrency checks via + :meth:`update_if_match` / :meth:`delete_if_match`. + + Works with both :class:`~sqlargon.mixins.UUIDVersionedMixin` (UUID + column, client-managed) and + :class:`~sqlargon.mixins.XminVersionedMixin` (PostgreSQL ``xmin``, + server-managed) — the strategy is read from the SQLAlchemy mapper at + class creation time:: + + class User(UUIDModelMixin, VersionedBase): ... + + + class UserRepository(VersionedRepository[User]): ... + + + users = UserRepository() + + user = await users.get(id=user_id) + updated = await users.update_if_match( + {"name": "jane"}, User.id == user_id, expected_version=user.version_id + ) + if updated is None: + # someone else modified the row first + + The version column is left out of the default ``ON CONFLICT DO + UPDATE`` set, so an upsert cannot silently clobber a version. Regular + ``update_one`` / ``update_many`` / ``bulk_update`` all auto-increment + the version but do **not** check it — use ``update_if_match`` / + ``delete_if_match`` for the optimistic guard, or add a manual + ``Model.version_id == expected`` filter to any method. + + The model type variable is bound to + :class:`~sqlargon.orm.VersionedBase`, so a type checker rejects a + model that cannot be versioned. At runtime the looser + :class:`~sqlargon.mixins.VersionedMixin` is enough; anything else + raises ``TypeError`` on subclassing. + """ + + def __init_subclass__(cls, *, abstract: bool = False, **kwargs: Any) -> None: + super().__init_subclass__(abstract=abstract, **kwargs) + if abstract: + return + mixin, name = cls._required_mixin() + if not issubclass(cls.model, mixin): + msg = ( + f"{cls.model.__name__} must inherit from {name} " + f"to be used with {cls.__name__}" + ) + raise TypeError(msg) + if cls._version_col() is None: + msg = ( + f"{cls.model.__name__} maps no version column: its " + f"__mapper_args__ carry no version_id_col, so " + f"{cls.__name__} could not tell a stale row from a fresh one" + ) + raise TypeError(msg) + + @classmethod + def _required_mixin(cls) -> tuple[type, str]: + """The mixin a model must carry, and how to name it when it does not.""" + return ( + VersionedMixin, + "VersionedMixin (UUIDVersionedMixin or XminVersionedMixin)", + ) + + @classmethod + def _version_col(cls) -> Any: + """The version column from the mapper, or ``None``.""" + return getattr(cls.model.__mapper__, "version_id_col", None) + + @classmethod + def _version_generator(cls) -> Any: + """The version generator: ``False`` (server-side), ``True`` + (integer increment), or a callable.""" + return getattr(cls.model.__mapper__, "version_id_generator", True) + + @classmethod + def _is_server_versioned(cls) -> bool: + return cls._version_generator() is False + + @classmethod + def _get_default_set(cls) -> set[str]: + col = cls._version_col() + excluded = {col.name} if col is not None else set() + return super()._get_default_set() - excluded + + def _with_version_increment(self, values: Values) -> Values: + """Return ``values`` with the version column bumped.""" + col = self._version_col() + generator = self._version_generator() + if col is None or generator is False: + return values + name = col.name + if callable(generator): + if isinstance(values, MappingABC): + return {**values, name: generator(None)} + return [{**row, name: generator(None)} for row in values] + # generator is True — integer increment via SQL expression + if isinstance(values, MappingABC): + return {**values, name: col + 1} + # integer increment is not supported for executemany + return values + + def _version_filter(self, expected: Any) -> Any: + """The guard matching ``expected`` against the version column. + + A server managed version is PostgreSQL's ``xmin``, an ``xid`` no + driver binds natively: asyncpg decodes it into an ``int`` while a + version round tripped through a client comes back as a ``str``, and + ``xid = varchar`` is not an operator PostgreSQL has. Comparing the + column as text accepts either. + """ + col = self._version_col() + if self._is_server_versioned(): + return cast(col, Text) == str(expected) + return col == expected + + def update(self, values: Values, *, return_results: bool = False) -> Self: + values = self._with_version_increment(values) + return super().update(values, return_results=return_results) + + async def bulk_update( + self, + values: MultipleValues, + *args: Any, + on_: set[str] | None = None, + **kwargs: Any, + ) -> None: + col = self._version_col() + generator = self._version_generator() + if col is not None and callable(generator): + name = col.name + values = [{**row, name: generator(None)} for row in values] + await super().bulk_update(values, *args, on_=on_, **kwargs) + + async def update_if_match( + self, + values: SingleValue, + *args: _ColumnExpressionArgument[bool], + expected_version: Any, + raise_on_mismatch: bool = False, + **kwargs: Any, + ) -> VersionedModel | None: + """Update with ``WHERE version_col = expected_version``. + + Returns the updated model, or ``None`` if no row matched (the + version was stale or the row is gone). Raises + :class:`ConcurrentModificationError` when ``raise_on_mismatch`` is + ``True`` and no row matched. + """ + filters = (*args, self._version_filter(expected_version)) + result = await self._update_returning(values, *filters, **kwargs) + row = result.one_or_none() + if row is None and raise_on_mismatch: + msg = ( + f"{self.model.__name__} with version {expected_version!r} " + "was modified or deleted by a concurrent transaction" + ) + raise ConcurrentModificationError(msg) + return row + + async def delete_if_match( + self, + *args: _ColumnExpressionArgument[bool], + expected_version: Any, + raise_on_mismatch: bool = False, + **kwargs: Any, + ) -> VersionedModel | None: + """Delete with ``WHERE version_col = expected_version``. + + Returns the deleted model, or ``None`` if no row matched. Raises + :class:`ConcurrentModificationError` when ``raise_on_mismatch`` is + ``True`` and no row matched. + """ + filters = (*args, self._version_filter(expected_version)) + result = await self._delete_returning(*filters, **kwargs) + row = result.one_or_none() + if row is None and raise_on_mismatch: + msg = ( + f"{self.model.__name__} with version {expected_version!r} " + "was modified or deleted by a concurrent transaction" + ) + raise ConcurrentModificationError(msg) + return row diff --git a/sqlargon/types/json.py b/sqlargon/types/json.py index 6279ff1..44225d7 100644 --- a/sqlargon/types/json.py +++ b/sqlargon/types/json.py @@ -1,13 +1,106 @@ -from typing import Any +from __future__ import annotations + +from typing import TYPE_CHECKING, Any import sqlalchemy as sa -from sqlalchemy import BOOLEAN, Dialect, FunctionElement, TypeDecorator +from sqlalchemy import Dialect, FunctionElement, TypeDecorator from sqlalchemy.dialects import postgresql, sqlite from sqlalchemy.ext.compiler import compiles -from sqlalchemy.types import TypeEngine +from sqlalchemy.sql import coercions, roles +from sqlalchemy.sql.elements import Grouping from sqlargon.utils import json_dumps +if TYPE_CHECKING: + from collections.abc import Callable, Mapping + + from sqlalchemy.sql.elements import ColumnElement + from sqlalchemy.types import TypeEngine + +__all__ = [ + "JSON", + "json_array_append", + "json_array_length", + "json_contains", + "json_get", + "json_has_all_keys", + "json_has_any_key", + "json_has_key", + "json_insert_key", + "json_keys", + "json_remove_key", + "json_replace_key", + "json_set_key", + "json_update", + "json_value", +] + + +def _json_path(key: str) -> str: + """``key`` as a top level JSON path, quotes and backslashes escaped.""" + escaped = key.replace("\\", "\\\\").replace('"', '\\"') + return f'$."{escaped}"' + + +def _operands(element: FunctionElement[Any]) -> tuple[ColumnElement[Any], ...]: + """The element's operands, read off its own clause list. + + Two reasons never to reach for an attribute cached on the instance. + Clone and adapt machinery -- ``ClauseAdapter``, ``with_loader_criteria``, + ``with_polymorphic`` -- rewrites only the traversed ``clauses``, so a + cached operand would still point at the pre-adaption column. And keys, + paths and values are bind parameters in this list rather than literals + baked in at compile time, which is what lets one cached statement serve + every key: a bind created inside a ``@compiles`` hook is invisible to the + statement cache, so the first key compiled would be reused for the rest. + """ + return tuple(element.clauses) + + +def _text(value: str) -> ColumnElement[str]: + return sa.literal(value, sa.String()) + + +def _json(value: Any) -> ColumnElement[Any]: + """``value`` as a JSON operand: an existing expression, or a bind.""" + return coercions.expect(roles.ExpressionElementRole, value, type_=JSON()) + + +def _sqlite_json(value: ColumnElement[Any]) -> ColumnElement[Any]: + """Re-parse a bound JSON string into a JSON value. + + Without it ``json_set`` would store the serialized text as a JSON + *string* rather than as the document it represents. + """ + return sa.func.json(value) + + +def _mysql_json(value: ColumnElement[Any]) -> ColumnElement[Any]: + """As :func:`_sqlite_json`, for the MySQL family. + + ``json_extract(:v, '$')`` rather than ``CAST(:v AS JSON)`` because + MariaDB's ``JSON`` is a ``LONGTEXT`` alias whose cast support differs + from MySQL's, while both spell the whole-document extract this way. + """ + return sa.func.json_extract(value, _text("$")) + + +def _pg_json(value: ColumnElement[Any]) -> ColumnElement[Any]: + return sa.cast(value, postgresql.JSONB) + + +def _pg_group(expression: ColumnElement[Any]) -> ColumnElement[Any]: + """Parenthesize a PostgreSQL JSON operator expression. + + These elements compile to infix operators, but to an enclosing compiler + they are opaque functions with no precedence to reason about. Nesting a + removal around a merge would otherwise emit ``a || b - c``, which + PostgreSQL reads as ``a || (b - c)`` -- binary ``-`` binds tighter than + ``||`` -- and so would drop the key from the patch instead of from the + merged document. + """ + return Grouping(expression) + class json_contains(FunctionElement): """ @@ -18,7 +111,7 @@ class json_contains(FunctionElement): https://www.postgresql.org/docs/current/functions-json.html """ - type = BOOLEAN + type = sa.Boolean() name = "json_contains" inherit_cache = False @@ -77,6 +170,7 @@ def _json_contains_sqlite(element: json_contains, compiler: Any, **kwargs: Any) @compiles(json_contains, "mysql") +@compiles(json_contains) def _json_contains_mysql(element: json_contains, compiler: Any, **kwargs: Any) -> str: return compiler.process( sa.func.json_contains( @@ -95,7 +189,7 @@ class json_has_any_key(FunctionElement): https://www.postgresql.org/docs/current/functions-json.html """ - type: Any = BOOLEAN + type: Any = sa.Boolean() name = "json_has_any_key" inherit_cache = False @@ -162,7 +256,7 @@ class json_has_all_keys(FunctionElement): https://www.postgresql.org/docs/current/functions-json.html """ - type: Any = BOOLEAN + type: Any = sa.Boolean() name = "json_has_all_keys" inherit_cache = False @@ -220,7 +314,7 @@ def _json_has_all_keys_mysql( ) -class json_value(FunctionElement): +class json_value(FunctionElement[str]): """Portable ``->>`` operator: text value at a JSON object key.""" name = "json_value" @@ -230,23 +324,21 @@ class json_value(FunctionElement): def __init__(self, column: Any, key: str) -> None: self.column = column self.key = key - super().__init__(column) + super().__init__(column, _text(key), _text(_json_path(key))) @compiles(json_value, "postgresql") def _json_value_postgresql(element: json_value, compiler: Any, **kwargs: Any) -> str: + column, key, _path = _operands(element) return compiler.process( - sa.type_coerce(element.column, postgresql.JSONB).op("->>")(element.key), - **kwargs, + sa.type_coerce(column, postgresql.JSONB).op("->>")(key), **kwargs ) @compiles(json_value) def _json_value_default(element: json_value, compiler: Any, **kwargs: Any) -> str: - return compiler.process( - sa.func.json_extract(element.column, sa.literal(f'$."{element.key}"')), - **kwargs, - ) + column, _key, path = _operands(element) + return compiler.process(sa.func.json_extract(column, path), **kwargs) class JSON(TypeDecorator): @@ -268,7 +360,43 @@ def load_dialect_impl(self, dialect: Dialect) -> TypeEngine[Any]: return dialect.type_descriptor(sqlite.JSON(none_as_null=True)) return dialect.type_descriptor(sa.JSON(none_as_null=True)) + def literal_processor(self, dialect: Dialect) -> Callable[[Any], str]: + """Render the value as an inline SQL string literal. + + Only reached under ``literal_binds``, so when a statement is being + printed or logged rather than executed. SQLAlchemy's own JSON types + ship no literal renderer, so without this a statement carrying a + JSON bind cannot be compiled at all. The serialized text is quoted by + the dialect's own string renderer rather than by hand, which is what + gets the backslash escaping right on MySQL. + """ + quote = sa.String().literal_processor(dialect) + if quote is None: # pragma: no cover - every dialect renders strings + msg = f"{dialect.name} cannot render a string literal" + raise NotImplementedError(msg) + + def process(value: Any) -> str: + if value is None: + return "NULL" + return quote(json_dumps(value)) + + return process + class ComparatorFactory(sa.JSON.Comparator): + """Type-specific SQL expression methods for a JSON column. + + Despite the name, this hook is not limited to comparisons -- the + ``Comparator`` base already carries ``concat``, ``collate``, + ``distinct`` and the arithmetic operators. The mutation methods below + follow ``postgresql.JSONB.Comparator.delete_path`` and the HSTORE + comparator's ``delete`` / ``slice`` / ``keys`` / ``vals``, which + likewise answer with a rewritten container rather than a boolean. + + Nothing here mutates anything: each method returns an expression that + is inert until it lands in an ``UPDATE``, at which point the document + is rewritten by the server. + """ + def contains(self, other: Any, **_kw: Any) -> json_contains: return json_contains(self, other) @@ -281,4 +409,465 @@ def has_all_keys(self, other: Any) -> json_has_all_keys: def json_value(self, other: Any) -> json_value: return json_value(self, other) + def has_key(self, other: str) -> json_has_key: + return json_has_key(self, other) + + def get(self, other: str) -> json_get: + return json_get(self, other) + + def keys(self) -> json_keys: + return json_keys(self) + + def array_length(self) -> json_array_length: + return json_array_length(self) + + def update(self, other: Mapping[str, Any]) -> json_update: + return json_update(self, other) + + def set_key(self, key: str, value: Any) -> json_update: + return json_set_key(self, key, value) + + def remove_key(self, *keys: str) -> json_remove_key: + return json_remove_key(self, *keys) + + def insert_key(self, key: str, value: Any) -> json_insert_key: + return json_insert_key(self, key, value) + + def replace_key(self, key: str, value: Any) -> json_replace_key: + return json_replace_key(self, key, value) + + def array_append(self, value: Any) -> json_array_append: + return json_array_append(self, value) + comparator_factory = ComparatorFactory + + +# --- mutations ------------------------------------------------------------- +# +# Every element below is typed ``JSON()``, so they nest freely: +# ``json_remove_key(json_update(col, {"a": 1}), "b")``. A SQL NULL column +# propagates -- ``jsonb_set`` and ``JSON_SET`` both return NULL for a NULL +# document and these do not paper over it -- and the ``JSON_SET`` family +# assumes the document is an object. + + +class json_update(FunctionElement[Any]): + """Shallow merge of ``mapping`` into a JSON object. + + Top level keys of ``mapping`` replace their counterparts wholesale; a + nested object is *not* merged recursively. Equivalent to the PostgreSQL + ``||`` operator, which is what the multi-pair ``JSON_SET`` on the other + backends reproduces -- deliberately not ``json_patch`` / + ``JSON_MERGE_PATCH``, whose recursive semantics PostgreSQL cannot express + without a recursive query. + """ + + name = "json_update" + type = JSON() + inherit_cache = True + + def __init__(self, column: Any, mapping: Mapping[str, Any]) -> None: + if not all(isinstance(key, str) for key in mapping): + msg = "json_update keys must be strings" + raise ValueError(msg) + pairs: list[ColumnElement[Any]] = [] + for key, value in mapping.items(): + pairs += [_text(_json_path(key)), _json(value)] + # the whole mapping rides along as a single operand for the + # PostgreSQL concat, the per-key pairs for the JSON_SET backends + super().__init__(_json(column), _json(dict(mapping)), *pairs) + + +@compiles(json_update, "postgresql") +def _json_update_postgresql(element: json_update, compiler: Any, **kwargs: Any) -> str: + column, mapping, *pairs = _operands(element) + if not pairs: + return compiler.process(column, **kwargs) + merged = sa.type_coerce(column, postgresql.JSONB).op( + "||", return_type=postgresql.JSONB + )(_pg_json(mapping)) + return compiler.process(_pg_group(merged), **kwargs) + + +@compiles(json_update, "sqlite") +def _json_update_sqlite(element: json_update, compiler: Any, **kwargs: Any) -> str: + column, _mapping, *pairs = _operands(element) + if not pairs: + return compiler.process(column, **kwargs) + args: list[ColumnElement[Any]] = [] + for path, value in zip(pairs[::2], pairs[1::2], strict=True): + args += [path, _sqlite_json(value)] + return compiler.process(sa.func.json_set(column, *args), **kwargs) + + +@compiles(json_update, "mysql") +@compiles(json_update) +def _json_update_mysql(element: json_update, compiler: Any, **kwargs: Any) -> str: + column, _mapping, *pairs = _operands(element) + if not pairs: + return compiler.process(column, **kwargs) + args: list[ColumnElement[Any]] = [] + for path, value in zip(pairs[::2], pairs[1::2], strict=True): + args += [path, _mysql_json(value)] + return compiler.process(sa.func.json_set(column, *args), **kwargs) + + +def json_set_key(column: Any, key: str, value: Any) -> json_update: + """Set one top level ``key``, replacing any value already there.""" + return json_update(column, {key: value}) + + +class json_remove_key(FunctionElement[Any]): + """Drop top level ``keys`` from a JSON object. + + On postgres this is the ``-`` operator over a ``text[]`` of keys. + """ + + name = "json_remove_key" + type = JSON() + inherit_cache = True + + def __init__(self, column: Any, *keys: str) -> None: + if not all(isinstance(key, str) for key in keys): + msg = "json_remove_key keys must be strings" + raise ValueError(msg) + super().__init__( + _json(column), + sa.literal(list(keys), postgresql.ARRAY(sa.Text)), + *(_text(_json_path(key)) for key in keys), + ) + + +@compiles(json_remove_key, "postgresql") +def _json_remove_key_postgresql( + element: json_remove_key, compiler: Any, **kwargs: Any +) -> str: + column, keys, *paths = _operands(element) + if not paths: + return compiler.process(column, **kwargs) + removed = sa.type_coerce(column, postgresql.JSONB).op( + "-", return_type=postgresql.JSONB + )(sa.cast(keys, postgresql.ARRAY(sa.Text))) + return compiler.process(_pg_group(removed), **kwargs) + + +@compiles(json_remove_key, "sqlite") +@compiles(json_remove_key, "mysql") +@compiles(json_remove_key) +def _json_remove_key_default( + element: json_remove_key, compiler: Any, **kwargs: Any +) -> str: + column, _keys, *paths = _operands(element) + # json_remove(col) with no path is a syntax error, and removing nothing + # is the column itself + if not paths: + return compiler.process(column, **kwargs) + return compiler.process(sa.func.json_remove(column, *paths), **kwargs) + + +class json_insert_key(FunctionElement[Any]): + """Set ``key`` to ``value``, but only where the key is absent.""" + + name = "json_insert_key" + type = JSON() + inherit_cache = True + + def __init__(self, column: Any, key: str, value: Any) -> None: + super().__init__( + _json(column), _json({key: value}), _text(_json_path(key)), _json(value) + ) + + +@compiles(json_insert_key, "postgresql") +def _json_insert_key_postgresql( + element: json_insert_key, compiler: Any, **kwargs: Any +) -> str: + column, mapping, _path, _value = _operands(element) + # the new key on the left, so an existing one on the right wins + inserted = _pg_json(mapping).op("||", return_type=postgresql.JSONB)( + sa.type_coerce(column, postgresql.JSONB) + ) + return compiler.process(_pg_group(inserted), **kwargs) + + +@compiles(json_insert_key, "sqlite") +def _json_insert_key_sqlite( + element: json_insert_key, compiler: Any, **kwargs: Any +) -> str: + column, _mapping, path, value = _operands(element) + return compiler.process( + sa.func.json_insert(column, path, _sqlite_json(value)), **kwargs + ) + + +@compiles(json_insert_key, "mysql") +@compiles(json_insert_key) +def _json_insert_key_mysql( + element: json_insert_key, compiler: Any, **kwargs: Any +) -> str: + column, _mapping, path, value = _operands(element) + return compiler.process( + sa.func.json_insert(column, path, _mysql_json(value)), **kwargs + ) + + +class json_replace_key(FunctionElement[Any]): + """Set ``key`` to ``value``, but only where the key is already present.""" + + name = "json_replace_key" + type = JSON() + inherit_cache = True + + def __init__(self, column: Any, key: str, value: Any) -> None: + super().__init__( + _json(column), _text(key), _text(_json_path(key)), _json(value) + ) + + +@compiles(json_replace_key, "postgresql") +def _json_replace_key_postgresql( + element: json_replace_key, compiler: Any, **kwargs: Any +) -> str: + column, key, _path, value = _operands(element) + return compiler.process( + sa.func.jsonb_set( + sa.type_coerce(column, postgresql.JSONB), + sa.cast(postgresql.array([key]), postgresql.ARRAY(sa.Text)), + _pg_json(value), + # create_missing => false is what makes this replace only + sa.false(), + ), + **kwargs, + ) + + +@compiles(json_replace_key, "sqlite") +def _json_replace_key_sqlite( + element: json_replace_key, compiler: Any, **kwargs: Any +) -> str: + column, _key, path, value = _operands(element) + return compiler.process( + sa.func.json_replace(column, path, _sqlite_json(value)), **kwargs + ) + + +@compiles(json_replace_key, "mysql") +@compiles(json_replace_key) +def _json_replace_key_mysql( + element: json_replace_key, compiler: Any, **kwargs: Any +) -> str: + column, _key, path, value = _operands(element) + return compiler.process( + sa.func.json_replace(column, path, _mysql_json(value)), **kwargs + ) + + +class json_array_append(FunctionElement[Any]): + """Append ``value`` as one element of a JSON array. + + A list ``value`` is appended as a single nested array, not concatenated, + so the three backends agree. + """ + + name = "json_array_append" + type = JSON() + inherit_cache = True + + def __init__(self, column: Any, value: Any) -> None: + super().__init__(_json(column), _json(value)) + + +@compiles(json_array_append, "postgresql") +def _json_array_append_postgresql( + element: json_array_append, compiler: Any, **kwargs: Any +) -> str: + column, value = _operands(element) + # build_array, not a bare concat: concatenating two arrays would merge + # them instead of appending one element + appended = sa.type_coerce(column, postgresql.JSONB).op( + "||", return_type=postgresql.JSONB + )(sa.func.jsonb_build_array(_pg_json(value))) + return compiler.process(_pg_group(appended), **kwargs) + + +@compiles(json_array_append, "sqlite") +def _json_array_append_sqlite( + element: json_array_append, compiler: Any, **kwargs: Any +) -> str: + column, value = _operands(element) + return compiler.process( + sa.func.json_insert(column, _text("$[#]"), _sqlite_json(value)), **kwargs + ) + + +@compiles(json_array_append, "mysql") +@compiles(json_array_append) +def _json_array_append_mysql( + element: json_array_append, compiler: Any, **kwargs: Any +) -> str: + column, value = _operands(element) + return compiler.process( + sa.func.json_array_append(column, _text("$"), _mysql_json(value)), **kwargs + ) + + +# --- reads ----------------------------------------------------------------- + + +class json_get(FunctionElement[Any]): + """Portable ``->`` operator: the JSON value at a top level key. + + The JSON-typed counterpart of :class:`json_value`, which is ``->>``. + """ + + name = "json_get" + type = JSON() + inherit_cache = True + + def __init__(self, column: Any, key: str) -> None: + super().__init__(_json(column), _text(key), _text(_json_path(key))) + + +@compiles(json_get, "postgresql") +def _json_get_postgresql(element: json_get, compiler: Any, **kwargs: Any) -> str: + column, key, _path = _operands(element) + got = sa.type_coerce(column, postgresql.JSONB).op( + "->", return_type=postgresql.JSONB + )(key) + return compiler.process(_pg_group(got), **kwargs) + + +@compiles(json_get) +def _json_get_default(element: json_get, compiler: Any, **kwargs: Any) -> str: + column, _key, path = _operands(element) + return compiler.process(sa.func.json_extract(column, path), **kwargs) + + +class json_has_key(FunctionElement[bool]): + """Whether a JSON object has ``key`` at the top level. + + On postgres this is the ``?`` existence operator. Unlike + :class:`json_has_any_key` and :class:`json_has_all_keys` this addresses + object *keys* on every backend, SQLite included. + """ + + name = "json_has_key" + type: Any = sa.Boolean() + inherit_cache = True + + def __init__(self, column: Any, key: str) -> None: + super().__init__(_json(column), _text(key), _text(_json_path(key))) + + +@compiles(json_has_key, "postgresql") +def _json_has_key_postgresql( + element: json_has_key, compiler: Any, **kwargs: Any +) -> str: + column, key, _path = _operands(element) + return compiler.process( + sa.type_coerce(column, postgresql.JSONB).has_key(key), **kwargs + ) + + +@compiles(json_has_key, "sqlite") +def _json_has_key_sqlite(element: json_has_key, compiler: Any, **kwargs: Any) -> str: + column, _key, path = _operands(element) + return compiler.process(sa.func.json_type(column, path).is_not(None), **kwargs) + + +@compiles(json_has_key, "mysql") +@compiles(json_has_key) +def _json_has_key_mysql(element: json_has_key, compiler: Any, **kwargs: Any) -> str: + column, _key, path = _operands(element) + return compiler.process( + sa.func.json_contains_path(column, _text("one"), path), **kwargs + ) + + +class json_array_length(FunctionElement[int]): + """How many elements a JSON array has. + + Only arrays are portable here: given an object, postgres raises while + SQLite answers 0 and MySQL answers 1. + """ + + name = "json_array_length" + type = sa.Integer() + inherit_cache = True + + def __init__(self, column: Any) -> None: + super().__init__(_json(column)) + + +@compiles(json_array_length, "postgresql") +def _json_array_length_postgresql( + element: json_array_length, compiler: Any, **kwargs: Any +) -> str: + (column,) = _operands(element) + return compiler.process( + sa.func.jsonb_array_length(sa.type_coerce(column, postgresql.JSONB)), **kwargs + ) + + +@compiles(json_array_length, "sqlite") +def _json_array_length_sqlite( + element: json_array_length, compiler: Any, **kwargs: Any +) -> str: + (column,) = _operands(element) + return compiler.process(sa.func.json_array_length(column), **kwargs) + + +@compiles(json_array_length, "mysql") +@compiles(json_array_length) +def _json_array_length_mysql( + element: json_array_length, compiler: Any, **kwargs: Any +) -> str: + (column,) = _operands(element) + return compiler.process(sa.func.json_length(column), **kwargs) + + +class json_keys(FunctionElement[Any]): + """The top level keys of a JSON object, as a JSON array of strings.""" + + name = "json_keys" + type = JSON() + inherit_cache = True + + def __init__(self, column: Any) -> None: + super().__init__(_json(column)) + + +@compiles(json_keys, "postgresql") +def _json_keys_postgresql(element: json_keys, compiler: Any, **kwargs: Any) -> str: + (column,) = _operands(element) + keys = sa.func.jsonb_object_keys(sa.type_coerce(column, postgresql.JSONB)).alias( + "json_keys" + ) + # jsonb_object_keys is set returning, and jsonb_agg over no rows is NULL + aggregated = ( + sa.select(sa.func.jsonb_agg(sa.literal_column("json_keys"))) + .select_from(keys) + .scalar_subquery() + ) + return compiler.process( + sa.func.coalesce(aggregated, _pg_json(_text("[]"))), **kwargs + ) + + +@compiles(json_keys, "sqlite") +def _json_keys_sqlite(element: json_keys, compiler: Any, **kwargs: Any) -> str: + (column,) = _operands(element) + each = sa.func.json_each(column).alias("json_each") + return compiler.process( + sa.select(sa.func.json_group_array(sa.literal_column("json_each.key"))) + .select_from(each) + .scalar_subquery(), + **kwargs, + ) + + +@compiles(json_keys, "mysql") +@compiles(json_keys) +def _json_keys_mysql(element: json_keys, compiler: Any, **kwargs: Any) -> str: + (column,) = _operands(element) + return compiler.process(sa.func.json_keys(column), **kwargs) diff --git a/sqlargon/types/vector.py b/sqlargon/types/vector.py new file mode 100644 index 0000000..c17bfe7 --- /dev/null +++ b/sqlargon/types/vector.py @@ -0,0 +1,230 @@ +from __future__ import annotations + +import struct +from enum import Enum +from typing import Any, ClassVar + +import sqlalchemy as sa +from sqlalchemy import Dialect, FunctionElement, TypeDecorator +from sqlalchemy.ext.compiler import compiles +from sqlalchemy.sql import coercions, roles +from sqlalchemy.types import TypeEngine + +from sqlargon.query_builder import UnsupportedDialectError + +try: + from pgvector.sqlalchemy import VECTOR +except ImportError as e: + msg = "Vector columns require the 'pgvector' package; install 'sqlargon[vectors]'" + raise ImportError(msg) from e + +__all__ = [ + "DistanceMetric", + "UnsupportedDialectError", + "Vector", + "cosine_distance", + "distance_for", + "l1_distance", + "l2_distance", + "max_inner_product", +] + + +class DistanceMetric(str, Enum): + """Distance metric of a vector similarity search. + + Maps each metric to the pgvector operator, the pgvector index + operator class and the sqlite-vector ``distance`` option. + """ + + COSINE = "cosine" + L2 = "l2" + DOT = "dot" + L1 = "l1" + + @property + def pg_operator(self) -> str: + return _PG_OPERATORS[self] + + @property + def pg_opclass(self) -> str: + return _PG_OPCLASSES[self] + + @property + def sqlite_option(self) -> str: + return _SQLITE_OPTIONS[self] + + +_PG_OPERATORS = { + DistanceMetric.COSINE: "<=>", + DistanceMetric.L2: "<->", + DistanceMetric.DOT: "<#>", + DistanceMetric.L1: "<+>", +} + +_PG_OPCLASSES = { + DistanceMetric.COSINE: "vector_cosine_ops", + DistanceMetric.L2: "vector_l2_ops", + DistanceMetric.DOT: "vector_ip_ops", + DistanceMetric.L1: "vector_l1_ops", +} + +_SQLITE_OPTIONS = { + DistanceMetric.COSINE: "COSINE", + DistanceMetric.L2: "L2", + DistanceMetric.DOT: "DOT", + DistanceMetric.L1: "L1", +} + + +class _vector_distance(FunctionElement): + """Distance between two vectors, ordered ascending by similarity. + + Compiles to the matching pgvector operator on PostgreSQL and raises + :class:`UnsupportedDialectError` elsewhere -- sqlite-vector exposes no + scalar distance function, use + :meth:`~sqlargon.vectors.VectorRepository.search` there. + """ + + type = sa.Float() + metric: ClassVar[DistanceMetric] + inherit_cache = True + + def __init__(self, left: Any, right: Any) -> None: + # both operands are coerced here rather than at compile time, so + # they belong to the element's own clause list -- that is what lets + # the statement be cached and still bind a fresh vector per + # execution. The query vector is typed after the column it is + # compared with, so a plain list of floats binds as a vector. + column = coercions.expect(roles.ExpressionElementRole, left) + super().__init__( + column, + coercions.expect(roles.ExpressionElementRole, right, type_=column.type), + ) + + @property + def operands(self) -> tuple[sa.ColumnElement[Any], sa.ColumnElement[Any]]: + left, right = self.clauses + return left, right + + +class cosine_distance(_vector_distance): + name = "cosine_distance" + metric = DistanceMetric.COSINE + inherit_cache = True + + +class l2_distance(_vector_distance): + name = "l2_distance" + metric = DistanceMetric.L2 + inherit_cache = True + + +class max_inner_product(_vector_distance): + """Negative inner product, so ascending order means most similar first.""" + + name = "max_inner_product" + metric = DistanceMetric.DOT + inherit_cache = True + + +class l1_distance(_vector_distance): + name = "l1_distance" + metric = DistanceMetric.L1 + inherit_cache = True + + +_DISTANCE_ELEMENTS: dict[DistanceMetric, type[_vector_distance]] = { + element.metric: element + for element in (cosine_distance, l2_distance, max_inner_product, l1_distance) +} + + +def distance_for(metric: DistanceMetric) -> type[_vector_distance]: + """The distance element matching ``metric``.""" + return _DISTANCE_ELEMENTS[metric] + + +@compiles(cosine_distance, "postgresql") +@compiles(l2_distance, "postgresql") +@compiles(max_inner_product, "postgresql") +@compiles(l1_distance, "postgresql") +def _distance_postgresql( + element: _vector_distance, compiler: Any, **kwargs: Any +) -> str: + left, right = element.operands + return compiler.process( + left.op(element.metric.pg_operator, return_type=sa.Float)(right), **kwargs + ) + + +@compiles(cosine_distance) +@compiles(l2_distance) +@compiles(max_inner_product) +@compiles(l1_distance) +def _distance_default(element: _vector_distance, compiler: Any, **_kwargs: Any) -> str: + msg = ( + f"{element.name} is not supported on the {compiler.dialect.name!r} dialect; " + "on sqlite use VectorRepository.search()" + ) + raise UnsupportedDialectError(msg) + + +class Vector(TypeDecorator): + """Embedding column of a fixed dimension. + + ``pgvector.VECTOR`` on PostgreSQL, a little-endian float32 BLOB on + SQLite (the layout sqlite-vector's ``FLOAT32`` columns expect) and + plain JSON storage on other dialects. Values are read back as + ``list[float]`` everywhere. + """ + + impl = VECTOR + cache_ok = True + + def __init__(self, dim: int) -> None: + super().__init__() + self.dim = dim + + def load_dialect_impl(self, dialect: Dialect) -> TypeEngine[Any]: + if dialect.name == "postgresql": + return dialect.type_descriptor(VECTOR(self.dim)) + if dialect.name == "sqlite": + return dialect.type_descriptor(sa.LargeBinary()) + return dialect.type_descriptor(sa.JSON(none_as_null=True)) + + def process_bind_param(self, value: Any, dialect: Dialect) -> Any: + if value is None: + return None + if dialect.name == "sqlite": + return struct.pack(f"<{len(value)}f", *value) + if isinstance(value, list): + return value + return [float(item) for item in value] + + def process_result_value(self, value: Any, dialect: Dialect) -> list[float] | None: + if value is None: + return None + if dialect.name == "sqlite": + return list(struct.unpack(f"<{len(value) // 4}f", value)) + if isinstance(value, list): + return value + return [float(item) for item in value] + + class ComparatorFactory(TypeEngine.Comparator): + def cosine_distance(self, other: Any) -> cosine_distance: + return cosine_distance(self.expr, other) + + def l2_distance(self, other: Any) -> l2_distance: + return l2_distance(self.expr, other) + + def max_inner_product(self, other: Any) -> max_inner_product: + return max_inner_product(self.expr, other) + + def l1_distance(self, other: Any) -> l1_distance: + return l1_distance(self.expr, other) + + def distance(self, other: Any, metric: DistanceMetric) -> _vector_distance: + return distance_for(metric)(self.expr, other) + + comparator_factory = ComparatorFactory diff --git a/sqlargon/vectors/__init__.py b/sqlargon/vectors/__init__.py new file mode 100644 index 0000000..cbf9961 --- /dev/null +++ b/sqlargon/vectors/__init__.py @@ -0,0 +1,61 @@ +from sqlargon.types.vector import ( + DistanceMetric, + UnsupportedDialectError, + Vector, + cosine_distance, + l1_distance, + l2_distance, + max_inner_product, +) + +from .loader import init_vectors, register_sqlite_vector +from .mixins import ( + AttributesMixin, + EmbeddingMixin, + TextMixin, + VectorCollectionMixin, +) +from .models import ( + EmbeddingBase, + HybridModel, + TextBase, + TextEmbeddingBase, + TextModel, + VectorCollection, + VectorDocument, + VectorModel, +) +from .repository import ( + HybridVectorRepository, + TextSearchRepository, + VectorCollectionRepository, + VectorRepository, +) + +__all__ = [ + "AttributesMixin", + "DistanceMetric", + "EmbeddingBase", + "EmbeddingMixin", + "HybridModel", + "HybridVectorRepository", + "TextBase", + "TextEmbeddingBase", + "TextMixin", + "TextModel", + "TextSearchRepository", + "UnsupportedDialectError", + "Vector", + "VectorCollection", + "VectorCollectionMixin", + "VectorCollectionRepository", + "VectorDocument", + "VectorModel", + "VectorRepository", + "cosine_distance", + "init_vectors", + "l1_distance", + "l2_distance", + "max_inner_product", + "register_sqlite_vector", +] diff --git a/sqlargon/vectors/loader.py b/sqlargon/vectors/loader.py new file mode 100644 index 0000000..f93734b --- /dev/null +++ b/sqlargon/vectors/loader.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +from inspect import isawaitable +from typing import TYPE_CHECKING, Any +from weakref import WeakSet + +import sqlalchemy as sa +from sqlalchemy import event +from sqlalchemy.util import await_only + +from sqlargon.types.vector import UnsupportedDialectError + +if TYPE_CHECKING: + from sqlalchemy.engine import Engine + from sqlalchemy.ext.asyncio import AsyncEngine + + from sqlargon.database import Database + +_registered_engines: WeakSet[Engine] = WeakSet() + + +async def init_vectors(db: Database) -> None: + """Prepare ``db`` for vector search; call before ``create_all()``. + + On PostgreSQL this creates the ``vector`` extension -- applications + managing their schema with alembic run + ``op.execute("CREATE EXTENSION IF NOT EXISTS vector")`` in a migration + instead. On SQLite it registers the sqlite-vector loadable extension + on every new pool connection. + """ + if db.dialect == "postgresql": + await db.execute(sa.text("CREATE EXTENSION IF NOT EXISTS vector")) + elif db.dialect == "sqlite": + register_sqlite_vector(db.engine) + else: + msg = f"vector search is not supported on the {db.dialect!r} dialect" + raise UnsupportedDialectError(msg) + + +def register_sqlite_vector(engine: AsyncEngine) -> None: + """Load the sqlite-vector extension on every new pool connection. + + Register at startup, before queries run -- connections checked out + earlier never get the extension. Registering the same engine twice + is a no-op. + """ + try: + import importlib.resources + + import sqlite_vector # noqa: F401 + except ImportError as e: + msg = ( + "SQLite vector search requires the 'sqliteai-vector' package; " + "install 'sqlargon[vectors-sqlite]'" + ) + raise ImportError(msg) from e + + if engine.sync_engine in _registered_engines: + return + + # SQLite appends the platform's shared library suffix itself + path = str(importlib.resources.files("sqlite_vector.binaries") / "vector") + + def _resolve(result: Any) -> None: + """Await what an async driver returns, pass a sync one's through. + + An async driver keeps its ``sqlite3`` object in a worker thread of + its own and rejects calls from any other, so loading has to go + through the coroutines it marshals -- awaited here on the greenlet + the ``connect`` event already runs in. + """ + if isawaitable(result): + await_only(result) + + def _load_extension(dbapi_connection: Any, _record: Any) -> None: + connection = getattr(dbapi_connection, "driver_connection", dbapi_connection) + _resolve(connection.enable_load_extension(True)) # noqa: FBT003 + try: + _resolve(connection.load_extension(path)) + finally: + _resolve(connection.enable_load_extension(False)) # noqa: FBT003 + + event.listen(engine.sync_engine, "connect", _load_extension) + _registered_engines.add(engine.sync_engine) diff --git a/sqlargon/vectors/mixins.py b/sqlargon/vectors/mixins.py new file mode 100644 index 0000000..a2f44a1 --- /dev/null +++ b/sqlargon/vectors/mixins.py @@ -0,0 +1,158 @@ +from typing import TYPE_CHECKING, Any, ClassVar +from uuid import UUID + +import sqlalchemy as sa +from sqlalchemy.orm import Mapped, declared_attr, mapped_column + +from sqlargon.types import GUID, JSON +from sqlargon.types.json import json_contains +from sqlargon.types.vector import DistanceMetric, Vector + + +class EmbeddingMixin: + """Adds a fixed-dimension ``embedding`` column to a model. + + Override ``__vector_dim__`` and ``__vector_distance__`` on the + concrete subclass to size the column and pick the metric its index + and default searches use:: + + class Document(EmbeddingMixin, Base): + __vector_dim__ = 384 + __vector_distance__ = DistanceMetric.L2 + """ + + __vector_dim__: ClassVar[int] = 1536 + __vector_distance__: ClassVar[DistanceMetric] = DistanceMetric.COSINE + + @declared_attr + def embedding(cls) -> Mapped[list[float]]: + return mapped_column(Vector(cls.__vector_dim__), nullable=False) + + @classmethod + def embedding_index( + cls, + name: str | None = None, + *, + m: int = 16, + ef_construction: int = 64, + ) -> sa.Index: + """An HNSW index over ``embedding``, tuned for ``__vector_distance__``. + + Add it to the concrete model's ``__table_args__``. The DDL only + runs on PostgreSQL -- sqlite-vector needs no index. With ``name`` + omitted the metadata naming convention applies. + """ + return sa.Index( + name, + "embedding", + postgresql_using="hnsw", + postgresql_with={"m": m, "ef_construction": ef_construction}, + postgresql_ops={"embedding": cls.__vector_distance__.pg_opclass}, + ).ddl_if(dialect="postgresql") + + +class TextMixin: + """Adds a nullable ``text`` column and its full-text search expressions. + + ``__text_regconfig__`` is the PostgreSQL text search configuration + the index and the search ranking share, so overriding it moves both:: + + class Document(TextMixin, Base): + __text_regconfig__ = "english" + """ + + if TYPE_CHECKING: + __tablename__: str + + __text_regconfig__: ClassVar[str] = "simple" + + text: Mapped[str | None] = mapped_column(sa.Text(), nullable=True) + + @classmethod + def _regconfig(cls) -> sa.ColumnElement[Any]: + """The search configuration, inlined rather than bound. + + Index DDL cannot carry bind parameters, and the planner only uses + a functional index when the query spells the expression exactly + as the index does -- so both have to inline it. + """ + regconfig = cls.__text_regconfig__ + if not regconfig.replace("_", "").isalnum(): + msg = f"{regconfig!r} is not a valid text search configuration name" + raise ValueError(msg) + return sa.literal_column(f"'{regconfig}'") + + @classmethod + def text_document(cls) -> sa.ColumnElement[Any]: + """The ``tsvector`` of ``text``, as indexed and as searched.""" + return sa.func.to_tsvector(cls._regconfig(), cls.text) + + @classmethod + def text_query(cls, query: str) -> sa.ColumnElement[Any]: + """``query`` parsed as a web search style ``tsquery``.""" + return sa.func.websearch_to_tsquery(cls._regconfig(), query) + + @classmethod + def text_index(cls, name: str | None = None) -> sa.Index: + """A GIN index over :meth:`text_document`, PostgreSQL only. + + Named explicitly rather than by the metadata convention, which + would derive the name from the inlined search configuration + instead of the column. + """ + return sa.Index( + name or f"ix_{cls.__tablename__}__text", + cls.text_document(), + postgresql_using="gin", + ).ddl_if(dialect="postgresql") + + +class AttributesMixin: + """Adds an ``attributes`` JSON column for application metadata.""" + + # the default is parenthesised because MySQL accepts one on a JSON + # column only as an expression, and every other backend reads + # ``DEFAULT ('{}')`` the same way as a bare literal + attributes: Mapped[dict[str, Any]] = mapped_column( + JSON(), nullable=False, default=dict, server_default=sa.text("('{}')") + ) + + @classmethod + def attributes_contain(cls, value: dict[str, Any]) -> sa.ColumnElement[bool]: + """Whether ``attributes`` contains every given key and value. + + Pass it to any search as an ordinary filter:: + + await docs.search(vector, Document.attributes_contain({"lang": "en"})) + """ + return json_contains(cls.attributes, value) + + @classmethod + def attributes_index(cls, name: str | None = None) -> sa.Index: + """A GIN index over ``attributes``, PostgreSQL only.""" + return sa.Index( + name, + "attributes", + postgresql_using="gin", + postgresql_ops={"attributes": "jsonb_path_ops"}, + ).ddl_if(dialect="postgresql") + + +class VectorCollectionMixin: + """Adds a nullable ``collection_id`` foreign key. + + Points at :class:`~sqlargon.vectors.VectorCollection` by default; + override ``__collection_table__`` to group documents by a table of + your own. + """ + + __collection_table__: ClassVar[str] = "vector_collection" + + @declared_attr + def collection_id(cls) -> Mapped[UUID | None]: + return mapped_column( + GUID(), + sa.ForeignKey(f"{cls.__collection_table__}.id", ondelete="CASCADE"), + nullable=True, + index=True, + ) diff --git a/sqlargon/vectors/models.py b/sqlargon/vectors/models.py new file mode 100644 index 0000000..40af8d1 --- /dev/null +++ b/sqlargon/vectors/models.py @@ -0,0 +1,95 @@ +from typing import TypeVar + +import sqlalchemy as sa +from sqlalchemy.orm import Mapped, mapped_column + +from sqlargon.mixins import CreatedUpdatedMixin, UUIDV7ModelMixin +from sqlargon.orm import Base + +from .mixins import ( + AttributesMixin, + EmbeddingMixin, + TextMixin, + VectorCollectionMixin, +) + + +class EmbeddingBase(EmbeddingMixin, Base): + """Declarative base for models carrying an embedding column. + + The minimum :class:`~sqlargon.vectors.VectorRepository` needs. Add + only the further mixins the application wants:: + + class Document(UUIDV7ModelMixin, EmbeddingBase): + __vector_dim__ = 384 + """ + + __abstract__ = True + + +VectorModel = TypeVar("VectorModel", bound=EmbeddingBase) + + +class TextBase(TextMixin, Base): + """Declarative base for models searchable by full text alone. + + What :class:`~sqlargon.vectors.TextSearchRepository` needs; combine + with :class:`EmbeddingBase` -- or inherit + :class:`TextEmbeddingBase` -- to also search by similarity. + """ + + __abstract__ = True + + +TextModel = TypeVar("TextModel", bound=TextBase) + + +class TextEmbeddingBase(TextMixin, EmbeddingBase): + """Declarative base for models searchable by similarity and by text. + + What :class:`~sqlargon.vectors.HybridVectorRepository` needs, so its + ``rrf_search`` has both rankings to fuse. + """ + + __abstract__ = True + + +HybridModel = TypeVar("HybridModel", bound=TextEmbeddingBase) + + +class VectorCollection(UUIDV7ModelMixin, CreatedUpdatedMixin, Base): + """A named grouping of vector documents. + + The default target of + :class:`~sqlargon.vectors.VectorCollectionMixin`; importing it + registers the table with the shared metadata, so ``create_all()`` + creates it. + """ + + name: Mapped[str] = mapped_column(sa.Unicode(255), unique=True, nullable=False) + + +class VectorDocument( + UUIDV7ModelMixin, + CreatedUpdatedMixin, + AttributesMixin, + VectorCollectionMixin, + TextEmbeddingBase, +): + """Ready-made document model: embedding, text, attributes, collection. + + The batteries-included option -- subclass it, set ``__vector_dim__`` + and add the indexes the columns deserve:: + + class Document(VectorDocument): + __vector_dim__ = 384 + + @declared_attr.directive + def __table_args__(cls) -> tuple[sa.Index, ...]: + return (cls.embedding_index(), cls.attributes_index()) + + Compose the mixins directly instead when only some of the columns + are wanted. + """ + + __abstract__ = True diff --git a/sqlargon/vectors/repository.py b/sqlargon/vectors/repository.py new file mode 100644 index 0000000..ec70454 --- /dev/null +++ b/sqlargon/vectors/repository.py @@ -0,0 +1,288 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, Literal, overload + +from sqlargon.orm import Model +from sqlargon.query_builder import Option, UnsupportedDialectError +from sqlargon.repository import SQLAlchemyRepository + +from .mixins import EmbeddingMixin, TextMixin +from .models import HybridModel, TextModel, VectorCollection, VectorModel + +if TYPE_CHECKING: + from collections.abc import Sequence + + import sqlalchemy as sa + from sqlalchemy.ext.asyncio import AsyncSession + + from sqlargon.types.vector import DistanceMetric + +_VECTOR_INIT_KEY = "sqlargon_vector_init" + + +class _SearchRepository(SQLAlchemyRepository[Model], abstract=True): + """Shared plumbing of the search repositories. + + Each concrete repository declares the mixins its model must carry in + :meth:`_required_mixins`; anything else raises ``TypeError`` on + subclassing, the way + :class:`~sqlargon.repository.SoftDeleteRepository` validates its own. + + The statements themselves are built by the dialect's + :class:`~sqlargon.query_builder.QueryBuilder`, so what is left here is + the orchestration: resolving filters, running the statement and + shaping the rows. + """ + + __slots__ = () + + def __init_subclass__(cls, *, abstract: bool = False, **kwargs: Any) -> None: + super().__init_subclass__(abstract=abstract, **kwargs) + if abstract: + return + for mixin, name in cls._required_mixins(): + if not issubclass(cls.model, mixin): + msg = ( + f"{cls.model.__name__} must inherit from {name} " + f"to be used with {cls.__name__}" + ) + raise TypeError(msg) + + @classmethod + def _required_mixins(cls) -> tuple[tuple[type, str], ...]: + return () + + def _filters( + self, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> list[Any]: + """Positional expressions and keyword equalities, as ``where()`` reads them.""" + filters: list[Any] = list(args) + filters.extend( + getattr(self.model, key) == value for key, value in kwargs.items() + ) + return filters + + def _require(self, option: Option, feature: str) -> None: + """Refuse a feature the dialect's query builder does not claim.""" + if not self.qb.supports(option): + msg = f"{feature} is not supported on the {self.db.dialect!r} dialect" + raise UnsupportedDialectError(msg) + + +class VectorRepository(_SearchRepository[VectorModel], abstract=True): + """Similarity search over a model's embedding column. + + The model type variable is bound to + :class:`~sqlargon.vectors.EmbeddingBase`, so a type checker rejects a + model without an embedding. At runtime + :class:`~sqlargon.vectors.EmbeddingMixin` is enough. + """ + + __slots__ = () + + @classmethod + def _required_mixins(cls) -> tuple[tuple[type, str], ...]: + return ((EmbeddingMixin, "EmbeddingMixin"),) + + def distance( + self, embedding: Sequence[float], metric: DistanceMetric | None = None + ) -> sa.ColumnElement[float]: + """The distance between the model's embedding and ``embedding``. + + PostgreSQL only -- sqlite-vector exposes no scalar distance + function, use :meth:`search` there. + """ + self._require(Option.VECTORS, "a distance expression") + return self.qb.vector_distance(self.model, embedding, metric) + + @overload + async def search( + self, + embedding: Sequence[float], + *filters: Any, + limit: int = ..., + metric: DistanceMetric | None = ..., + with_distance: Literal[False] = ..., + **kwargs: Any, + ) -> Sequence[VectorModel]: ... + + @overload + async def search( + self, + embedding: Sequence[float], + *filters: Any, + limit: int = ..., + metric: DistanceMetric | None = ..., + with_distance: Literal[True], + **kwargs: Any, + ) -> list[tuple[VectorModel, float]]: ... + + async def search( + self, + embedding: Sequence[float], + *filters: Any, + limit: int = 10, + metric: DistanceMetric | None = None, + with_distance: bool = False, + **kwargs: Any, + ) -> Sequence[VectorModel] | list[tuple[VectorModel, float]]: + """The ``limit`` models nearest to ``embedding``, most similar first. + + Positional expressions and keyword equalities narrow the search + the way :meth:`where` does, which is what makes it hybrid:: + + await docs.search( + vector, Document.attributes_contain({"lang": "en"}), limit=5 + ) + + ``metric`` overrides the model's distance metric per query on + PostgreSQL; SQLite fixes the metric per column, so passing a + different one there raises. ``with_distance=True`` returns + ``(model, distance)`` pairs. + """ + self._require(Option.VECTORS, "vector search") + query = self.qb.vector_search( + self.model, + embedding, + *self._filters(filters, kwargs), + limit=limit, + metric=metric, + ) + result = await self._execute_search(query) + if with_distance: + return [(row[0], row[1]) for row in result.all()] + return result.scalars().all() + + async def _execute_search(self, query: sa.Select[Any]) -> sa.Result[Any]: + """Run ``query``, declaring the column first where that is needed. + + The declaration has to reach the same connection as the search, so + both go through one session. + """ + statement = self.qb.vector_init(self.model) + if statement is None: + return await self.execute_query(query) + async with self.session(query) as session: + await self._init_vector_column(session, statement) + return await session.execute(query) + + async def _init_vector_column( + self, session: AsyncSession, statement: sa.Executable + ) -> None: + """Run the declaration once per connection, not once per search.""" + connection = await session.connection() + info = (await connection.get_raw_connection()).info + key = f"{_VECTOR_INIT_KEY}:{self.model.__table__.name}" + if info.get(key): + return + await connection.execute(statement) + info[key] = True + + +class TextSearchRepository(_SearchRepository[TextModel], abstract=True): + """Full-text search over a model's text column, PostgreSQL only.""" + + __slots__ = () + + @classmethod + def _required_mixins(cls) -> tuple[tuple[type, str], ...]: + return ((TextMixin, "TextMixin"),) + + @overload + async def text_search( + self, + query: str, + *filters: Any, + limit: int = ..., + with_score: Literal[False] = ..., + **kwargs: Any, + ) -> Sequence[TextModel]: ... + + @overload + async def text_search( + self, + query: str, + *filters: Any, + limit: int = ..., + with_score: Literal[True], + **kwargs: Any, + ) -> list[tuple[TextModel, float]]: ... + + async def text_search( + self, + query: str, + *filters: Any, + limit: int = 10, + with_score: bool = False, + **kwargs: Any, + ) -> Sequence[TextModel] | list[tuple[TextModel, float]]: + """The ``limit`` best full-text matches of ``query``, best first. + + Positional expressions and keyword equalities narrow the search + the way :meth:`where` does. PostgreSQL only. + """ + self._require(Option.FULL_TEXT, "text_search") + statement = self.qb.text_search( + self.model, query, *self._filters(filters, kwargs), limit=limit + ) + result = await self.execute_query(statement) + if with_score: + return [(row[0], row[1]) for row in result.all()] + return result.scalars().all() + + +class HybridVectorRepository( + VectorRepository[HybridModel], TextSearchRepository[HybridModel], abstract=True +): + """Similarity search, full-text search, and their fusion. + + Requires a model carrying both an embedding and a text column, so + :meth:`rrf_search` has two rankings to fuse. + """ + + __slots__ = () + + @classmethod + def _required_mixins(cls) -> tuple[tuple[type, str], ...]: + return ( + (EmbeddingMixin, "EmbeddingMixin"), + (TextMixin, "TextMixin"), + ) + + async def rrf_search( + self, + embedding: Sequence[float], + query: str, + *filters: Any, + k: int = 60, + limit: int = 10, + candidates: int = 50, + **kwargs: Any, + ) -> list[tuple[HybridModel, float]]: + """Hybrid search fusing the two rankings with reciprocal rank fusion. + + Ranks the ``candidates`` nearest rows and the ``candidates`` best + full-text matches of ``query``, then scores each row + ``sum(1 / (k + rank))`` over the rankings it appears in, so a row + both agree on outranks one either alone prefers. PostgreSQL only. + """ + self._require(Option.VECTORS | Option.FULL_TEXT, "rrf_search") + statement = self.qb.rrf_search( + self.model, + embedding, + query, + *self._filters(filters, kwargs), + k=k, + limit=limit, + candidates=candidates, + ) + result = await self.execute_query(statement) + return [(row[0], row[1]) for row in result.all()] + + +class VectorCollectionRepository(SQLAlchemyRepository[VectorCollection]): + """Repository over the ready-made :class:`VectorCollection` model.""" + + __slots__ = () diff --git a/tests/conftest.py b/tests/conftest.py index 177d27d..512a5d8 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -46,6 +46,12 @@ def pytest_collection_modifyitems( item.add_marker(skip_e2e) +@pytest.fixture(scope="session") +def anyio_backend() -> str: + """Run the suite on asyncio only, the one loop SQLAlchemy supports.""" + return "asyncio" + + @pytest.fixture(scope="session", autouse=True) def db(): return Database.from_env() diff --git a/tests/e2e/backends.py b/tests/e2e/backends.py index 337bcc7..89e2524 100644 --- a/tests/e2e/backends.py +++ b/tests/e2e/backends.py @@ -18,8 +18,10 @@ from contextlib import AbstractContextManager from pathlib import Path -# ``uuidv7()`` is a PostgreSQL 18 builtin, so GenerateUUIDV7 needs at least it -POSTGRES_IMAGE = os.environ.get("SQLARGON_E2E_POSTGRES_IMAGE", "postgres:18-alpine") +# ``uuidv7()`` is a PostgreSQL 18 builtin, so GenerateUUIDV7 needs at least it. +# The pgvector image is that PostgreSQL plus the extension the vector suite +# needs, so it stands in for the plain one rather than adding a backend. +POSTGRES_IMAGE = os.environ.get("SQLARGON_E2E_POSTGRES_IMAGE", "pgvector/pgvector:pg18") MYSQL_IMAGE = os.environ.get("SQLARGON_E2E_MYSQL_IMAGE", "mysql:8.4") # RANDOM_BYTES, which the UUID server defaults use, needs MariaDB 10.10 MARIADB_IMAGE = os.environ.get("SQLARGON_E2E_MARIADB_IMAGE", "mariadb:11.4") @@ -99,6 +101,9 @@ class Backend: #: an upsert may leave a column of the conflict set out of its values partial_upsert: bool = True is_mariadb: bool = False + #: the server can search vectors -- pgvector, or the sqlite-vector + #: loadable extension + vector_search: bool = False @property def is_mysql_family(self) -> bool: @@ -119,6 +124,7 @@ def is_mysql_family(self) -> bool: skip_locked=False, # the SQLite key operators match JSON values, not object keys json_key_operators=False, + vector_search=True, ), "postgres": Backend( name="postgres", @@ -132,6 +138,7 @@ def is_mysql_family(self) -> bool: server_side_uuid=True, skip_locked=True, json_key_operators=True, + vector_search=True, ), "mysql": Backend( name="mysql", diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index b292190..e00d874 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -3,27 +3,38 @@ Every test runs once per selected backend. A container is started once per session and shared, while the :class:`~sqlargon.Database` is per test: an engine's pool holds connections bound to the event loop that opened them, and -pytest-asyncio gives each test a fresh loop. +the anyio plugin gives each test a fresh loop. """ from __future__ import annotations -import asyncio +from functools import partial from typing import TYPE_CHECKING +import anyio import pytest +import sqlalchemy as sa from sqlalchemy.ext.asyncio import create_async_engine from sqlargon import Base, Database +from sqlargon.vectors import init_vectors from .backends import Backend, parse_backends from .models import ( SERVER_DEFAULT_TABLES, TABLES, XMIN_TABLES, + AuditArticleRepository, + AuditCommentRepository, + AuditFollowRepository, DocumentRepository, + OutboxUserRepository, + RawAuditArticleRepository, SoftUserRepository, UserRepository, + UUIDAuditArticleRepository, + VectorDocRepository, + VectorNoteRepository, VersionedUserRepository, XminUserRepository, ) @@ -31,8 +42,6 @@ if TYPE_CHECKING: from collections.abc import AsyncGenerator, Generator - import sqlalchemy as sa - def pytest_generate_tests(metafunc: pytest.Metafunc) -> None: """Run every e2e test once per selected backend.""" @@ -41,7 +50,10 @@ def pytest_generate_tests(metafunc: pytest.Metafunc) -> None: backends = parse_backends(metafunc.config.getoption("e2e_backends")) metafunc.parametrize( "backend", - backends, + [ + pytest.param(backend, marks=pytest.mark.xdist_group(backend.name)) + for backend in backends + ], ids=[backend.name for backend in backends], indirect=True, scope="session", @@ -76,17 +88,26 @@ def tables(backend: Backend) -> tuple[sa.Table, ...]: @pytest.fixture(scope="session") -def schema(database_url: str, tables: tuple[sa.Table, ...]) -> Generator[None]: +def schema( + backend: Backend, database_url: str, tables: tuple[sa.Table, ...] +) -> Generator[None]: """Create the e2e tables once per backend, and drop them afterwards. DDL runs in its own throwaway engine and event loop, so the schema can be session scoped without an engine outliving the loop that built it. + + The ``vector`` extension comes first: a ``VECTOR`` column cannot be + declared before the type it names exists. """ async def run(*, create: bool) -> None: engine = create_async_engine(database_url) try: async with engine.begin() as connection: + if create and backend.dialect == "postgresql": + await connection.execute( + sa.text("CREATE EXTENSION IF NOT EXISTS vector") + ) await connection.run_sync( Base.metadata.create_all if create else Base.metadata.drop_all, tables=list(tables), @@ -94,9 +115,9 @@ async def run(*, create: bool) -> None: finally: await engine.dispose() - asyncio.run(run(create=True)) + anyio.run(partial(run, create=True)) yield - asyncio.run(run(create=False)) + anyio.run(partial(run, create=False)) @pytest.fixture(autouse=True) @@ -181,3 +202,70 @@ def versioned_users() -> VersionedUserRepository: @pytest.fixture def xmin_users() -> XminUserRepository: return XminUserRepository() + + +@pytest.fixture +def outbox_users() -> OutboxUserRepository: + return OutboxUserRepository() + + +@pytest.fixture +def audit_articles() -> AuditArticleRepository: + return AuditArticleRepository() + + +@pytest.fixture +def raw_audit_articles() -> RawAuditArticleRepository: + return RawAuditArticleRepository() + + +@pytest.fixture +def uuid_audit_articles() -> UUIDAuditArticleRepository: + return UUIDAuditArticleRepository() + + +@pytest.fixture +def audit_comments() -> AuditCommentRepository: + return AuditCommentRepository() + + +@pytest.fixture +def audit_follows() -> AuditFollowRepository: + return AuditFollowRepository() + + +@pytest.fixture +def needs_foreign_keys(backend: Backend) -> None: + if backend.name == "sqlite": + pytest.skip("sqlite does not enforce foreign keys unless asked to") + + +@pytest.fixture +def needs_vector_search(backend: Backend) -> None: + if not backend.vector_search: + pytest.skip(f"{backend.name} cannot search vectors") + if backend.dialect == "sqlite": + pytest.importorskip( + "sqlite_vector", reason="sqlite vector search needs 'sqliteai-vector'" + ) + + +@pytest.fixture +async def vector_db(db: Database) -> Database: + """The database under test, prepared for vector search. + + On PostgreSQL this creates the extension; on SQLite it registers the + loadable one on the pool, which every later connection then gets. + """ + await init_vectors(db) + return db + + +@pytest.fixture +def vector_notes(vector_db: Database) -> VectorNoteRepository: + return VectorNoteRepository().using(db=vector_db) + + +@pytest.fixture +def vector_docs(vector_db: Database) -> VectorDocRepository: + return VectorDocRepository().using(db=vector_db) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index a10b49d..6d3e42e 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -11,21 +11,35 @@ import sqlalchemy as sa from pydantic import BaseModel -from sqlalchemy.orm import Mapped, mapped_column +from sqlalchemy.orm import Mapped, declared_attr, mapped_column, relationship from sqlargon import ( + AuditableBase, + AuditableRepository, Base, SoftDeleteBase, SoftDeleteRepository, SQLAlchemyRepository, + UUIDAuditableBase, VersionedBase, VersionedRepository, XminVersionedBase, + latest_relationship, + version_foreign_key, + version_mapped_column, ) from sqlargon.mixins import CreatedUpdatedMixin, UUIDModelMixin +from sqlargon.outbox import OutboxConfig, OutboxEvent, OutboxRepository from sqlargon.types import GUID, JSON, GenerateUUID, GenerateUUIDV7, Timestamp, now from sqlargon.types.pydantic import Pydantic from sqlargon.typing import OnConflictOptions +from sqlargon.vectors import ( + EmbeddingBase, + HybridVectorRepository, + VectorCollection, + VectorDocument, + VectorRepository, +) class Address(BaseModel): @@ -134,16 +148,152 @@ class VersionedUserRepository(VersionedRepository[VersionedUser]): default_order_by = VersionedUser.name -class XminUserRepository(VersionedRepository[XminUser]): # type: ignore[type-var] +class XminUserRepository(VersionedRepository[XminUser]): default_order_by = XminUser.name +class AuditArticle(UUIDModelMixin, AuditableBase): + """Append-only, versioned by a counter.""" + + __tablename__ = "e2e_audit_article" + + name: Mapped[str] = mapped_column(sa.Unicode(64)) + tag: Mapped[str | None] = mapped_column(sa.Unicode(64), nullable=True) + + +class UUIDAuditArticle(UUIDModelMixin, UUIDAuditableBase): + """The same shape, versioned by UUIDv7.""" + + __tablename__ = "e2e_uuid_audit_article" + + name: Mapped[str] = mapped_column(sa.Unicode(64)) + + +class AuditComment(UUIDModelMixin, Base): + """A child pinned to one exact version of an article.""" + + __tablename__ = "e2e_audit_comment" + + article_id: Mapped[UUID] = mapped_column(GUID()) + article_version = version_mapped_column(AuditArticle) + body: Mapped[str] = mapped_column(sa.Unicode(64)) + + __table_args__ = ( + version_foreign_key(AuditArticle, "article_id", "article_version"), + ) + + article: Mapped[AuditArticle] = relationship() + + +class AuditFollow(UUIDModelMixin, Base): + """A child following whichever version of an article is newest.""" + + __tablename__ = "e2e_audit_follow" + + article_id: Mapped[UUID] = mapped_column(GUID()) + + article: Mapped[AuditArticle] = latest_relationship(AuditArticle, "article_id") + + +class AuditArticleRepository(AuditableRepository[AuditArticle]): + pass + + +class RawAuditArticleRepository(SQLAlchemyRepository[AuditArticle]): + """Unscoped view of the same table, to observe what is physically stored.""" + + default_order_by = AuditArticle.version + + +class UUIDAuditArticleRepository(AuditableRepository[UUIDAuditArticle]): + pass + + +class AuditCommentRepository(SQLAlchemyRepository[AuditComment]): + pass + + +class AuditFollowRepository(SQLAlchemyRepository[AuditFollow]): + pass + + +class OutboxUser(UUIDModelMixin, CreatedUpdatedMixin, Base): + __tablename__ = "e2e_outbox_user" + + name: Mapped[str] = mapped_column(sa.Unicode(64), unique=True) + password: Mapped[str | None] = mapped_column(sa.Unicode(64), nullable=True) + tenant_id: Mapped[UUID | None] = mapped_column(GUID(), nullable=True) + + +class OutboxUserRepository(OutboxRepository[OutboxUser]): + default_order_by = OutboxUser.name + + outbox = OutboxConfig( + topic="users", + type_prefix="user", + source="e2e", + exclude={"password", "tenant_id"}, + attributes={"tenant_id": "tenant_id"}, + ) + + +class VectorNote(UUIDModelMixin, EmbeddingBase): + """An embedding and nothing else, to prove the mixins are optional.""" + + __tablename__ = "e2e_vector_note" + __vector_dim__ = 3 + + @declared_attr.directive + def __table_args__(cls) -> tuple[sa.Index, ...]: + return (cls.embedding_index(),) + + name: Mapped[str] = mapped_column(sa.Unicode(64)) + + +class VectorDoc(VectorDocument): + """Every column the extension offers, indexes included.""" + + __tablename__ = "e2e_vector_doc" + __vector_dim__ = 3 + __text_regconfig__ = "english" + + @declared_attr.directive + def __table_args__(cls) -> tuple[sa.Index, ...]: + return (cls.embedding_index(), cls.attributes_index(), cls.text_index()) + + +class VectorNoteRepository(VectorRepository[VectorNote]): + default_order_by = VectorNote.name + + +class VectorDocRepository(HybridVectorRepository[VectorDoc]): + default_order_by = VectorDoc.created_at + + def _tables(*models: type[Base]) -> tuple[sa.Table, ...]: return tuple(Base.metadata.tables[model.__tablename__] for model in models) #: Tables every backend can hold; the only ones the e2e suite creates. -TABLES: tuple[sa.Table, ...] = _tables(User, Document, SoftUser, VersionedUser) +#: +#: A child referencing an audited version comes before the table it points at, +#: so the per test cleanup can empty them in this order without tripping the +#: foreign key. +TABLES: tuple[sa.Table, ...] = _tables( + User, + Document, + SoftUser, + VersionedUser, + OutboxUser, + OutboxEvent, + AuditComment, + AuditFollow, + AuditArticle, + UUIDAuditArticle, + VectorNote, + VectorDoc, + VectorCollection, +) #: Tables whose DDL carries a server side UUID default, which not every #: backend accepts -- see :attr:`~tests.e2e.backends.Backend.server_side_uuid`. diff --git a/tests/e2e/test_auditable.py b/tests/e2e/test_auditable.py new file mode 100644 index 0000000..c16ecb3 --- /dev/null +++ b/tests/e2e/test_auditable.py @@ -0,0 +1,330 @@ +"""The append-only repository against a real server. + +What only a real backend can prove: the correlated subquery scoping every +read, the ``INSERT ... SELECT`` an append compiles to, the composite foreign +key pinning a child to one version, and -- on MySQL, which has no RETURNING +at all -- the path that has to find the appended rows again. +""" + +from __future__ import annotations + +from uuid import uuid4 + +import pytest +import sqlalchemy as sa +from sqlalchemy.exc import IntegrityError + +from sqlargon import AppendOnlyError, ConcurrentModificationError + +from .models import AuditArticle, AuditComment, AuditFollow, UUIDAuditArticle + + +@pytest.fixture +async def article(audit_articles): + """One article carried up to version 3.""" + created = await audit_articles.create(name="draft", tag="news") + await audit_articles.update_one({"name": "revised"}, AuditArticle.id == created.id) + await audit_articles.update_one({"name": "final"}, AuditArticle.id == created.id) + return created.id + + +# --- appending --- + + +async def test_update_appends_a_row_instead_of_rewriting_one( + article, raw_audit_articles +): + stored = await raw_audit_articles.list(AuditArticle.id == article) + + assert [(row.version, row.name) for row in stored] == [ + (1, "draft"), + (2, "revised"), + (3, "final"), + ] + # a column the caller never named rides along + assert {row.tag for row in stored} == {"news"} + + +async def test_reads_are_scoped_to_the_newest_version(article, audit_articles): + head = await audit_articles.get(id=article) + + assert (head.version, head.name) == (3, "final") + assert await audit_articles.count() == 1 + assert await audit_articles.versions().count() == 3 + + +async def test_update_many_appends_one_version_per_entity( + audit_articles, raw_audit_articles +): + first = await audit_articles.create(name="a") + second = await audit_articles.create(name="b") + + appended = await audit_articles.update_many( + {"tag": "shared"}, AuditArticle.id.in_([first.id, second.id]) + ) + + assert {row.version for row in appended} == {2} + assert {row.tag for row in appended} == {"shared"} + assert len(await raw_audit_articles.list()) == 4 + + +async def test_the_fluent_builder_appends_in_one_statement( + article, audit_articles, raw_audit_articles +): + await ( + audit_articles.update({"name": "fluent"}) + .filter(AuditArticle.id == article) + .execute() + ) + + assert (await audit_articles.get(id=article)).name == "fluent" + assert len(await raw_audit_articles.list(AuditArticle.id == article)) == 4 + + +async def test_create_or_update_appends_or_creates(article, audit_articles): + appended = await audit_articles.create_or_update(id=article, name="fourth") + created = await audit_articles.create_or_update(id=uuid4(), name="fresh") + + assert (appended.version, appended.name) == (4, "fourth") + assert created.version == 1 + + +async def test_bulk_create_or_update_appends_and_creates(article, audit_articles): + fresh = uuid4() + + await audit_articles.bulk_create_or_update( + [{"id": article, "name": "appended"}, {"id": fresh, "name": "created"}] + ) + + assert (await audit_articles.get(id=article)).version == 4 + assert (await audit_articles.get(id=fresh)).version == 1 + + +async def test_bulk_update_appends_one_version_per_row(audit_articles): + first = await audit_articles.create(name="a") + second = await audit_articles.create(name="b") + + await audit_articles.bulk_update( + [{"id": first.id, "name": "a2"}, {"id": second.id, "name": "b2"}] + ) + + assert (await audit_articles.get(id=first.id)).name == "a2" + assert (await audit_articles.get(id=second.id)).name == "b2" + + +async def test_upsert_is_refused(audit_articles): + with pytest.raises(AppendOnlyError, match="cannot resolve a conflict"): + audit_articles.upsert([{"id": uuid4(), "name": "x"}]) + + +# --- deletion is an appended tombstone --- + + +async def test_remove_appends_a_tombstoned_version( + article, audit_articles, raw_audit_articles +): + await audit_articles.remove(AuditArticle.id == article) + + stored = await raw_audit_articles.list(AuditArticle.id == article) + assert [(row.version, row.tombstone) for row in stored[-1:]] == [(4, True)] + assert await audit_articles.get(id=article) is None + assert await audit_articles.versions().count() == 4 + + +async def test_restore_appends_a_live_version(article, audit_articles): + await audit_articles.remove(AuditArticle.id == article) + + restored = await audit_articles.restore(AuditArticle.id == article) + + assert [(row.version, row.tombstone) for row in restored] == [(5, False)] + assert (await audit_articles.get(id=article)).name == "final" + + +# --- inspecting the history --- + + +async def test_history_returns_every_version_oldest_first(article, audit_articles): + history = await audit_articles.history(id=article) + + assert [(row.version, row.name) for row in history] == [ + (1, "draft"), + (2, "revised"), + (3, "final"), + ] + + +async def test_get_version_returns_one_exact_version(article, audit_articles): + assert (await audit_articles.get_version(2, id=article)).name == "revised" + assert await audit_articles.get_version(9, id=article) is None + + +async def test_at_reads_the_state_of_that_moment( + article, audit_articles, raw_audit_articles +): + second = (await raw_audit_articles.list(AuditArticle.id == article))[1] + + as_of = await audit_articles.at(second.created_at).get(id=article) + + assert (as_of.version, as_of.name) == (2, "revised") + + +async def test_at_hides_an_entity_already_deleted_by_then( + article, audit_articles, raw_audit_articles +): + await audit_articles.remove(AuditArticle.id == article) + tombstone = (await raw_audit_articles.list(AuditArticle.id == article))[-1] + + assert await audit_articles.at(tombstone.created_at).get(id=article) is None + + +# --- optimistic concurrency --- + + +async def test_update_if_match_appends_when_the_version_is_current( + article, audit_articles +): + appended = await audit_articles.update_if_match( + {"name": "guarded"}, AuditArticle.id == article, expected_version=3 + ) + + assert (appended.version, appended.name) == (4, "guarded") + + +async def test_update_if_match_refuses_a_stale_version(article, audit_articles): + with pytest.raises(ConcurrentModificationError): + await audit_articles.update_if_match( + {"name": "stale"}, + AuditArticle.id == article, + expected_version=1, + raise_on_mismatch=True, + ) + + +async def test_a_duplicate_version_collides_on_the_primary_key( + article, raw_audit_articles +): + with pytest.raises(IntegrityError): + await raw_audit_articles.insert( + [{"id": article, "version": 3, "name": "racing"}] + ).execute() + + +# --- relationships --- + + +async def test_pinned_relationship_stays_on_its_version( + article, audit_articles, audit_comments +): + await audit_comments.create(article_id=article, article_version=2, body="on v2") + await audit_articles.update_one({"name": "later"}, AuditArticle.id == article) + + comment = await audit_comments.load(AuditComment.article).one() + + assert (comment.article.version, comment.article.name) == (2, "revised") + + +async def test_latest_relationship_follows_the_entity_forward( + article, audit_articles, audit_follows +): + await audit_follows.create(article_id=article) + + before = await audit_follows.load(AuditFollow.article).one() + assert (before.article.version, before.article.name) == (3, "final") + + await audit_articles.update_one({"name": "newest"}, AuditArticle.id == article) + after = await audit_follows.load(AuditFollow.article).one() + + assert (after.article.version, after.article.name) == (4, "newest") + + +async def test_latest_relationship_joins_in_sql(article, audit_follows): + await audit_follows.create(article_id=article) + + names = ( + await audit_follows.select(AuditArticle.name).join(AuditFollow.article).all() + ) + + assert names == ["final"] + + +# --- purging superseded versions --- + + +async def test_purge_keeps_the_newest_version_only( + article, audit_articles, raw_audit_articles +): + await audit_articles.purge(id=article) + + stored = await raw_audit_articles.list(AuditArticle.id == article) + assert [(row.version, row.name) for row in stored] == [(3, "final")] + + +async def test_purge_leaves_other_entities_alone( + article, audit_articles, raw_audit_articles +): + other = await audit_articles.create(name="other") + await audit_articles.update_one({"name": "other2"}, AuditArticle.id == other.id) + + await audit_articles.purge(id=article) + + assert len(await raw_audit_articles.list(AuditArticle.id == other.id)) == 2 + + +@pytest.mark.usefixtures("needs_foreign_keys") +async def test_a_pinned_version_cannot_be_purged( + article, audit_articles, audit_comments +): + await audit_comments.create(article_id=article, article_version=1, body="pinned") + + with pytest.raises(IntegrityError): + await audit_articles.purge(id=article) + + +# --- the UUIDv7 strategy --- + + +async def test_uuid_strategy_appends_sortable_versions(uuid_audit_articles): + entity_id = uuid4() + + await uuid_audit_articles.create(id=entity_id, name="draft") + await uuid_audit_articles.update_one( + {"name": "final"}, UUIDAuditArticle.id == entity_id + ) + + history = await uuid_audit_articles.history(id=entity_id) + versions = [row.version for row in history] + + assert [row.name for row in history] == ["draft", "final"] + assert versions == sorted(versions) + assert (await uuid_audit_articles.get(id=entity_id)).name == "final" + + +async def test_uuid_strategy_deletes_by_appending_a_tombstone(uuid_audit_articles): + entity_id = uuid4() + await uuid_audit_articles.create(id=entity_id, name="draft") + + await uuid_audit_articles.remove(UUIDAuditArticle.id == entity_id) + + assert await uuid_audit_articles.get(id=entity_id) is None + assert len(await uuid_audit_articles.history(id=entity_id)) == 2 + + +# --- the scope reaches the server, not just the ORM --- + + +async def test_the_latest_scope_is_one_sql_statement(article, audit_articles): + statement = audit_articles.select().filter(AuditArticle.id == article).query + compiled = str(statement.compile()) + + assert "ORDER BY" in compiled + assert compiled.count("e2e_audit_article") >= 2 + + +async def test_count_of_a_history_spanning_two_entities(audit_articles): + first = await audit_articles.create(name="a") + await audit_articles.create(name="b") + await audit_articles.update_one({"name": "a2"}, AuditArticle.id == first.id) + + assert await audit_articles.count() == 2 + assert await audit_articles.versions().count() == 3 + assert await audit_articles.count(sa.true()) == 2 diff --git a/tests/e2e/test_cron.py b/tests/e2e/test_cron.py index 07a4545..f7e8bfe 100644 --- a/tests/e2e/test_cron.py +++ b/tests/e2e/test_cron.py @@ -78,8 +78,8 @@ async def changed() -> None: ... before = {task.name: task.next_run_at for task in await cron.tasks()} rescheduling = Cron(namespace=NAMESPACE) - rescheduling.task(EVERY_MINUTE, name="kept")(kept) - rescheduling.task("*/5 * * * *", name="changed")(changed) + rescheduling.task(EVERY_MINUTE, "kept", kept) + rescheduling.task("*/5 * * * *", "changed", changed) await rescheduling.sync() after = {task.name: task for task in await rescheduling.tasks()} diff --git a/tests/e2e/test_outbox.py b/tests/e2e/test_outbox.py new file mode 100644 index 0000000..9677de4 --- /dev/null +++ b/tests/e2e/test_outbox.py @@ -0,0 +1,221 @@ +"""The transactional outbox against a real server. + +Capture runs on every backend: the repository reads its written rows back +through the same RETURNING fallback the rest of the library uses, so MySQL and +MariaDB -- which have no RETURNING at all -- take the select-write-refetch path +instead. See :mod:`tests.e2e.test_capabilities`. +""" + +from datetime import timedelta +from typing import TYPE_CHECKING +from uuid import uuid4 + +import anyio +import pytest + +from sqlargon import Database +from sqlargon.outbox import OutboxEvent, OutboxEventRepository, OutboxRelay +from sqlargon.utils import utc_now + +from .models import OutboxUser, OutboxUserRepository + +if TYPE_CHECKING: + from collections.abc import Sequence + +HASHED = "not-a-real-secret" + + +class Failure(RuntimeError): + pass + + +async def create_then_fail(db: Database, users: OutboxUserRepository) -> None: + async with db.session_context(): + await users.create(name="John") + raise Failure + + +@pytest.fixture +def events() -> OutboxEventRepository: + return OutboxEventRepository() + + +# --- capture --- + + +async def test_create_records_an_event( + outbox_users: OutboxUserRepository, events: OutboxEventRepository +): + user = await outbox_users.create(name="John", password=HASHED) + assert user is not None + + (event,) = await events.select().all() + assert event.topic == "users" + assert event.type == "user.created" + assert event.source == "e2e" + assert event.data["name"] == "John" + assert event.data["id"] == str(user.id) + assert "password" not in event.data + assert event.published_at is None + + +async def test_update_and_delete_are_recorded( + outbox_users: OutboxUserRepository, events: OutboxEventRepository +): + user = await outbox_users.create(name="John") + assert user is not None + await outbox_users.update_one({"name": "Jane"}, id=user.id) + await outbox_users.remove(OutboxUser.id == user.id) + + types = [ + event.type + for event in await events.select().order_by(OutboxEvent.created_at).all() + ] + assert types == ["user.created", "user.updated", "user.deleted"] + + +async def test_an_upsert_records_what_each_row_turned_out_to_be( + outbox_users: OutboxUserRepository, events: OutboxEventRepository +): + """``is_new`` tells the inserted rows from the updated ones on every backend.""" + existing = await outbox_users.create(name="John") + assert existing is not None + + # every column of the default conflict set is named, so the statement is + # portable -- see test_capabilities.test_upsert_of_a_partial_row + await outbox_users.bulk_create_or_update( + [ + {"id": existing.id, "name": "Jane", "password": HASHED, "tenant_id": None}, + {"id": uuid4(), "name": "Jack", "password": HASHED, "tenant_id": None}, + ] + ) + + recorded = { + event.data["name"]: event.type + for event in await events.select().order_by(OutboxEvent.created_at).all() + } + assert recorded == { + "John": "user.created", + "Jane": "user.updated", + "Jack": "user.created", + } + + +async def test_an_upsert_leaves_the_creation_time_alone( + outbox_users: OutboxUserRepository, +): + user = await outbox_users.create(name="John") + assert user is not None + + updated = await outbox_users.create_or_update( + id=user.id, name="Jane", password=HASHED, tenant_id=None + ) + + assert updated.created_at == user.created_at + assert updated.updated_at > updated.created_at + + +async def test_bulk_writes_record_one_event_per_row( + outbox_users: OutboxUserRepository, events: OutboxEventRepository +): + await outbox_users.bulk_create([{"name": "a"}, {"name": "b"}]) + + created = await events.select().all() + assert sorted(event.data["name"] for event in created) == ["a", "b"] + + +async def test_the_event_rolls_back_with_the_row( + db: Database, outbox_users: OutboxUserRepository, events: OutboxEventRepository +): + with pytest.raises(Failure): + await create_then_fail(db, outbox_users) + + assert await outbox_users.count() == 0 + assert await events.count() == 0 + + +async def test_the_payload_survives_a_round_trip( + outbox_users: OutboxUserRepository, events: OutboxEventRepository +): + """The JSON column and the engine serializer agree on the row's types.""" + user = await outbox_users.create(name="John") + assert user is not None + + (event,) = await events.select().all() + assert event.data["id"] == str(user.id) + assert event.data["created_at"].startswith(str(user.created_at.year)) + + +async def test_the_attributes_survive_a_round_trip( + outbox_users: OutboxUserRepository, events: OutboxEventRepository +): + tenant = uuid4() + user = await outbox_users.create(name="John", tenant_id=tenant) + assert user is not None + + (event,) = await events.select().all() + assert event.attributes == {"tenant_id": str(tenant)} + assert "tenant_id" not in event.data + + +# --- relay --- + + +async def test_the_relay_publishes_and_marks( + outbox_users: OutboxUserRepository, events: OutboxEventRepository +): + published: list[OutboxEvent] = [] + + async def publisher(event: OutboxEvent) -> None: + published.append(event) + + relay = OutboxRelay(publisher) + await outbox_users.create(name="John") + + assert await relay.dispatch_once() == 1 + assert [event.data["name"] for event in published] == ["John"] + assert await events.pending_count() == 0 + + +async def test_the_relay_purges_published_events( + outbox_users: OutboxUserRepository, events: OutboxEventRepository +): + async def publisher(_event: OutboxEvent) -> None: ... + + relay = OutboxRelay(publisher, retention=timedelta(days=7)) + await outbox_users.create(name="John") + (event,) = await events.select().all() + await events.mark_published([event.id], utc_now() - timedelta(days=8)) + + assert await relay.purge() == 1 + assert await events.count() == 0 + + +@pytest.mark.usefixtures("needs_skip_locked") +async def test_concurrent_relays_never_claim_the_same_event( + database_url: str, outbox_users: OutboxUserRepository +): + await outbox_users.create(name="John") + claimed: list[Sequence[OutboxEvent]] = [] + + async def claim(database: Database) -> None: + repository = OutboxEventRepository().using(db=database) + claimed.append( + await repository.claim_pending(utc_now(), lease=timedelta(minutes=5)) + ) + + first, second = Database(database_url), Database(database_url) + try: + async with anyio.create_task_group() as group: + group.start_soon(claim, first) + group.start_soon(claim, second) + finally: + await first.dispose() + await second.dispose() + + assert sorted(len(batch) for batch in claimed) == [0, 1] + + +# The background poll loop is dialect independent and covered by the unit +# suite; cancelling it here would abort whichever statement the driver has in +# flight and hand the next test a poisoned connection. diff --git a/tests/e2e/test_types.py b/tests/e2e/test_types.py index 4743e8d..73748ec 100644 --- a/tests/e2e/test_types.py +++ b/tests/e2e/test_types.py @@ -6,6 +6,19 @@ from sqlalchemy.exc import StatementError from sqlargon import Database +from sqlargon.types.json import ( + json_array_append, + json_array_length, + json_get, + json_has_key, + json_insert_key, + json_keys, + json_remove_key, + json_replace_key, + json_set_key, + json_update, +) +from sqlargon.utils import json_loads from .backends import Backend from .models import Address, Document, DocumentRepository, ServerDefaults, User @@ -137,3 +150,146 @@ async def test_json_has_any_key_rejects_unknown_keys(documents: DocumentReposito documents.select().where(Document.payload.has_any_key(["country"])).all() ) assert matched == [] + + +async def mutate(documents: DocumentRepository, expression: object) -> Document: + """Apply ``expression`` to every document and read the row back.""" + await documents.update({Document.payload: expression}).execute() + stored = await documents.first() + assert stored is not None + return stored + + +async def test_json_set_key_adds_a_key(documents: DocumentRepository): + await store(documents, payload={"a": 1}) + stored = await mutate(documents, json_set_key(Document.payload, "b", 2)) + assert stored.payload == {"a": 1, "b": 2} + + +async def test_json_set_key_overwrites_a_key(documents: DocumentRepository): + await store(documents, payload={"a": 1}) + stored = await mutate(documents, json_set_key(Document.payload, "a", 9)) + assert stored.payload == {"a": 9} + + +async def test_json_set_key_stores_a_nested_document(documents: DocumentRepository): + await store(documents, payload={"a": 1}) + stored = await mutate(documents, json_set_key(Document.payload, "b", {"x": [1, 2]})) + # the value is a document, not the serialized text as a JSON string + assert stored.payload == {"a": 1, "b": {"x": [1, 2]}} + + +async def test_json_update_merges_shallowly(documents: DocumentRepository): + await store(documents, payload={"a": {"x": 1}, "b": 2}) + stored = await mutate( + documents, json_update(Document.payload, {"a": {"y": 9}, "c": 3}) + ) + # "a" is replaced wholesale rather than merged into + assert stored.payload == {"a": {"y": 9}, "b": 2, "c": 3} + + +async def test_json_update_with_no_keys_leaves_the_document( + documents: DocumentRepository, +): + await store(documents, payload={"a": 1}) + stored = await mutate(documents, json_update(Document.payload, {})) + assert stored.payload == {"a": 1} + + +async def test_json_remove_key_drops_keys(documents: DocumentRepository): + await store(documents, payload={"a": 1, "b": 2, "c": 3}) + stored = await mutate(documents, json_remove_key(Document.payload, "a", "c")) + assert stored.payload == {"b": 2} + + +async def test_json_remove_key_ignores_a_missing_key(documents: DocumentRepository): + await store(documents, payload={"a": 1}) + stored = await mutate(documents, json_remove_key(Document.payload, "nope")) + assert stored.payload == {"a": 1} + + +async def test_json_insert_key_only_adds_a_missing_key(documents: DocumentRepository): + await store(documents, payload={"a": 1}) + stored = await mutate(documents, json_insert_key(Document.payload, "b", 2)) + assert stored.payload == {"a": 1, "b": 2} + + +async def test_json_insert_key_leaves_an_existing_key(documents: DocumentRepository): + await store(documents, payload={"a": 1}) + stored = await mutate(documents, json_insert_key(Document.payload, "a", 9)) + assert stored.payload == {"a": 1} + + +async def test_json_replace_key_only_updates_an_existing_key( + documents: DocumentRepository, +): + await store(documents, payload={"a": 1}) + stored = await mutate(documents, json_replace_key(Document.payload, "a", 9)) + assert stored.payload == {"a": 9} + + +async def test_json_replace_key_does_not_add_a_missing_key( + documents: DocumentRepository, +): + await store(documents, payload={"a": 1}) + stored = await mutate(documents, json_replace_key(Document.payload, "b", 2)) + assert stored.payload == {"a": 1} + + +async def test_json_mutations_compose_in_one_statement(documents: DocumentRepository): + await store(documents, payload={"a": 1, "b": 2}) + stored = await mutate( + documents, json_remove_key(json_update(Document.payload, {"c": 3}), "a") + ) + assert stored.payload == {"b": 2, "c": 3} + + +async def test_json_array_append_appends_one_element(documents: DocumentRepository): + await store(documents, tags=["red"]) + await documents.update( + {Document.tags: json_array_append(Document.tags, "green")} + ).execute() + stored = await documents.first() + assert stored is not None + assert stored.tags == ["red", "green"] + + +async def test_json_array_append_nests_a_list(documents: DocumentRepository): + await store(documents, tags=["red"]) + await documents.update( + {Document.tags: json_array_append(Document.tags, ["a", "b"])} + ).execute() + stored = await documents.first() + assert stored is not None + # appended as one element rather than concatenated + assert stored.tags == ["red", ["a", "b"]] + + +async def test_json_has_key_addresses_object_keys(documents: DocumentRepository): + # the gap has_any_key / has_all_keys leave on sqlite, whose json_each + # fallback matches values instead of keys + await store(documents, payload={"a": "b"}) + assert await documents.select().where(json_has_key(Document.payload, "a")).all() + assert not await documents.select().where(json_has_key(Document.payload, "b")).all() + + +async def test_json_get_reads_a_nested_document(documents: DocumentRepository): + await store(documents, payload={"a": {"x": 1}}) + value = await documents.select(json_get(Document.payload, "a")).scalar() + # sqlite and mysql hand back the JSON text, postgres a decoded document + if isinstance(value, str): + value = json_loads(value) + assert value == {"x": 1} + + +async def test_json_array_length_counts_elements(documents: DocumentRepository): + await store(documents, tags=["a", "b", "c"]) + assert await documents.select(json_array_length(Document.tags)).scalar() == 3 + + +async def test_json_keys_lists_top_level_keys(documents: DocumentRepository): + await store(documents, payload={"a": 1, "b": 2}) + keys = await documents.select(json_keys(Document.payload)).scalar() + if isinstance(keys, str): + keys = json_loads(keys) + assert sorted(keys) == ["a", "b"] diff --git a/tests/e2e/test_vectors.py b/tests/e2e/test_vectors.py new file mode 100644 index 0000000..4cce22f --- /dev/null +++ b/tests/e2e/test_vectors.py @@ -0,0 +1,253 @@ +"""Vector search against a real server. + +Runs on the backends that can search vectors: PostgreSQL through pgvector, +SQLite through the sqlite-vector loadable extension. The two take entirely +different paths -- an ORDER BY over a distance operator against a join over a +table valued scan -- so every similarity test is worth running on both. + +Full text search and reciprocal rank fusion are PostgreSQL only. +""" + +from typing import TYPE_CHECKING + +import pytest +import sqlalchemy as sa + +from sqlargon import Database +from sqlargon.vectors import ( + DistanceMetric, + UnsupportedDialectError, + VectorCollectionRepository, +) + +from .backends import Backend +from .models import VectorDoc, VectorDocRepository, VectorNoteRepository + +if TYPE_CHECKING: + from collections.abc import Sequence + +pytestmark = pytest.mark.usefixtures("needs_vector_search") + +# unit vectors along the three axes, so cosine distances are exactly known +X = [1.0, 0.0, 0.0] +Y = [0.0, 1.0, 0.0] +Z = [0.0, 0.0, 1.0] + +DOCUMENTS = [ + {"text": "red apple fruit", "embedding": X, "attributes": {"lang": "en"}}, + {"text": "blue sky above", "embedding": Y, "attributes": {"lang": "en"}}, + {"text": "green apple tree", "embedding": Z, "attributes": {"lang": "de"}}, +] + + +@pytest.fixture +async def documents(vector_docs: VectorDocRepository) -> VectorDocRepository: + await vector_docs.create_many(DOCUMENTS) + return vector_docs + + +def _texts(models: "Sequence[VectorDoc]") -> list[str]: + return [model.text for model in models] + + +async def test_embedding_round_trips_as_floats(vector_docs: VectorDocRepository): + await vector_docs.create(text="one", embedding=[0.25, 0.5, 0.75]) + stored = await vector_docs.one() + assert stored.embedding == pytest.approx([0.25, 0.5, 0.75]) + + +@pytest.mark.parametrize( + ("query", "expected"), + [(X, "red apple fruit"), (Y, "blue sky above"), (Z, "green apple tree")], +) +async def test_search_returns_the_nearest_first( + documents: VectorDocRepository, query, expected +): + found = await documents.search(query, limit=1) + assert _texts(found) == [expected] + + +async def test_search_orders_every_row_by_distance(documents: VectorDocRepository): + found = await documents.search([1.0, 0.1, 0.0], limit=3) + assert _texts(found)[0] == "red apple fruit" + assert _texts(found)[1] == "blue sky above" + + +async def test_search_honours_the_limit(documents: VectorDocRepository): + assert len(await documents.search(X, limit=2)) == 2 + + +async def test_search_reports_distances(documents: VectorDocRepository): + found = await documents.search(X, limit=2, with_distance=True) + nearest, distance = found[0] + assert nearest.text == "red apple fruit" + # cosine distance of a vector to itself + assert distance == pytest.approx(0.0, abs=1e-5) + assert found[1][1] > distance + + +async def test_hybrid_search_filters_out_the_nearest_neighbour( + documents: VectorDocRepository, +): + """The filter has to run before the limit, not after it. + + ``X`` is the embedding of the English row, so a scan that took its top + hit first and filtered afterwards would come back empty. + """ + found = await documents.search( + X, VectorDoc.attributes_contain({"lang": "de"}), limit=2 + ) + assert _texts(found) == ["green apple tree"] + + +async def test_hybrid_search_accepts_arbitrary_expressions( + documents: VectorDocRepository, +): + found = await documents.search(X, VectorDoc.text.like("%tree%"), limit=2) + assert _texts(found) == ["green apple tree"] + + +async def test_hybrid_search_accepts_keyword_equalities( + documents: VectorDocRepository, +): + found = await documents.search(X, collection_id=None, limit=3) + assert len(found) == 3 + + +async def test_search_works_without_the_optional_columns( + vector_notes: VectorNoteRepository, +): + """A model carrying nothing but an embedding still searches.""" + await vector_notes.create_many( + [ + {"name": "x-axis", "embedding": X}, + {"name": "z-axis", "embedding": Z}, + ] + ) + found = await vector_notes.search(Z, limit=1) + assert [note.name for note in found] == ["z-axis"] + + +async def test_search_by_collection( + documents: VectorDocRepository, vector_db: Database +): + collections = VectorCollectionRepository().using(db=vector_db) + collection = await collections.create(name="library") + assert collection is not None + await documents.update_one({"collection_id": collection.id}, text="blue sky above") + found = await documents.search(X, collection_id=collection.id, limit=3) + assert _texts(found) == ["blue sky above"] + + +# --- PostgreSQL only --- + + +@pytest.fixture +def needs_postgresql(backend: Backend) -> None: + if backend.dialect != "postgresql": + pytest.skip(f"{backend.name} has no PostgreSQL full text search") + + +@pytest.mark.usefixtures("needs_postgresql") +async def test_text_search_ranks_matches(documents: VectorDocRepository): + found = await documents.text_search("apple", limit=3) + assert sorted(_texts(found)) == ["green apple tree", "red apple fruit"] + + +@pytest.mark.usefixtures("needs_postgresql") +async def test_text_search_reports_scores(documents: VectorDocRepository): + found = await documents.text_search("apple", limit=3, with_score=True) + assert all(score > 0 for _, score in found) + + +@pytest.mark.usefixtures("needs_postgresql") +async def test_text_search_ignores_non_matches(documents: VectorDocRepository): + assert await documents.text_search("submarine", limit=3) == [] + + +@pytest.mark.usefixtures("needs_postgresql") +async def test_rrf_search_prefers_what_both_rankings_agree_on( + documents: VectorDocRepository, +): + """The row both rankings find outranks the rows only one of them does. + + ``Z`` is the embedding of "green apple tree" and it is the only row + matching "tree", so it takes the top of both rankings while the other + two appear in the vector ranking alone. Searching for "apple" instead + would leave the top two tied, since ``ts_rank`` scores both rows + matching it the same. + """ + found = await documents.rrf_search(Z, "tree", limit=3) + assert found[0][0].text == "green apple tree" + assert found[0][1] > found[1][1] + + +@pytest.mark.usefixtures("needs_postgresql") +async def test_rrf_search_keeps_rows_only_one_ranking_found( + documents: VectorDocRepository, +): + found = await documents.rrf_search(Y, "apple", limit=3) + assert set(_texts([model for model, _ in found])) == { + "red apple fruit", + "blue sky above", + "green apple tree", + } + + +@pytest.mark.usefixtures("needs_postgresql") +async def test_metric_override_changes_the_ordering( + documents: VectorDocRepository, +): + found = await documents.search(X, limit=1, metric=DistanceMetric.L2) + assert _texts(found) == ["red apple fruit"] + + +@pytest.mark.usefixtures("needs_postgresql") +async def test_the_indexes_exist(vector_db: Database): + result = await vector_db.execute( + sa.text("SELECT indexdef FROM pg_indexes WHERE tablename = 'e2e_vector_doc'") + ) + definitions = " ".join(row[0] for row in result) + assert "USING hnsw" in definitions + assert "vector_cosine_ops" in definitions + assert "USING gin" in definitions + assert "to_tsvector" in definitions + + +# --- SQLite only --- + + +@pytest.fixture +def needs_sqlite(backend: Backend) -> None: + if backend.dialect != "sqlite": + pytest.skip(f"{backend.name} is not sqlite") + + +@pytest.mark.usefixtures("needs_sqlite") +async def test_search_survives_a_fresh_pooled_connection( + documents: VectorDocRepository, vector_db: Database +): + """The extension is loaded per connection, so a new one must get it too.""" + await vector_db.dispose() + found = await documents.search(X, limit=1) + assert _texts(found) == ["red apple fruit"] + + +@pytest.mark.usefixtures("needs_sqlite") +async def test_repeated_searches_reuse_the_initialised_column( + documents: VectorDocRepository, +): + assert _texts(await documents.search(X, limit=1)) == ["red apple fruit"] + assert _texts(await documents.search(Z, limit=1)) == ["green apple tree"] + + +@pytest.mark.usefixtures("needs_sqlite") +async def test_text_search_is_rejected(documents: VectorDocRepository): + with pytest.raises(UnsupportedDialectError, match="text_search"): + await documents.text_search("apple") + + +@pytest.mark.usefixtures("needs_sqlite") +async def test_metric_override_is_rejected(documents: VectorDocRepository): + with pytest.raises(UnsupportedDialectError, match="per column"): + await documents.search(X, metric=DistanceMetric.L2) diff --git a/tests/test_auditable.py b/tests/test_auditable.py new file mode 100644 index 0000000..005055e --- /dev/null +++ b/tests/test_auditable.py @@ -0,0 +1,846 @@ +from uuid import UUID, uuid4 + +import pytest +import sqlalchemy as sa +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Mapped, mapped_column, relationship + +from sqlargon import ( + AppendOnlyError, + AuditableBase, + AuditableRepository, + Base, + ConcurrentModificationError, + Database, + SQLAlchemyRepository, + UUIDAuditableBase, + latest_relationship, + version_foreign_key, + version_mapped_column, +) +from sqlargon.dialects.sqlite import SQLiteQueryBuilder +from sqlargon.mixins import AuditableMixin, IntegerAuditableMixin, UUIDModelMixin +from sqlargon.query_builder import Option +from sqlargon.types import GUID +from tests import MEMORY_URL + +# Models defined at module level to avoid re-registration with --count=3 + + +class Article(AuditableBase): + __tablename__ = "test_auditable_article" + id = sa.Column(sa.Integer, primary_key=True) + name = sa.Column(sa.Unicode(255), nullable=True) + tag = sa.Column(sa.Unicode(255), nullable=True) + + +class UUIDArticle(UUIDModelMixin, UUIDAuditableBase): + """The same shape, versioned by UUIDv7 rather than by a counter.""" + + __tablename__ = "test_auditable_uuid_article" + name: Mapped[str | None] = mapped_column(sa.Unicode(255), nullable=True) + + +class MixedIn(IntegerAuditableMixin, Base): + """The mixin combined with ``Base`` by hand, rather than AuditableBase.""" + + __tablename__ = "test_auditable_mixed_in" + id = sa.Column(sa.Integer, primary_key=True) + + +class VersionOnly(AuditableBase): + """Keyed by its version alone, so no entity can be told from another.""" + + __tablename__ = "test_auditable_version_only" + + +class Plain(Base): + __tablename__ = "test_auditable_plain" + id = sa.Column(sa.Integer, primary_key=True) + + +class Comment(Base): + """A child pinned to one exact version of an article.""" + + __tablename__ = "test_auditable_comment" + id = sa.Column(sa.Integer, primary_key=True) + article_id = sa.Column(sa.Integer) + article_version = version_mapped_column(Article) + body = sa.Column(sa.Unicode(255), nullable=True) + + __table_args__ = (version_foreign_key(Article, "article_id", "article_version"),) + + article: Mapped[Article] = relationship() + + +class Follow(Base): + """A child following whichever version of an article is newest.""" + + __tablename__ = "test_auditable_follow" + id = sa.Column(sa.Integer, primary_key=True) + article_id = sa.Column(sa.Integer) + + article: Mapped[Article] = latest_relationship(Article, "article_id") + + +class ArticleRepository(AuditableRepository[Article]): + pass + + +class RawArticleRepository(SQLAlchemyRepository[Article]): + """Unscoped view of the same table, to observe what is physically stored.""" + + default_order_by = Article.version + + +class UUIDArticleRepository(AuditableRepository[UUIDArticle]): + pass + + +class CommentRepository(SQLAlchemyRepository[Comment]): + pass + + +class FollowRepository(SQLAlchemyRepository[Follow]): + pass + + +TABLES = ( + Article.__table__, + UUIDArticle.__table__, + Comment.__table__, + Follow.__table__, +) + + +@pytest.fixture(autouse=True) +async def tables(db: Database): + async with db.engine.begin() as conn: + for table in TABLES: + await conn.run_sync(table.create, checkfirst=True) + yield + async with db.engine.begin() as conn: + for table in reversed(TABLES): + await conn.run_sync(table.drop, checkfirst=True) + + +@pytest.fixture +def repository(): + return ArticleRepository() + + +@pytest.fixture +def raw(): + return RawArticleRepository() + + +@pytest.fixture +async def articles(repository): + """A repository holding one article carried up to version 3.""" + await repository.create(id=1, name="draft", tag="news") + await repository.update_one({"name": "revised"}, Article.id == 1) + await repository.update_one({"name": "final"}, Article.id == 1) + return repository + + +# --- model validation --- + + +def test_model_without_auditable_mixin_is_rejected(): + with pytest.raises(TypeError, match="Plain must inherit from AuditableMixin"): + + class BadRepository(AuditableRepository[Plain]): # type: ignore[type-var] + pass + + +def test_model_keyed_by_version_alone_is_rejected(): + with pytest.raises(TypeError, match="VersionOnly is keyed by its version alone"): + + class BadRepository(AuditableRepository[VersionOnly]): + pass + + +def test_model_declaring_the_mixin_by_hand_is_accepted(): + # the static bound asks for AnyAuditableBase, but the version column and + # its expressions are all the repository actually needs + class MixedInRepository(AuditableRepository[MixedIn]): # type: ignore[type-var] + pass + + assert MixedInRepository.model is MixedIn + + +def test_abstract_subclass_needs_no_model(): + class Shared(AuditableRepository[Article], abstract=True): + pass + + class Concrete(Shared): + pass + + assert Concrete.model is Article + + +def test_model_can_be_passed_explicitly(): + class Explicit(AuditableRepository, model=Article): + pass + + assert Explicit.model is Article + + +def test_version_joins_the_primary_key(): + assert {c.name for c in Article.__table__.primary_key.columns} == {"id", "version"} + assert Article.audit_key() == ("id",) + + +def test_version_and_tombstone_are_left_out_of_the_conflict_set(): + default_set = ArticleRepository._get_default_set() + assert "version" not in default_set + assert "tombstone" not in default_set + + +# --- appending instead of updating --- + + +async def test_create_starts_at_version_one(repository, raw): + created = await repository.create(id=1, name="draft") + + assert created.version == 1 + assert len(await raw.list()) == 1 + + +@pytest.mark.usefixtures("articles") +async def test_update_appends_a_row_and_leaves_the_old_one_alone(raw): + stored = await raw.list() + + assert [(row.id, row.version, row.name) for row in stored] == [ + (1, 1, "draft"), + (1, 2, "revised"), + (1, 3, "final"), + ] + + +@pytest.mark.usefixtures("articles") +async def test_append_carries_columns_the_caller_did_not_name(raw): + assert [row.tag for row in await raw.list()] == ["news", "news", "news"] + + +@pytest.mark.usefixtures("articles") +async def test_each_version_is_timestamped_on_its_own(raw): + stored = await raw.list() + timestamps = [row.created_at for row in stored] + + assert timestamps == sorted(timestamps) + assert len(set(timestamps)) == len(timestamps) + # nothing is ever updated, so the two timestamps never diverge + assert all(row.created_at == row.updated_at for row in stored) + + +async def test_update_many_appends_one_version_per_entity(repository, raw): + await repository.bulk_create([{"id": 1, "name": "a"}, {"id": 2, "name": "b"}]) + + appended = await repository.update_many({"tag": "shared"}, Article.id.in_([1, 2])) + + assert {row.version for row in appended} == {2} + assert len(await raw.list()) == 4 + + +async def test_create_or_update_appends_the_next_version(articles): + appended = await articles.create_or_update(id=1, name="fifth") + + assert appended.version == 4 + assert appended.name == "fifth" + + +async def test_create_or_update_creates_an_absent_entity(repository): + created = await repository.create_or_update(id=7, name="fresh") + + assert (created.id, created.version) == (7, 1) + + +async def test_create_or_update_revives_a_deleted_entity(repository): + await repository.create(id=1, name="draft") + await repository.remove(Article.id == 1) + + revived = await repository.create_or_update(id=1, name="back") + + assert (revived.version, revived.tombstone) == (3, False) + assert await repository.get(id=1) is not None + + +async def test_bulk_update_appends_one_version_per_row(repository, raw): + await repository.bulk_create([{"id": 1, "name": "a"}, {"id": 2, "name": "b"}]) + + await repository.bulk_update([{"id": 1, "name": "a2"}, {"id": 2, "name": "b2"}]) + + assert [(row.id, row.version, row.name) for row in await raw.list()] == [ + (1, 1, "a"), + (2, 1, "b"), + (1, 2, "a2"), + (2, 2, "b2"), + ] + + +# --- reads see the newest version --- + + +@pytest.mark.parametrize("method", ["all", "list"]) +async def test_reads_return_only_the_newest_version(articles, method): + rows = await getattr(articles.select() if method == "all" else articles, method)() + + assert [(row.version, row.name) for row in rows] == [(3, "final")] + + +async def test_get_returns_the_newest_version(articles): + assert (await articles.get(id=1)).name == "final" + + +async def test_count_counts_entities_while_versions_counts_rows(articles): + assert await articles.count() == 1 + assert await articles.versions().count() == 3 + + +async def test_selecting_columns_is_scoped_too(articles): + assert await articles.select(Article.name).all() == ["final"] + + +async def test_reads_of_two_entities_pick_each_ones_newest(repository): + await repository.bulk_create([{"id": 1, "name": "a"}, {"id": 2, "name": "b"}]) + await repository.update_one({"name": "a2"}, Article.id == 1) + + rows = await repository.list() + + assert sorted((row.id, row.version, row.name) for row in rows) == [ + (1, 2, "a2"), + (2, 1, "b"), + ] + + +# --- delete appends a tombstone --- + + +@pytest.mark.parametrize("method", ["remove", "delete_one", "delete_many"]) +async def test_delete_appends_a_tombstoned_version(articles, raw, method): + await getattr(articles, method)(Article.id == 1) + + stored = await raw.list() + assert [(row.version, row.tombstone) for row in stored[-1:]] == [(4, True)] + assert len(stored) == 4 + + +async def test_a_deleted_entity_leaves_reads_but_keeps_its_history(articles): + await articles.remove(Article.id == 1) + + assert await articles.list() == [] + assert await articles.count() == 0 + assert await articles.versions().count() == 4 + + +async def test_with_deleted_sees_the_tombstoned_head(articles): + await articles.remove(Article.id == 1) + + head = await articles.with_deleted().get(id=1) + + assert (head.version, head.tombstone) == (4, True) + + +async def test_only_deleted_scopes_to_entities_whose_head_is_a_tombstone(articles): + await articles.bulk_create([{"id": 2, "name": "live"}]) + await articles.remove(Article.id == 1) + + assert [row.id for row in await articles.only_deleted().list()] == [1] + + +async def test_restore_appends_a_live_version(articles, raw): + await articles.remove(Article.id == 1) + + restored = await articles.restore(Article.id == 1) + + assert [(row.version, row.tombstone) for row in restored] == [(5, False)] + assert (await articles.get(id=1)).name == "final" + assert len(await raw.list()) == 5 + + +# --- inspecting the history --- + + +async def test_history_returns_every_version_oldest_first(articles): + assert [(row.version, row.name) for row in await articles.history(id=1)] == [ + (1, "draft"), + (2, "revised"), + (3, "final"), + ] + + +async def test_history_covers_tombstoned_versions(articles): + await articles.remove(Article.id == 1) + + assert [row.version for row in await articles.history(id=1)] == [1, 2, 3, 4] + + +async def test_get_version_returns_one_exact_version(articles): + assert (await articles.get_version(2, id=1)).name == "revised" + assert await articles.get_version(9, id=1) is None + + +async def test_versions_is_an_unscoped_view(articles): + rows = await articles.versions().list() + + assert sorted(row.version for row in rows) == [1, 2, 3] + + +async def test_at_reads_the_state_of_that_moment(articles, raw): + second = (await raw.list())[1] + + as_of = await articles.at(second.created_at).get(id=1) + + assert (as_of.version, as_of.name) == (2, "revised") + + +async def test_at_hides_an_entity_already_deleted_by_then(articles, raw): + await articles.remove(Article.id == 1) + tombstone = (await raw.list())[-1] + + assert await articles.at(tombstone.created_at).get(id=1) is None + + +async def test_at_predates_an_entity_that_did_not_exist_yet(articles, raw): + first = (await raw.list())[0] + await articles.create(id=2, name="later") + + assert [row.id for row in await articles.at(first.created_at).list()] == [1] + + +# --- the builders append rather than rewrite --- + + +async def test_update_builder_appends_in_one_statement(articles, raw): + await articles.update({"name": "fluent"}).filter(Article.id == 1).execute() + + assert [(row.version, row.name) for row in await raw.list()][-1:] == [(4, "fluent")] + + +async def test_update_builder_is_awaitable_directly(articles): + await articles.update({"name": "awaited"}).filter(Article.id == 1) + + assert (await articles.get(id=1)).name == "awaited" + + +async def test_update_builder_only_touches_matched_entities(repository, raw): + await repository.bulk_create([{"id": 1, "name": "a"}, {"id": 2, "name": "b"}]) + + await repository.update({"name": "only one"}).filter(Article.id == 1).execute() + + assert sorted((row.id, row.version) for row in await raw.list()) == [ + (1, 1), + (1, 2), + (2, 1), + ] + + +async def test_update_builder_can_return_the_appended_rows(articles): + appended = ( + await articles.update({"name": "returned"}, return_results=True) + .filter(Article.id == 1) + .all() + ) + + assert [(row.version, row.name) for row in appended] == [(4, "returned")] + + +async def test_update_builder_reads_the_scope_it_was_built_from(articles, raw): + await articles.remove(Article.id == 1) + + # the entity is tombstoned, so the live scope matches nothing to append to + await articles.update({"name": "ignored"}).filter(Article.id == 1).execute() + + assert len(await raw.list()) == 4 + + +async def test_delete_builder_appends_a_tombstone(articles, raw): + await articles.delete().filter(Article.id == 1).execute() + + assert [(row.version, row.tombstone) for row in await raw.list()][-1:] == [ + (4, True) + ] + + +async def test_update_builder_accepts_a_sql_expression(articles): + appended = await articles.update_one({"name": Article.tag}, Article.id == 1) + + assert appended.name == "news" + + +async def test_update_builder_rejects_many_value_sets(repository): + with pytest.raises(AppendOnlyError, match="takes a single mapping"): + repository.update([{"name": "a"}, {"name": "b"}]) + + +def test_upsert_is_refused(repository): + with pytest.raises(AppendOnlyError, match="cannot resolve a conflict"): + repository.upsert([{"id": 1}]) + + +# --- bulk create or update --- + + +async def test_bulk_create_or_update_appends_and_creates(articles, raw): + await articles.bulk_create_or_update( + [{"id": 1, "name": "appended"}, {"id": 9, "name": "created"}] + ) + + assert sorted((row.id, row.version, row.name) for row in await raw.list()) == [ + (1, 1, "draft"), + (1, 2, "revised"), + (1, 3, "final"), + (1, 4, "appended"), + (9, 1, "created"), + ] + + +async def test_bulk_create_or_update_can_return_the_written_rows(articles): + written = await articles.bulk_create_or_update( + [{"id": 1, "name": "appended"}, {"id": 9, "name": "created"}], + return_results=True, + ) + + assert sorted((row.id, row.version) for row in written) == [(1, 4), (9, 1)] + + +async def test_bulk_create_or_update_revives_a_deleted_entity(repository): + await repository.create(id=1, name="draft") + await repository.remove(Article.id == 1) + + await repository.bulk_create_or_update([{"id": 1, "name": "back"}]) + + assert (await repository.get(id=1)).name == "back" + + +async def test_bulk_create_or_update_needs_the_entity_key(repository): + with pytest.raises(AppendOnlyError, match="entity key"): + await repository.bulk_create_or_update([{"name": "keyless"}]) + + +async def test_bulk_create_or_update_of_nothing_is_a_no_op(repository, raw): + await repository.bulk_create_or_update([]) + + assert await raw.list() == [] + + +# --- optimistic concurrency --- + + +async def test_update_if_match_appends_when_the_version_is_current(articles): + appended = await articles.update_if_match( + {"name": "guarded"}, Article.id == 1, expected_version=3 + ) + + assert (appended.version, appended.name) == (4, "guarded") + + +async def test_update_if_match_refuses_a_stale_version(articles, raw): + assert ( + await articles.update_if_match( + {"name": "stale"}, Article.id == 1, expected_version=2 + ) + is None + ) + assert len(await raw.list()) == 3 + + +async def test_update_if_match_can_raise_on_a_stale_version(articles): + with pytest.raises(ConcurrentModificationError, match="version 2"): + await articles.update_if_match( + {"name": "stale"}, + Article.id == 1, + expected_version=2, + raise_on_mismatch=True, + ) + + +async def test_delete_if_match_appends_a_tombstone_when_current(articles): + deleted = await articles.delete_if_match(Article.id == 1, expected_version=3) + + assert (deleted.version, deleted.tombstone) == (4, True) + + +async def test_delete_if_match_refuses_a_stale_version(articles): + assert await articles.delete_if_match(Article.id == 1, expected_version=1) is None + + +@pytest.mark.usefixtures("articles") +async def test_a_duplicate_version_collides_on_the_primary_key(raw): + with pytest.raises(IntegrityError): + await raw.insert([{"id": 1, "version": 3, "name": "racing"}]).execute() + + +# --- purging superseded versions --- + + +async def test_purge_keeps_the_newest_version_only(articles, raw): + await articles.purge(id=1) + + assert [(row.version, row.name) for row in await raw.list()] == [(3, "final")] + assert (await articles.get(id=1)).name == "final" + + +async def test_purge_leaves_other_entities_alone(articles, raw): + await articles.create(id=2, name="other") + await articles.update_one({"name": "other2"}, Article.id == 2) + + await articles.purge(id=1) + + assert sorted((row.id, row.version) for row in await raw.list()) == [ + (1, 3), + (2, 1), + (2, 2), + ] + + +async def test_purge_keeps_a_tombstoned_head(articles, raw): + await articles.remove(Article.id == 1) + + await articles.purge(id=1) + + assert [(row.version, row.tombstone) for row in await raw.list()] == [(4, True)] + + +async def test_purge_on_a_single_version_entity_is_a_no_op(repository, raw): + await repository.create(id=1, name="only") + + await repository.purge(id=1) + + assert len(await raw.list()) == 1 + + +# --- relationships --- + + +async def test_pinned_relationship_stays_on_its_version(articles): + comments = CommentRepository() + await comments.create(id=1, article_id=1, article_version=2, body="on the draft") + + await articles.update_one({"name": "even later"}, Article.id == 1) + comment = await comments.load(Comment.article).one() + + assert (comment.article.version, comment.article.name) == (2, "revised") + + +async def test_latest_relationship_follows_the_entity_forward(articles): + follows = FollowRepository() + await follows.create(id=1, article_id=1) + + before = await follows.load(Follow.article).one() + assert (before.article.version, before.article.name) == (3, "final") + + await articles.update_one({"name": "newest"}, Article.id == 1) + after = await FollowRepository().load(Follow.article).one() + + assert (after.article.version, after.article.name) == (4, "newest") + + +@pytest.mark.usefixtures("articles") +async def test_latest_relationship_joins_without_eager_loading(): + follows = FollowRepository() + await follows.create(id=1, article_id=1) + + rows = await follows.select(Article.name).join(Follow.article).all() + + assert rows == ["final"] + + +# --- scope escapes keep the routing preference --- + + +async def test_scope_escapes_retain_the_routing_preference(repository): + other = Database(MEMORY_URL) + try: + bound = repository.using(db=other) + + assert bound.versions().db is other + assert bound.at(sa.func.now()).db is other + assert bound.with_deleted().db is other + finally: + await other.dispose() + + +# --- the UUIDv7 strategy behaves the same --- + + +async def test_uuid_strategy_appends_a_sortable_version(): + repository = UUIDArticleRepository() + entity_id = uuid4() + + await repository.create(id=entity_id, name="draft") + await repository.update_one({"name": "final"}, UUIDArticle.id == entity_id) + + history = await repository.history(id=entity_id) + versions = [row.version for row in history] + + assert [row.name for row in history] == ["draft", "final"] + assert all(isinstance(version, UUID) for version in versions) + assert versions == sorted(versions) + + +async def test_uuid_strategy_reads_the_newest_version(): + repository = UUIDArticleRepository() + entity_id = uuid4() + + await repository.create(id=entity_id, name="draft") + await repository.update_one({"name": "final"}, UUIDArticle.id == entity_id) + + assert (await repository.get(id=entity_id)).name == "final" + assert await repository.count() == 1 + assert await repository.versions().count() == 2 + + +async def test_uuid_strategy_appends_client_side_in_bulk(): + """The bulk paths mint the successor in Python, not in SQL.""" + repository = UUIDArticleRepository() + first, second = uuid4(), uuid4() + await repository.create(id=first, name="a") + await repository.create(id=second, name="b") + + await repository.bulk_update( + [{"id": first, "name": "a2"}, {"id": second, "name": "b2"}] + ) + + assert (await repository.get(id=first)).name == "a2" + assert (await repository.get(id=second)).name == "b2" + versions = [row.version for row in await repository.history(id=first)] + assert versions == sorted(versions) + + +async def test_uuid_strategy_bulk_create_or_update_appends_and_creates(): + repository = UUIDArticleRepository() + known, fresh = uuid4(), uuid4() + await repository.create(id=known, name="a") + + written = await repository.bulk_create_or_update( + [{"id": known, "name": "a2"}, {"id": fresh, "name": "new"}], + return_results=True, + ) + + assert len(written) == 2 + assert (await repository.get(id=known)).name == "a2" + assert (await repository.get(id=fresh)).name == "new" + assert await repository.versions().count() == 3 + + +async def test_uuid_strategy_deletes_by_appending_a_tombstone(): + repository = UUIDArticleRepository() + entity_id = uuid4() + await repository.create(id=entity_id, name="draft") + + await repository.remove(UUIDArticle.id == entity_id) + + assert await repository.list() == [] + assert len(await repository.history(id=entity_id)) == 2 + + +# --- a dialect without RETURNING has to find the appended rows again --- + + +class NoReturningQueryBuilder(SQLiteQueryBuilder): + supported_options = Option.CONFLICTS + + +@pytest.fixture +def no_returning(db: Database): + """Strip the RETURNING clause off the dialect, as MySQL has none.""" + builder = db.query_builder + db.query_builder = NoReturningQueryBuilder() + yield db + db.query_builder = builder + + +@pytest.mark.usefixtures("no_returning") +async def test_the_fallback_tests_see_a_dialect_without_returning(repository): + """Guards every test below: without this they exercise the native path.""" + assert not repository.qb.supports(Option.RETURNING) + + +@pytest.mark.usefixtures("no_returning") +async def test_update_one_refetches_the_appended_row(articles, raw): + appended = await articles.update_one({"name": "refetched"}, Article.id == 1) + + assert (appended.version, appended.name) == (4, "refetched") + assert len(await raw.list()) == 4 + + +@pytest.mark.usefixtures("no_returning") +async def test_delete_one_refetches_the_tombstoned_row(articles): + deleted = await articles.delete_one(Article.id == 1) + + assert (deleted.version, deleted.tombstone) == (4, True) + assert await articles.get(id=1) is None + + +@pytest.mark.usefixtures("no_returning") +async def test_update_many_refetches_every_appended_row(repository): + await repository.bulk_create([{"id": 1, "name": "a"}, {"id": 2, "name": "b"}]) + + appended = await repository.update_many({"tag": "t"}, Article.id.in_([1, 2])) + + assert sorted((row.id, row.version, row.tag) for row in appended) == [ + (1, 2, "t"), + (2, 2, "t"), + ] + + +@pytest.mark.usefixtures("no_returning") +async def test_an_append_matching_nothing_returns_nothing(repository): + assert await repository.update_one({"name": "x"}, Article.id == 404) is None + + +@pytest.mark.usefixtures("no_returning") +async def test_update_if_match_still_guards_without_returning(articles): + assert ( + await articles.update_if_match( + {"name": "stale"}, Article.id == 1, expected_version=1 + ) + is None + ) + + +# --- streaming --- + + +async def test_stream_yields_the_newest_versions(articles): + rows = [row async for row in articles.select().stream()] + + assert [row[0].name for row in rows] == ["final"] + + +async def test_stream_runs_a_staged_append(articles, raw): + staged = articles.update({"name": "streamed"}, return_results=True).filter( + Article.id == 1 + ) + + rows = [row async for row in staged.stream()] + + assert [row[0].name for row in rows] == ["streamed"] + assert len(await raw.list()) == 4 + + +async def test_bulk_update_of_nothing_is_a_no_op(repository, raw): + await repository.bulk_update([]) + + assert await raw.list() == [] + + +# --- relationship helpers --- + + +@pytest.mark.parametrize("helper", [latest_relationship, version_foreign_key]) +def test_relationship_helpers_check_the_column_count(helper): + with pytest.raises(ValueError, match="identified by \\('id',\\)"): + helper(Article, "article_id", "extra", "surplus") + + +def test_version_mapped_column_takes_the_type_of_the_parent(): + assert isinstance(Comment.__table__.c.article_version.type, sa.Integer) + assert isinstance(UUIDArticle.__table__.c.version.type, GUID) + + +def test_the_marker_mixin_declares_no_version_strategy(): + with pytest.raises(NotImplementedError): + AuditableMixin.next_version_expression() diff --git a/tests/test_cron.py b/tests/test_cron.py index 8e6d657..a6fb201 100644 --- a/tests/test_cron.py +++ b/tests/test_cron.py @@ -54,6 +54,27 @@ async def other() -> None: ... assert tasks["cleanup"].next_run_at > utc_now() +async def test_task_declares_a_function_given_outright(cron): + async def purge() -> None: ... + + assert cron.task("0 3 * * *", "purge_outbox", purge) is purge + + await cron.sync() + (task,) = await cron.tasks() + assert task.name == "purge_outbox" + assert task.schedule == "0 3 * * *" + assert task.declarative is True + + +async def test_a_function_given_outright_keeps_its_own_name(cron): + async def purge() -> None: ... + + cron.task(EVERY_MINUTE, func=purge) + + await cron.sync() + assert [task.name for task in await cron.tasks()] == ["purge"] + + async def test_sync_updates_and_deletes(cron): @cron.task("*/5 * * * *") async def cleanup() -> None: ... @@ -264,6 +285,26 @@ async def job() -> None: await ran.wait() +async def test_run_executes_a_bound_method_declared_outright(cron): + ran = anyio.Event() + + class Relay: + async def purge(self) -> None: + ran.set() + + cron.task(EVERY_MINUTE, "purge_outbox", Relay().purge) + + await cron.sync() + await cron.repository.update_one( + {"next_run_at": utc_now() - timedelta(seconds=1)}, + namespace="test", + name="purge_outbox", + ) + async with cron.running(): + with anyio.fail_after(2): + await ran.wait() + + async def test_run_executes_task_with_stored_args(cron): received = {} ran = anyio.Event() diff --git a/tests/test_database.py b/tests/test_database.py index 6ee559c..a59f9f1 100644 --- a/tests/test_database.py +++ b/tests/test_database.py @@ -1,7 +1,7 @@ -import asyncio import sqlite3 from datetime import datetime, timezone +import anyio import pytest import sqlalchemy as sa @@ -161,10 +161,12 @@ async def test_lock_serializes_tasks_under_one_name(db: Database): async def work(tag: str) -> None: async with db.lock("serialized"): order.append(f"enter {tag}") - await asyncio.sleep(0.01) + await anyio.sleep(0.01) order.append(f"exit {tag}") - await asyncio.gather(work("a"), work("b"), work("c")) + async with anyio.create_task_group() as tg: + for tag in "abc": + tg.start_soon(work, tag) # every enter is immediately followed by the matching exit assert [entry.split()[0] for entry in order] == ["enter", "exit"] * 3 diff --git a/tests/test_dialects.py b/tests/test_dialects.py index fb58516..5c1a35e 100644 --- a/tests/test_dialects.py +++ b/tests/test_dialects.py @@ -72,7 +72,14 @@ def test_get_query_builder_is_cached(): @pytest.mark.parametrize( ("dialect", "expected"), [ - ("postgresql", Option.RETURNING | Option.CONFLICTS | Option.LOCKS), + ( + "postgresql", + Option.RETURNING + | Option.CONFLICTS + | Option.LOCKS + | Option.VECTORS + | Option.FULL_TEXT, + ), ("mysql", Option.CONFLICTS | Option.LOCKS), ("oracle", Option.NONE), ], diff --git a/tests/test_eventiq.py b/tests/test_eventiq.py new file mode 100644 index 0000000..fbf1844 --- /dev/null +++ b/tests/test_eventiq.py @@ -0,0 +1,223 @@ +"""Tests for the eventiq integration layer.""" + +from datetime import timedelta +from typing import Any +from uuid import UUID, uuid4 + +import pytest +import sqlalchemy as sa +from eventiq import CloudEvent +from pydantic import ValidationError + +from sqlargon import Base, Database +from sqlargon.integrations.eventiq import eventiq_publisher, to_cloud_event +from sqlargon.mixins import CreatedUpdatedMixin, UUIDModelMixin +from sqlargon.outbox import ( + OutboxConfig, + OutboxEvent, + OutboxEventRepository, + OutboxRelay, + OutboxRepository, +) +from sqlargon.types import GUID +from sqlargon.utils import utc_now + +# Models defined at module level to avoid re-registration with --count=3 + + +class PublishedOrder(UUIDModelMixin, CreatedUpdatedMixin, Base): + __tablename__ = "test_eventiq_order" + + reference = sa.Column(sa.Unicode(255), nullable=True) + tenant_id = sa.Column(GUID(), nullable=True) + + +class OrderRepository(OutboxRepository[PublishedOrder]): + outbox = OutboxConfig( + topic="orders", + type_prefix="order", + exclude={"tenant_id"}, + attributes={"tenant_id": "tenant_id"}, + ) + + +class TenantEvent(CloudEvent[Any]): + """The base class a service of its own would publish.""" + + tenant_id: UUID + + +class FakeService: + """Stands in for ``eventiq.Service``, which needs a broker to build.""" + + def __init__(self) -> None: + self.published: list[CloudEvent] = [] + + async def publish(self, message: CloudEvent) -> None: + self.published.append(message) + + +@pytest.fixture(autouse=True) +async def tables(db: Database): + created = (PublishedOrder.__table__, OutboxEvent.__table__) + async with db.engine.begin() as conn: + for table in created: + await conn.run_sync(table.create, checkfirst=True) + yield + async with db.engine.begin() as conn: + for table in reversed(created): + await conn.run_sync(table.drop, checkfirst=True) + + +@pytest.fixture +def orders(): + return OrderRepository() + + +@pytest.fixture +def events(): + return OutboxEventRepository() + + +def build_event(**overrides) -> OutboxEvent: + values = { + "id": uuid4(), + "created_at": utc_now(), + "updated_at": utc_now(), + "available_at": utc_now(), + "topic": "orders", + "type": "order.created", + "source": None, + "data": {"reference": "abc"}, + "attributes": None, + "headers": None, + "attempts": 0, + } + return OutboxEvent(**{**values, **overrides}) + + +def test_to_cloud_event_maps_every_attribute(): + event = build_event(source="orders-service") + + message = to_cloud_event(event) + + assert message.id == event.id + assert message.time == event.created_at + assert message.topic == "orders" + assert message.type == "order.created" + assert message.source == "orders-service" + assert message.data == {"reference": "abc"} + assert message.specversion == "1.0" + + +def test_topic_serialises_as_the_cloud_events_subject(): + message = to_cloud_event(build_event()) + + assert message.model_dump(by_alias=True)["subject"] == "orders" + + +def test_the_row_source_wins_over_the_fallback(): + message = to_cloud_event(build_event(source="row"), source="fallback") + + assert message.source == "row" + + +def test_the_fallback_source_is_used_when_the_row_has_none(): + message = to_cloud_event(build_event(source=None), source="fallback") + + assert message.source == "fallback" + + +def test_stored_attributes_become_top_level_attributes(): + tenant = uuid4() + event = build_event(attributes={"tenant_id": str(tenant)}) + + message = to_cloud_event(event, event_class=TenantEvent) + + assert isinstance(message, TenantEvent) + assert message.tenant_id == tenant + assert message.model_dump()["tenant_id"] == tenant + + +def test_the_base_cloud_event_keeps_unknown_attributes(): + message = to_cloud_event(build_event(attributes={"tenant_id": "acme"})) + + assert message.model_dump()["tenant_id"] == "acme" + + +def test_an_event_class_missing_an_attribute_is_rejected(): + with pytest.raises(ValidationError): + to_cloud_event(build_event(), event_class=TenantEvent) + + +def test_stored_attributes_cannot_override_the_core_ones(): + event = build_event(attributes={"id": str(uuid4()), "type": "spoofed"}) + + message = to_cloud_event(event) + + assert message.id == event.id + assert message.type == "order.created" + + +def test_headers_are_handed_over_as_strings(): + message = to_cloud_event(build_event(headers={"retries": 2})) + + assert message.headers == {"retries": "2"} + + +async def test_the_publisher_hands_recorded_events_to_the_service(orders, events): + service = FakeService() + relay = OutboxRelay( + eventiq_publisher(service, source="orders-service"), + repository=OutboxEventRepository(), + ) + await orders.create(reference="abc") + + assert await relay.dispatch_once() == 1 + + (message,) = service.published + assert message.topic == "orders" + assert message.type == "order.created" + assert message.source == "orders-service" + assert message.data["reference"] == "abc" + assert await events.pending_count() == 0 + + +async def test_the_publisher_builds_the_configured_event_class(orders): + service = FakeService() + relay = OutboxRelay( + eventiq_publisher(service, event_class=TenantEvent, source="orders-service"), + repository=OutboxEventRepository(), + ) + tenant = uuid4() + await orders.create(reference="abc", tenant_id=tenant) + + assert await relay.dispatch_once() == 1 + + (message,) = service.published + assert isinstance(message, TenantEvent) + assert message.tenant_id == tenant + assert message.data["reference"] == "abc" + assert "tenant_id" not in message.data + + +async def test_a_service_failure_leaves_the_event_pending(orders, events): + class BrokenService(FakeService): + async def publish(self, message: CloudEvent) -> None: + msg = "broker down" + raise RuntimeError(msg) + + relay = OutboxRelay( + eventiq_publisher(BrokenService()), + repository=OutboxEventRepository(), + retry_backoff=1.0, + max_retry_delay=1.0, + ) + await orders.create(reference="abc") + + assert await relay.dispatch_once() == 0 + + assert await events.pending_count() == 1 + (event,) = await events.select().all() + assert event.last_error == "RuntimeError: broker down" + assert event.available_at <= utc_now() + timedelta(seconds=1) diff --git a/tests/test_i18n.py b/tests/test_i18n.py new file mode 100644 index 0000000..33b6bfe --- /dev/null +++ b/tests/test_i18n.py @@ -0,0 +1,190 @@ +"""Unit tests for the locale getter / fallback-chain callable slots +and for ``TranslatedRepository``. + +Mirrors ``test_outbox.py``: in-memory SQLite via the shared ``db`` fixture, +module-level models to survive ``--count=3`` re-registration. +""" + +import pytest +import sqlalchemy as sa +from sqlalchemy.orm import relationship + +import sqlargon.i18n.expression as _expr +import sqlargon.i18n.translation as _trans +from sqlargon import Base, Database +from sqlargon.i18n import ( + TranslatedRepository, + fallback_chain, + get_locale, + set_fallback_chain, + set_locale_getter, +) + +# ── models ────────────────────────────────────────────────────────────────── + + +class RepoModel(Base): + __tablename__ = "test_i18n_repo_model" + + id = sa.Column(sa.Integer, primary_key=True) + + +class RepoTranslation(Base): + __tablename__ = "test_i18n_repo_trans" + + id = sa.Column( + sa.Integer, + sa.ForeignKey("test_i18n_repo_model.id"), + primary_key=True, + ) + + +RepoModel._current_translation = relationship( + RepoTranslation, + primaryjoin=RepoModel.id == RepoTranslation.id, + uselist=False, + viewonly=True, + lazy="raise", +) + + +class TestRepo(TranslatedRepository): + """Concrete repository on a model whose ``_current_translation`` is a + plain relationship -- not one created by ``TranslatableMixin`` -- which + is enough to verify the join behaviour. + """ + + __test__ = False + model = RepoModel + + +# ── fixtures ──────────────────────────────────────────────────────────────── + + +@pytest.fixture(autouse=True) +def _reset_locale_slots(): + """Reset the locale and fallback callable slots before every test.""" + _expr._get_locale = None + _trans._get_fallback = None + + +@pytest.fixture(autouse=True) +async def _create_tables(db: Database): + created = (RepoModel.__table__, RepoTranslation.__table__) + async with db.engine.begin() as conn: + for table in created: + await conn.run_sync(table.create, checkfirst=True) + yield + async with db.engine.begin() as conn: + for table in reversed(created): + await conn.run_sync(table.drop, checkfirst=True) + + +# --- locale getter ----------------------------------------------------------- + + +def test_locale_getter_raises_before_configuration(): + with pytest.raises(RuntimeError, match="No locale getter"): + get_locale() + + +def test_locale_getter_returns_the_registered_value(): + set_locale_getter(lambda: "pl") + + assert get_locale() == "pl" + + +def test_locale_getter_replacing_is_honoured(): + set_locale_getter(lambda: "pl") + set_locale_getter(lambda: "de") + + assert get_locale() == "de" + + +def test_locale_getter_slot_is_cleared_between_tests(): + """The autouse fixture clears the slot, so a fresh test starts clean.""" + assert _expr._get_locale is None + + set_locale_getter(lambda: "fr") + + assert get_locale() == "fr" + + +# --- fallback chain ---------------------------------------------------------- + + +def test_fallback_chain_raises_before_configuration(): + with pytest.raises(RuntimeError, match="No fallback chain"): + fallback_chain() + + +def test_fallback_chain_returns_the_registered_chain(): + set_fallback_chain(lambda _: ("en", "en-US")) + + assert fallback_chain() == ("en", "en-US") + + +def test_fallback_chain_passes_the_explicit_locale_through(): + called_with: list[str | None] = [] + + def capture(locale: str | None) -> tuple[str, ...]: + called_with.append(locale) + return (locale or "en",) + + set_fallback_chain(capture) + + fallback_chain("de") + + assert called_with == ["de"] + + +def test_fallback_chain_passes_none_when_no_locale_is_given(): + called_with: list[str | None] = [] + + def capture(locale: str | None) -> tuple[str, ...]: + called_with.append(locale) + return ("en",) + + set_fallback_chain(capture) + + fallback_chain() + + assert called_with == [None] + + +def test_fallback_chain_replacing_is_honoured(): + set_fallback_chain(lambda _: ("en",)) + set_fallback_chain(lambda _: ("de", "en")) + + assert fallback_chain() == ("de", "en") + + +def test_fallback_chain_slot_is_cleared_between_tests(): + assert _trans._get_fallback is None + + set_fallback_chain(lambda _: ("fr", "en")) + + assert fallback_chain() == ("fr", "en") + + +# --- TranslatedRepository ---------------------------------------------------- + + +def test_translated_repository_select_includes_the_outer_join(): + repo = TestRepo() + stmt = repo.select().query + + compiled = str(stmt.compile(compile_kwargs={"literal_binds": True})) + + assert "LEFT OUTER JOIN" in compiled + assert "test_i18n_repo_trans" in compiled + + +def test_translated_repository_select_accepts_column_args(): + repo = TestRepo() + stmt = repo.select(RepoModel.id).query + + compiled = str(stmt.compile(compile_kwargs={"literal_binds": True})) + + assert "LEFT OUTER JOIN" in compiled + assert "test_i18n_repo_model" in compiled diff --git a/tests/test_mixins.py b/tests/test_mixins.py index d279f4a..f2f133b 100644 --- a/tests/test_mixins.py +++ b/tests/test_mixins.py @@ -206,6 +206,23 @@ async def test_is_new_false_after_update(repository): assert updated.is_new is False +@pytest.mark.usefixtures("tables") +async def test_is_new_false_after_an_upsert(repository): + """An upsert leaves ``created_at`` alone, so it cannot fake a new row.""" + obj = await repository.create(name="john", created_at=OLD, updated_at=OLD) + + upserted = await repository.create_or_update(id=obj.id, name="jane") + + assert upserted.created_at == OLD + assert upserted.updated_at > upserted.created_at + assert upserted.is_new is False + + +def test_the_conflict_set_leaves_created_at_alone(): + assert "created_at" not in _MixinRepository._get_default_set() + assert "updated_at" in _MixinRepository._get_default_set() + + def test_is_new_sql_expression(): assert ( str(_MixinModel.is_new.expression) diff --git a/tests/test_outbox.py b/tests/test_outbox.py new file mode 100644 index 0000000..7899d61 --- /dev/null +++ b/tests/test_outbox.py @@ -0,0 +1,795 @@ +"""Unit tests for OutboxRepository, OutboxEventRepository and OutboxRelay. + +Mirrors ``test_versioned.py``: in-memory SQLite via the shared ``db`` +fixture, module-level models to survive ``--count=3`` re-registration. +""" + +from contextvars import ContextVar +from datetime import timedelta +from uuid import uuid4 + +import anyio +import pytest +import sqlalchemy as sa + +from sqlargon import Base, Database, atomic +from sqlargon.mixins import CreatedUpdatedMixin, UUIDModelMixin +from sqlargon.outbox import ( + Operation, + OutboxConfig, + OutboxEvent, + OutboxEventRepository, + OutboxRelay, + OutboxRepository, + format_topic, +) +from sqlargon.types import GUID +from sqlargon.typing import OnConflictOptions +from sqlargon.utils import utc_now + +HASHED = "not-a-real-secret" + +traceparent: ContextVar[str] = ContextVar("traceparent", default="none") + +# Models defined at module level to avoid re-registration with --count=3 + + +class OutboxUser(UUIDModelMixin, CreatedUpdatedMixin, Base): + __tablename__ = "test_outbox_user" + + name = sa.Column(sa.Unicode(255), nullable=True) + password = sa.Column(sa.Unicode(255), nullable=True) + tenant_id = sa.Column(GUID(), nullable=True) + + +class InsertOnlyNote(UUIDModelMixin, CreatedUpdatedMixin, Base): + __tablename__ = "test_outbox_note" + + body = sa.Column(sa.Unicode(255), nullable=True) + + +class OutboxTag(UUIDModelMixin, CreatedUpdatedMixin, Base): + """A model with a unique key, so a conflicting insert really writes nothing.""" + + __tablename__ = "test_outbox_tag" + __table_args__ = (sa.UniqueConstraint("label"),) + + label = sa.Column(sa.Unicode(255), nullable=False) + + +class OrganizationUser(UUIDModelMixin, CreatedUpdatedMixin, Base): + """A row whose events land on a topic templated from its own key.""" + + __tablename__ = "test_outbox_org_user" + + name = sa.Column(sa.Unicode(255), nullable=True) + organization_id = sa.Column(GUID(), nullable=True) + + +class BareModel(UUIDModelMixin, CreatedUpdatedMixin, Base): + __tablename__ = "test_outbox_bare" + + +class PlainModel(CreatedUpdatedMixin, Base): + __tablename__ = "test_outbox_plain" + id = sa.Column(sa.Integer, primary_key=True) + + +class UntimestampedModel(UUIDModelMixin, Base): + """A model the outbox cannot tell an insert from an update on.""" + + __tablename__ = "test_outbox_untimestamped" + + +class OutboxUserRepository(OutboxRepository[OutboxUser]): + default_order_by = OutboxUser.name + + outbox = OutboxConfig( + topic="users", type_prefix="user", source="tests", exclude={"password"} + ) + + +class BareRepository(OutboxRepository[BareModel]): + """The repository with no ``outbox`` of its own.""" + + +class NarrowedUserRepository(OutboxUserRepository): + """Tightens the exclusions further, the way ``on_conflict`` is.""" + + outbox = OutboxConfig(topic="users", include=frozenset({"name"})) + + +class TenantUserRepository(OutboxRepository[OutboxUser]): + """Promotes a column and a context variable to CloudEvent attributes.""" + + outbox = OutboxConfig( + topic="users", + type_prefix="user", + exclude={"password", "tenant_id"}, + attributes={ + "tenant_id": "tenant_id", + "traceparent": lambda _: traceparent.get(), + }, + ) + + +class NoteRepository(OutboxRepository[InsertOnlyNote]): + outbox = OutboxConfig(operations=frozenset({Operation.CREATED})) + + +class TagRepository(OutboxRepository[OutboxTag]): + @property + def on_conflict(self) -> OnConflictOptions: + return {"index_elements": {"label"}, "set_": {"label"}} + + +class OrganizationUserRepository(OutboxRepository[OrganizationUser]): + outbox = OutboxConfig( + topic="events.organizations.{organization_id}.users.{id}.created", + type_prefix="org_user", + ) + + +class EventRepository(OutboxEventRepository): + pass + + +# --- fixtures --- + + +@pytest.fixture(autouse=True) +async def tables(db: Database): + created = ( + OutboxUser.__table__, + InsertOnlyNote.__table__, + OutboxTag.__table__, + OrganizationUser.__table__, + OutboxEvent.__table__, + ) + async with db.engine.begin() as conn: + for table in created: + await conn.run_sync(table.create, checkfirst=True) + yield + async with db.engine.begin() as conn: + for table in reversed(created): + await conn.run_sync(table.drop, checkfirst=True) + + +@pytest.fixture +def users(): + return OutboxUserRepository() + + +@pytest.fixture +def tenant_users(): + return TenantUserRepository() + + +@pytest.fixture +def notes(): + return NoteRepository() + + +@pytest.fixture +def tags(): + return TagRepository() + + +@pytest.fixture +def org_users(): + return OrganizationUserRepository() + + +@pytest.fixture +def events(): + return EventRepository() + + +async def stored(events: EventRepository) -> list[OutboxEvent]: + return list(await events.select().order_by(OutboxEvent.created_at).all()) + + +# --- configuration --- + + +def payload_of(repository: OutboxRepository) -> tuple[str, ...]: + return tuple(name for name, _ in repository.payload_columns) + + +def test_any_timestamped_model_can_be_recorded(): + """Nothing else marks a model as recorded -- the repository decides.""" + + class Repository(OutboxRepository[PlainModel]): + pass + + assert Repository.model is PlainModel + + +def test_a_model_without_timestamps_is_rejected(): + with pytest.raises(TypeError, match="CreatedUpdatedMixin"): + + class Repository(OutboxRepository[UntimestampedModel]): + pass + + +def test_a_repository_needs_no_configuration(): + assert isinstance(BareRepository().outbox, OutboxConfig) + assert BareRepository().topic == "test_outbox_bare" + + +def test_topic_and_type_default_to_the_table_name(): + notes = NoteRepository() + + assert notes.topic == "test_outbox_note" + assert notes.event_types[Operation.CREATED] == "test_outbox_note.created" + + +def test_repository_config_names_the_topic_and_type(users): + assert users.topic == "users" + assert users.event_types[Operation.DELETED] == "user.deleted" + + +def test_excluded_columns_are_kept_out_of_the_payload(users): + assert "password" not in payload_of(users) + assert "name" in payload_of(users) + + +def test_include_wins_over_exclude(): + class Repository(OutboxRepository[OutboxUser]): + outbox = OutboxConfig(include=frozenset({"name"}), exclude=frozenset({"name"})) + + assert payload_of(Repository()) == ("name",) + + +def test_only_declared_operations_get_an_event_type(): + types = NoteRepository().event_types + + assert Operation.CREATED in types + assert Operation.DELETED not in types + + +def test_the_config_belongs_to_the_class(users): + assert users.outbox is OutboxUserRepository.outbox + + +def test_a_subclass_overrides_the_config_of_its_base(): + assert NarrowedUserRepository().outbox is not OutboxUserRepository.outbox + assert payload_of(NarrowedUserRepository()) == ("name",) + assert "password" not in payload_of(OutboxUserRepository()) + + +def test_attributes_cannot_shadow_the_core_ones(): + with pytest.raises(ValueError, match="source, type"): + OutboxConfig(attributes={"type": "name", "source": "name", "tenant": "name"}) + + +# --- topic templating --- + + +def test_format_topic_passes_a_plain_topic_through(): + assert format_topic("users", object()) == "users" + + +def test_format_topic_fills_placeholders_from_the_row(): + row = OrganizationUser(id=uuid4(), organization_id=uuid4()) + + assert ( + format_topic("events.organizations.{organization_id}.deleted", row) + == f"events.organizations.{row.organization_id}.deleted" + ) + + +def test_format_topic_fills_an_id_placeholder(): + row = OrganizationUser(id=uuid4(), organization_id=uuid4()) + + assert format_topic("events.users.{id}", row) == f"events.users.{row.id}" + + +def test_a_missing_attribute_raises_like_str_format_does(): + row = OrganizationUser() + + with pytest.raises(KeyError, match="organization_id"): + format_topic("events.organizations.{organization_id}", row) + + +async def test_a_templated_topic_is_filled_per_event(org_users, events): + organization = uuid4() + user = await org_users.create(name="John", organization_id=organization) + + (event,) = await stored(events) + assert event.topic == ( + f"events.organizations.{organization}.users.{user.id}.created" + ) + + +async def test_each_row_of_a_bulk_write_gets_its_own_topic(org_users, events): + first, second = uuid4(), uuid4() + + await org_users.bulk_create( + [ + {"name": "a", "organization_id": first}, + {"name": "b", "organization_id": second}, + ] + ) + + recorded = {event.data["name"]: event.topic for event in await stored(events)} + assert recorded["a"].startswith(f"events.organizations.{first}.") + assert recorded["b"].startswith(f"events.organizations.{second}.") + assert recorded["a"] != recorded["b"] + + +async def test_a_templated_topic_is_filled_at_write_time(org_users, events): + organization = uuid4() + user = await org_users.create(name="John", organization_id=organization) + + await org_users.update_one({"name": "Jane"}, id=user.id) + + topics = [event.topic for event in await stored(events)] + assert topics == [ + f"events.organizations.{organization}.users.{user.id}.created", + f"events.organizations.{organization}.users.{user.id}.created", + ] + + +# --- capture --- + + +async def test_create_records_an_event(users, events): + user = await users.create(name="John", password=HASHED) + + (event,) = await stored(events) + assert event.topic == "users" + assert event.type == "user.created" + assert event.source == "tests" + assert event.data["name"] == "John" + assert event.data["id"] == str(user.id) + assert "password" not in event.data + assert event.published_at is None + assert event.attempts == 0 + + +async def test_an_upsert_that_updates_records_an_update(users, events): + user = await users.create(name="John") + await users.create_or_update(id=user.id, name="Jane") + + types = [event.type for event in await stored(events)] + assert types == ["user.created", "user.updated"] + + +async def test_an_upsert_that_inserts_records_a_create(users, events): + await users.create_or_update(id=uuid4(), name="John") + + assert [event.type for event in await stored(events)] == ["user.created"] + + +async def test_one_upsert_records_what_each_row_turned_out_to_be(users, events): + existing = await users.create(name="John") + + await users.bulk_create_or_update( + [{"id": existing.id, "name": "Jane"}, {"id": uuid4(), "name": "Jack"}] + ) + + recorded = {event.data["name"]: event.type for event in await stored(events)} + assert recorded["Jane"] == "user.updated" + assert recorded["Jack"] == "user.created" + + +async def test_an_upsert_leaves_the_creation_time_alone(users): + user = await users.create(name="John") + + updated = await users.create_or_update(id=user.id, name="Jane") + + assert updated.created_at == user.created_at + assert updated.updated_at > updated.created_at + + +async def test_update_one_records_an_update(users, events): + user = await users.create(name="John") + await users.update_one({"name": "Jane"}, id=user.id) + + event = (await stored(events))[-1] + assert event.type == "user.updated" + assert event.data["name"] == "Jane" + + +async def test_update_many_records_one_event_per_row(users, events): + await users.create_many([{"name": "a"}, {"name": "b"}]) + await users.update_many({"password": HASHED}, OutboxUser.name.in_(["a", "b"])) + + updated = [e for e in await stored(events) if e.type == "user.updated"] + assert sorted(event.data["name"] for event in updated) == ["a", "b"] + + +async def test_delete_one_records_a_delete(users, events): + user = await users.create(name="John") + await users.delete_one(id=user.id) + + event = (await stored(events))[-1] + assert event.type == "user.deleted" + assert event.data["name"] == "John" + + +@pytest.mark.parametrize("method", ["remove", "delete_many"]) +async def test_bare_deletes_record_events(users, events, method): + await users.create_many([{"name": "a"}, {"name": "b"}]) + + await getattr(users, method)(OutboxUser.name == "a") + + deleted = [e for e in await stored(events) if e.type == "user.deleted"] + assert [event.data["name"] for event in deleted] == ["a"] + assert await users.count() == 1 + + +async def test_bulk_create_records_events_without_returning_results(users, events): + result = await users.bulk_create([{"name": "a"}, {"name": "b"}]) + + assert result is None + created = [e for e in await stored(events) if e.type == "user.created"] + assert sorted(event.data["name"] for event in created) == ["a", "b"] + + +async def test_bulk_create_returning_results_still_returns_rows(users, events): + rows = await users.bulk_create([{"name": "a"}], return_results=True) + + assert [row.name for row in rows] == ["a"] + assert len(await stored(events)) == 1 + + +async def test_bulk_create_or_update_records_events(users, events): + user = await users.create(name="a") + result = await users.bulk_create_or_update([{"id": user.id, "name": "b"}]) + + assert [row.name for row in result.scalars()] == ["b"] + assert [event.type for event in await stored(events)] == [ + "user.created", + "user.updated", + ] + + +async def test_bulk_update_reads_the_rows_back(users, events): + created = await users.create_many([{"name": "a"}, {"name": "b"}]) + await users.bulk_update( + [{"id": row.id, "password": HASHED} for row in created], + ) + + updated = [e for e in await stored(events) if e.type == "user.updated"] + assert sorted(event.data["name"] for event in updated) == ["a", "b"] + + +async def test_get_or_create_records_only_the_created_row(tags, events): + first = await tags.get_or_create(label="x") + again = await tags.get_or_create(label="x") + + assert again.id == first.id + assert [event.type for event in await stored(events)] == ["test_outbox_tag.created"] + + +async def test_undeclared_operations_are_not_recorded(notes, events): + note = await notes.create(body="hello") + await notes.delete_one(id=note.id) + + assert [event.type for event in await stored(events)] == [ + "test_outbox_note.created" + ] + + +async def test_repository_can_narrow_the_payload(events): + await NarrowedUserRepository().create(name="John", password=HASHED) + + (event,) = await stored(events) + assert event.data == {"name": "John"} + + +async def test_events_roll_back_with_the_row(users, events): + class Failure(RuntimeError): + pass + + @atomic + async def write(repository) -> None: + await repository.create(name="John") + raise Failure + + with pytest.raises(Failure): + await write(users) + + assert await users.count() == 0 + assert await events.count() == 0 + + +# --- extra attributes --- + + +async def test_attributes_come_from_the_row_and_the_context(tenant_users, events): + tenant = uuid4() + traceparent.set("00-trace-span-01") + + await tenant_users.create(name="John", tenant_id=tenant) + + (event,) = await stored(events) + assert event.attributes == { + "tenant_id": str(tenant), + "traceparent": "00-trace-span-01", + } + + +async def test_a_promoted_column_leaves_the_payload(tenant_users, events): + await tenant_users.create(name="John", tenant_id=uuid4()) + + (event,) = await stored(events) + assert "tenant_id" not in event.data + assert event.data["name"] == "John" + + +async def test_events_carry_no_attributes_without_configuration(users, events): + await users.create(name="John") + + (event,) = await stored(events) + assert event.attributes is None + + +@pytest.mark.parametrize( + ("operation", "event_type"), + [("create", "user.created"), ("update_one", "user.updated")], +) +async def test_every_operation_carries_the_attributes( + tenant_users, events, operation, event_type +): + tenant = uuid4() + user = await tenant_users.create(name="John", tenant_id=tenant) + if operation == "update_one": + await tenant_users.update_one({"name": "Jane"}, id=user.id) + + event = (await stored(events))[-1] + assert event.type == event_type + assert event.attributes["tenant_id"] == str(tenant) + + +async def test_deletes_carry_the_attributes(tenant_users, events): + tenant = uuid4() + user = await tenant_users.create(name="John", tenant_id=tenant) + + await tenant_users.delete_one(id=user.id) + + event = (await stored(events))[-1] + assert event.type == "user.deleted" + assert event.attributes["tenant_id"] == str(tenant) + + +async def test_bulk_writes_carry_the_attributes_of_each_row(tenant_users, events): + tenants = [uuid4(), uuid4()] + + await tenant_users.bulk_create( + [{"name": "a", "tenant_id": tenants[0]}, {"name": "b", "tenant_id": tenants[1]}] + ) + + recorded = {event.data["name"]: event.attributes for event in await stored(events)} + assert recorded["a"]["tenant_id"] == str(tenants[0]) + assert recorded["b"]["tenant_id"] == str(tenants[1]) + + +# --- event repository --- + + +async def test_claim_pending_leases_and_counts_attempts(users, events): + await users.create(name="John") + now = utc_now() + + claimed = await events.claim_pending(now, lease=timedelta(minutes=5)) + + assert len(claimed) == 1 + assert claimed[0].attempts == 1 + assert claimed[0].available_at == now + timedelta(minutes=5) + assert await events.claim_pending(now, lease=timedelta(minutes=5)) == [] + + +async def test_claim_pending_skips_exhausted_events(users, events): + await users.create(name="John") + now = utc_now() + + await events.claim_pending(now, lease=timedelta(0), max_attempts=1) + + assert await events.claim_pending(now, lease=timedelta(0), max_attempts=1) == [] + + +async def test_claim_pending_returns_events_in_write_order(users, events): + for name in ("a", "b", "c"): + await users.create(name=name) + + claimed = await events.claim_pending(utc_now(), lease=timedelta(minutes=5)) + + assert [event.data["name"] for event in claimed] == ["a", "b", "c"] + + +async def test_mark_published_clears_the_error(users, events): + await users.create(name="John") + (event,) = await stored(events) + await events.mark_failed(event.id, "boom", utc_now()) + + await events.mark_published([event.id], utc_now()) + + (stored_event,) = await stored(events) + assert stored_event.published_at is not None + assert stored_event.last_error is None + assert await events.pending_count() == 0 + + +async def test_mark_published_without_ids_is_a_no_op(events): + await events.mark_published([], utc_now()) + + assert await events.count() == 0 + + +async def test_purge_only_deletes_published_events(users, events): + await users.create_many([{"name": "a"}, {"name": "b"}]) + first, second = await stored(events) + cutoff = utc_now() + await events.mark_published([first.id], cutoff - timedelta(days=8)) + await events.mark_published([second.id], cutoff) + + deleted = await events.purge(cutoff - timedelta(days=7)) + + assert deleted == 1 + assert [event.id for event in await stored(events)] == [second.id] + + +# --- relay --- + + +@pytest.fixture +def published(): + return [] + + +@pytest.fixture +def relay(published): + async def publisher(event: OutboxEvent) -> None: + published.append(event) + + return OutboxRelay(publisher, repository=EventRepository(), poll_interval=0.05) + + +async def test_dispatch_once_publishes_and_marks(users, events, relay, published): + await users.create(name="John") + + assert await relay.dispatch_once() == 1 + + assert [event.data["name"] for event in published] == ["John"] + assert await events.pending_count() == 0 + assert await relay.dispatch_once() == 0 + + +async def test_dispatch_once_preserves_order(users, relay, published): + for name in ("a", "b", "c"): + await users.create(name=name) + + await relay.dispatch_once() + + assert [event.data["name"] for event in published] == ["a", "b", "c"] + + +async def test_a_failing_publisher_stops_the_batch(users, events, published): + for name in ("a", "b"): + await users.create(name=name) + + async def publisher(event: OutboxEvent) -> None: + published.append(event) + msg = "broker down" + raise RuntimeError(msg) + + relay = OutboxRelay( + publisher, + repository=EventRepository(), + retry_backoff=1.0, + max_retry_delay=1.0, + ) + + assert await relay.dispatch_once() == 0 + assert [event.data["name"] for event in published] == ["a"] + failed = next(e for e in await stored(events) if e.last_error) + assert failed.data["name"] == "a" + assert failed.last_error == "RuntimeError: broker down" + assert failed.published_at is None + + +async def test_retry_delay_grows_and_is_capped(relay): + relay.retry_backoff = 2.0 + relay.max_retry_delay = 4.0 + + delays = [relay.retry_delay(attempts).total_seconds() for attempts in (1, 2, 3, 9)] + + assert delays == [1.0, 2.0, 4.0, 4.0] + + +async def test_purge_uses_the_configured_retention(users, events, published): + await users.create(name="John") + (event,) = await stored(events) + await events.mark_published([event.id], utc_now() - timedelta(days=8)) + + async def publisher(_event: OutboxEvent) -> None: + published.append(_event) + + relay = OutboxRelay( + publisher, repository=EventRepository(), retention=timedelta(days=7) + ) + + assert await relay.purge() == 1 + + +async def test_a_tick_leaves_published_events_alone(users, events, relay): + await users.create(name="John") + (event,) = await stored(events) + await events.mark_published([event.id], utc_now() - timedelta(days=8)) + + await relay._tick() + + assert await events.count() == 1 + + +async def parked(events: EventRepository) -> None: + """Wait until the relay has finished its batch and is back in its sleep. + + Leaving ``running()`` cancels the relay wherever it happens to be, and a + cancel landing inside a statement terminates the one connection the + in-memory SQLite pool hands out -- which the next test then finds closed. + The publishers below stretch the poll interval on their way through, so + once the batch is marked published the relay stays parked long enough for + the cancel to be harmless. + """ + for _ in range(200): + if not await events.pending_count(): + return + await anyio.sleep(0.01) + pytest.fail("the relay never drained the table") + + +async def test_run_publishes_in_the_background(users, events, published): + await users.create(name="John") + arrived = anyio.Event() + + async def publisher(event: OutboxEvent) -> None: + published.append(event) + relay.poll_interval = 5.0 + arrived.set() + + relay = OutboxRelay(publisher, repository=EventRepository(), poll_interval=0.05) + async with relay.running(): + with anyio.fail_after(2): + await arrived.wait() + await parked(events) + + assert [event.data["name"] for event in published] == ["John"] + + +async def test_run_survives_a_failing_poll( + users, events, published, monkeypatch, caplog +): + await users.create(name="John") + arrived = anyio.Event() + calls = 0 + + async def publisher(event: OutboxEvent) -> None: + published.append(event) + relay.poll_interval = 5.0 + arrived.set() + + relay = OutboxRelay(publisher, repository=EventRepository(), poll_interval=0.05) + original = relay.repository.claim_pending + + async def flaky(*args, **kwargs): + nonlocal calls + calls += 1 + if calls == 1: + msg = "database gone" + raise RuntimeError(msg) + return await original(*args, **kwargs) + + monkeypatch.setattr(relay.repository, "claim_pending", flaky) + + with caplog.at_level("ERROR", logger="sqlargon.outbox.relay"): + async with relay.running(): + with anyio.fail_after(2): + await arrived.wait() + await parked(events) + + assert "Dispatching outbox events failed" in caplog.text + assert [event.data["name"] for event in published] == ["John"] diff --git a/tests/test_types.py b/tests/test_types.py index eb8cc40..4fcbf0e 100644 --- a/tests/test_types.py +++ b/tests/test_types.py @@ -1,3 +1,4 @@ +import contextlib from datetime import datetime, timedelta, timezone from uuid import UUID, uuid4 @@ -10,12 +11,23 @@ from sqlargon.orm import Base as _Base from sqlargon.types import GUID, JSON, GenerateUUID, GenerateUUIDV7, Timestamp, now from sqlargon.types.json import ( + json_array_append, + json_array_length, json_contains, + json_get, json_has_all_keys, json_has_any_key, + json_has_key, + json_insert_key, + json_keys, + json_remove_key, + json_replace_key, + json_set_key, + json_update, json_value, ) from sqlargon.types.pydantic import Pydantic, ValidatedType +from sqlargon.utils import json_loads def _compile(expr, dialect, *, literal_binds=True) -> str: @@ -346,7 +358,12 @@ def test_json_value_init(): assert element.key == "key" assert element.name == "json_value" assert isinstance(element.type, sa.String) - assert list(element.clauses) == [_json_col] + # the key and its JSON path are operands, not compile time literals, so + # that one cached statement can serve every key + column, key, path = element.clauses + assert column is _json_col + assert key.value == "key" + assert path.value == '$."key"' @pytest.mark.parametrize( @@ -362,6 +379,303 @@ def test_json_value_compiles_per_dialect(dialect, expected): assert _compile(json_value(_json_col, "k"), dialect) == expected +# --- JSON mutations and reads, per-dialect compilation --- + + +@pytest.mark.parametrize( + ("dialect", "expected"), + [ + (postgresql.dialect(), """(data || CAST('{"a":1}' AS JSONB))"""), + (sqlite.dialect(), """json_set(data, '$."a"', json('1'))"""), + (mysql.dialect(), """json_set(data, '$."a"', json_extract('1', '$'))"""), + (DefaultDialect(), """json_set(data, '$."a"', json_extract('1', '$'))"""), + ], +) +def test_json_update_compiles_per_dialect(dialect, expected): + assert _compile(json_update(_json_col, {"a": 1}), dialect) == expected + + +def test_json_set_key_is_a_single_key_update(): + assert _compile(json_set_key(_json_col, "a", 1), sqlite.dialect()) == _compile( + json_update(_json_col, {"a": 1}), sqlite.dialect() + ) + + +@pytest.mark.parametrize( + "dialect", [sqlite.dialect(), mysql.dialect(), DefaultDialect()] +) +def test_json_update_merges_every_key_in_one_call(dialect): + sql = _compile(json_update(_json_col, {"a": 1, "b": 2}), dialect) + # one json_set, not a nested pair -- and shallow, so a nested object + # replaces rather than merges + assert sql.count("json_set(") == 1 + assert """'$."a"'""" in sql + assert """'$."b"'""" in sql + + +@pytest.mark.parametrize( + ("dialect", "expected"), + [ + (postgresql.dialect(), "(data - CAST(ARRAY['a', 'b'] AS TEXT[]))"), + (sqlite.dialect(), """json_remove(data, '$."a"', '$."b"')"""), + (mysql.dialect(), """json_remove(data, '$."a"', '$."b"')"""), + (DefaultDialect(), """json_remove(data, '$."a"', '$."b"')"""), + ], +) +def test_json_remove_key_compiles_per_dialect(dialect, expected): + assert _compile(json_remove_key(_json_col, "a", "b"), dialect) == expected + + +@pytest.mark.parametrize( + ("dialect", "expected"), + [ + (postgresql.dialect(), """(CAST('{"a":1}' AS JSONB) || data)"""), + (sqlite.dialect(), """json_insert(data, '$."a"', json('1'))"""), + (mysql.dialect(), """json_insert(data, '$."a"', json_extract('1', '$'))"""), + (DefaultDialect(), """json_insert(data, '$."a"', json_extract('1', '$'))"""), + ], +) +def test_json_insert_key_compiles_per_dialect(dialect, expected): + # the patch goes on the left of the postgres concat so an existing key wins + assert _compile(json_insert_key(_json_col, "a", 1), dialect) == expected + + +@pytest.mark.parametrize( + ("dialect", "expected"), + [ + ( + postgresql.dialect(), + "jsonb_set(data, CAST(ARRAY['a'] AS TEXT[]), CAST('1' AS JSONB), false)", + ), + (sqlite.dialect(), """json_replace(data, '$."a"', json('1'))"""), + (mysql.dialect(), """json_replace(data, '$."a"', json_extract('1', '$'))"""), + (DefaultDialect(), """json_replace(data, '$."a"', json_extract('1', '$'))"""), + ], +) +def test_json_replace_key_compiles_per_dialect(dialect, expected): + # create_missing => false is what keeps postgres from inserting the key + assert _compile(json_replace_key(_json_col, "a", 1), dialect) == expected + + +@pytest.mark.parametrize( + ("dialect", "expected"), + [ + ( + postgresql.dialect(), + """(data || jsonb_build_array(CAST('"x"' AS JSONB)))""", + ), + (sqlite.dialect(), """json_insert(data, '$[#]', json('"x"'))"""), + ( + mysql.dialect(), + """json_array_append(data, '$', json_extract('"x"', '$'))""", + ), + ( + DefaultDialect(), + """json_array_append(data, '$', json_extract('"x"', '$'))""", + ), + ], +) +def test_json_array_append_compiles_per_dialect(dialect, expected): + assert _compile(json_array_append(_json_col, "x"), dialect) == expected + + +@pytest.mark.parametrize( + ("dialect", "expected"), + [ + (postgresql.dialect(), "(data -> 'k')"), + (sqlite.dialect(), """json_extract(data, '$."k"')"""), + (mysql.dialect(), """json_extract(data, '$."k"')"""), + (DefaultDialect(), """json_extract(data, '$."k"')"""), + ], +) +def test_json_get_compiles_per_dialect(dialect, expected): + assert _compile(json_get(_json_col, "k"), dialect) == expected + + +@pytest.mark.parametrize( + ("dialect", "expected"), + [ + (postgresql.dialect(), "data ? 'k'"), + (sqlite.dialect(), """json_type(data, '$."k"') IS NOT NULL"""), + (mysql.dialect(), """json_contains_path(data, 'one', '$."k"')"""), + (DefaultDialect(), """json_contains_path(data, 'one', '$."k"')"""), + ], +) +def test_json_has_key_compiles_per_dialect(dialect, expected): + assert _compile(json_has_key(_json_col, "k"), dialect) == expected + + +@pytest.mark.parametrize( + ("dialect", "expected"), + [ + (postgresql.dialect(), "jsonb_array_length(data)"), + (sqlite.dialect(), "json_array_length(data)"), + (mysql.dialect(), "json_length(data)"), + (DefaultDialect(), "json_length(data)"), + ], +) +def test_json_array_length_compiles_per_dialect(dialect, expected): + assert _compile(json_array_length(_json_col), dialect) == expected + + +def test_json_keys_compiles_per_dialect(): + assert _compile(json_keys(_json_col), mysql.dialect()) == "json_keys(data)" + assert _compile(json_keys(_json_col), DefaultDialect()) == "json_keys(data)" + + postgres = _compile(json_keys(_json_col), postgresql.dialect()) + assert "jsonb_object_keys(data)" in postgres + # jsonb_agg over an object with no keys is NULL, not an empty array + assert "coalesce" in postgres + assert "CAST('[]' AS JSONB)" in postgres + + assert "json_group_array(json_each.key)" in _compile( + json_keys(_json_col), sqlite.dialect() + ) + + +@pytest.mark.parametrize( + "dialect", + [postgresql.dialect(), sqlite.dialect(), mysql.dialect(), DefaultDialect()], +) +def test_json_update_with_no_keys_is_the_column(dialect): + # json_set(col) with no pair is a syntax error, and merging nothing is + # the column itself + assert _compile(json_update(_json_col, {}), dialect) == "data" + + +@pytest.mark.parametrize( + "dialect", + [postgresql.dialect(), sqlite.dialect(), mysql.dialect(), DefaultDialect()], +) +def test_json_remove_key_with_no_keys_is_the_column(dialect): + assert _compile(json_remove_key(_json_col), dialect) == "data" + + +@pytest.mark.parametrize( + ("factory", "message"), + [(json_update, "json_update keys"), (json_remove_key, "json_remove_key keys")], +) +def test_json_mutation_keys_must_be_strings(factory, message): + argument = {1: "a"} if factory is json_update else 1 + with pytest.raises(ValueError, match=message): + factory(_json_col, argument) + + +def test_json_mutations_nest(): + expression = json_remove_key(json_update(_json_col, {"a": 1}), "b") + assert ( + _compile(expression, sqlite.dialect()) + == """json_remove(json_set(data, '$."a"', json('1')), '$."b"')""" + ) + + +def test_json_mutations_nest_on_postgresql_without_losing_precedence(): + # binary "-" binds tighter than "||" in postgres, so an unparenthesized + # "a || b - c" would drop the key from the patch, not from the result + expression = json_remove_key(json_update(_json_col, {"a": 1}), "b") + assert ( + _compile(expression, postgresql.dialect()) + == """((data || CAST('{"a":1}' AS JSONB)) - CAST(ARRAY['b'] AS TEXT[]))""" + ) + + +@pytest.mark.parametrize( + ("key", "expected"), + [ + ("plain", '$."plain"'), + ('we"ird', '$."we\\"ird"'), + ("back\\slash", '$."back\\\\slash"'), + ], +) +def test_json_path_escapes_the_key(key, expected): + # the path is an operand, so assert the value we hand the driver; how it + # is then quoted into SQL text differs per dialect and is SQLAlchemy's job + _column, _mapping, path, _value = json_set_key(_json_col, key, 1).clauses + assert path.value == expected + + +def test_json_mutation_values_are_bound_not_inlined(): + compiled = json_update(_json_col, {"a": {"nested": True}}).compile( + dialect=sqlite.dialect() + ) + assert {"nested": True} in compiled.params.values() + + +def test_json_mutation_values_serialize_to_a_json_document(): + # json() / json_extract() re-parse this text, so the value lands as a + # document rather than as a JSON string holding the serialized text. + # A bare dialect serializes with the stdlib, an engine with orjson, so + # compare the document and not the spacing. + serialized = JSON().bind_processor(sqlite.dialect())({"nested": True}) + assert json_loads(serialized) == {"nested": True} + + +# the comparator is the documented entry point, so every method has to +# resolve through an InstrumentedAttribute and compile +_COMPARATOR_CALLS = [ + ("set_key", lambda c: c.set_key("a", 1), True), + ("update", lambda c: c.update({"a": 1}), True), + ("remove_key", lambda c: c.remove_key("a"), True), + ("insert_key", lambda c: c.insert_key("a", 1), True), + ("replace_key", lambda c: c.replace_key("a", 1), True), + ("array_append", lambda c: c.array_append(1), True), + ("get", lambda c: c.get("a"), True), + ("keys", lambda c: c.keys(), True), + # these answer with a boolean, an int and text, so they are ends of a + # chain rather than links in one + ("has_key", lambda c: c.has_key("a"), False), + ("array_length", lambda c: c.array_length(), False), + ("json_value", lambda c: c.json_value("a"), False), + ("contains", lambda c: c.contains(["a"]), False), +] + + +_COMPARATOR_CASES = [ + (call, returns_json) for _name, call, returns_json in _COMPARATOR_CALLS +] +_COMPARATOR_IDS = [name for name, _call, _returns_json in _COMPARATOR_CALLS] + + +@pytest.mark.parametrize( + ("call", "returns_json"), _COMPARATOR_CASES, ids=_COMPARATOR_IDS +) +def test_json_comparator_methods_resolve_and_compile(call, returns_json): + expression = call(_JsonMutationModel.data) + assert _compile(expression, sqlite.dialect()) + # a JSON-typed result carries the comparator again, which is what lets + # the mutations chain: col.set_key(...).remove_key(...) + assert isinstance(expression.type, JSON) is returns_json + + +@pytest.mark.parametrize( + ("call", "returns_json"), _COMPARATOR_CASES, ids=_COMPARATOR_IDS +) +def test_json_comparator_methods_chain_when_they_return_json(call, returns_json): + expression = call(_JsonMutationModel.data) + assert hasattr(expression, "remove_key") is returns_json + + +@pytest.mark.parametrize( + "factory", + [ + lambda key: json_update(_json_col, {key: 1}), + lambda key: json_remove_key(_json_col, key), + lambda key: json_get(_json_col, key), + lambda key: json_has_key(_json_col, key), + lambda key: json_value(_json_col, key), + ], +) +def test_json_keys_are_bound_so_a_cached_statement_serves_any_key(factory): + # the keys live in the clause list, so two expressions share one compiled + # statement and each execution binds its own key. A key baked in by a + # @compiles hook would instead be reused for every later key. + first, second = sa.select(factory("a")), sa.select(factory("b")) + assert first._generate_cache_key() == second._generate_cache_key() + assert str(first.compile(dialect=sqlite.dialect())) == str( + second.compile(dialect=sqlite.dialect()) + ) + + # --- JSON integration (SQLite) --- # Models defined at module level to avoid re-registration with --count=3 @@ -396,6 +710,12 @@ class _JsonValueModel(_Base): data = sa.Column(JSON()) +class _JsonMutationModel(_Base): + __tablename__ = "test_json_mutation" + id = sa.Column(sa.Integer, primary_key=True, autoincrement=True) + data = sa.Column(JSON()) + + async def test_json_column_crud(db): async with db.engine.begin() as conn: await conn.run_sync(_JsonCrudModel.__table__.create, checkfirst=True) @@ -616,3 +936,168 @@ def test_validated_type_no_validation(): def test_validated_type_custom_sa_column_type(): vtype = ValidatedType(list[int], sa_column_type=sa.JSON) assert vtype.impl == sa.JSON + + +# --- JSON mutations (SQLite) --- + + +@contextlib.asynccontextmanager +async def _mutation_table(db, initial): + """The mutation table holding one row of ``initial``.""" + async with db.engine.begin() as conn: + await conn.run_sync(_JsonMutationModel.__table__.create, checkfirst=True) + try: + async with db.session() as session: + session.add(_JsonMutationModel(id=1, data=initial)) + yield + finally: + async with db.engine.begin() as conn: + await conn.run_sync(_JsonMutationModel.__table__.drop, checkfirst=True) + + +async def _mutate(db, expression): + """Apply ``expression`` to the row's data column and read it back.""" + async with db.session() as session: + await session.execute( + sa.update(_JsonMutationModel).values({_JsonMutationModel.data: expression}) + ) + await session.commit() + async with db.session() as session: + return await session.scalar(sa.select(_JsonMutationModel.data)) + + +async def test_json_set_key_adds_a_key(db): + async with _mutation_table(db, {"a": 1}): + column = _JsonMutationModel.data + assert await _mutate(db, json_set_key(column, "b", 2)) == {"a": 1, "b": 2} + + +async def test_json_set_key_overwrites_a_key(db): + async with _mutation_table(db, {"a": 1}): + column = _JsonMutationModel.data + assert await _mutate(db, json_set_key(column, "a", 9)) == {"a": 9} + + +async def test_json_set_key_stores_a_nested_document(db): + async with _mutation_table(db, {"a": 1}): + column = _JsonMutationModel.data + # not the serialized text as a JSON string + assert await _mutate(db, json_set_key(column, "b", {"x": [1, 2]})) == { + "a": 1, + "b": {"x": [1, 2]}, + } + + +async def test_json_update_merges_shallowly(db): + async with _mutation_table(db, {"a": {"x": 1}, "b": 2}): + column = _JsonMutationModel.data + # "a" is replaced wholesale rather than merged into + assert await _mutate(db, json_update(column, {"a": {"y": 9}, "c": 3})) == { + "a": {"y": 9}, + "b": 2, + "c": 3, + } + + +async def test_json_update_with_no_keys_leaves_the_document(db): + async with _mutation_table(db, {"a": 1}): + column = _JsonMutationModel.data + assert await _mutate(db, json_update(column, {})) == {"a": 1} + + +async def test_json_remove_key_drops_keys(db): + async with _mutation_table(db, {"a": 1, "b": 2, "c": 3}): + column = _JsonMutationModel.data + assert await _mutate(db, json_remove_key(column, "a", "c")) == {"b": 2} + + +async def test_json_remove_key_ignores_a_missing_key(db): + async with _mutation_table(db, {"a": 1}): + column = _JsonMutationModel.data + assert await _mutate(db, json_remove_key(column, "nope")) == {"a": 1} + + +async def test_json_insert_key_only_adds_a_missing_key(db): + async with _mutation_table(db, {"a": 1}): + column = _JsonMutationModel.data + assert await _mutate(db, json_insert_key(column, "b", 2)) == {"a": 1, "b": 2} + + +async def test_json_insert_key_leaves_an_existing_key(db): + async with _mutation_table(db, {"a": 1}): + column = _JsonMutationModel.data + assert await _mutate(db, json_insert_key(column, "a", 9)) == {"a": 1} + + +async def test_json_replace_key_only_updates_an_existing_key(db): + async with _mutation_table(db, {"a": 1}): + column = _JsonMutationModel.data + assert await _mutate(db, json_replace_key(column, "a", 9)) == {"a": 9} + + +async def test_json_replace_key_does_not_add_a_missing_key(db): + async with _mutation_table(db, {"a": 1}): + column = _JsonMutationModel.data + assert await _mutate(db, json_replace_key(column, "b", 2)) == {"a": 1} + + +async def test_json_array_append_appends_one_element(db): + async with _mutation_table(db, [1, 2]): + column = _JsonMutationModel.data + assert await _mutate(db, json_array_append(column, 3)) == [1, 2, 3] + + +async def test_json_array_append_nests_a_list_rather_than_concatenating(db): + async with _mutation_table(db, [1]): + column = _JsonMutationModel.data + assert await _mutate(db, json_array_append(column, [2, 3])) == [1, [2, 3]] + + +async def test_json_mutations_compose_in_one_statement(db): + async with _mutation_table(db, {"a": 1, "b": 2}): + column = _JsonMutationModel.data + expression = json_remove_key(json_update(column, {"c": 3}), "a") + assert await _mutate(db, expression) == {"b": 2, "c": 3} + + +async def test_json_mutation_of_a_null_column_propagates_null(db): + async with _mutation_table(db, None): + column = _JsonMutationModel.data + assert await _mutate(db, json_set_key(column, "a", 1)) is None + + +async def test_json_reads_on_sqlite(db): + async with _mutation_table(db, {"a": {"x": 1}, "b": [1, 2, 3]}): + column = _JsonMutationModel.data + async with db.session() as session: + assert await session.scalar(sa.select(json_get(column, "a"))) == {"x": 1} + assert await session.scalar(sa.select(json_has_key(column, "a"))) + assert not await session.scalar(sa.select(json_has_key(column, "nope"))) + assert ( + await session.scalar( + sa.select(json_array_length(json_get(column, "b"))) + ) + == 3 + ) + assert sorted(await session.scalar(sa.select(json_keys(column)))) == [ + "a", + "b", + ] + + +async def test_json_has_key_addresses_object_keys_not_values(db): + # the gap json_has_any_key / json_has_all_keys leave on sqlite, where + # their json_each fallback matches values instead + async with _mutation_table(db, {"a": "b"}): + column = _JsonMutationModel.data + async with db.session() as session: + assert await session.scalar(sa.select(json_has_key(column, "a"))) + assert not await session.scalar(sa.select(json_has_key(column, "b"))) + + +async def test_json_comparator_methods_reach_the_mutations(db): + async with _mutation_table(db, {"a": 1}): + column = _JsonMutationModel.data + assert await _mutate(db, column.set_key("b", 2)) == {"a": 1, "b": 2} + assert await _mutate(db, column.remove_key("a")) == {"b": 2} + assert await _mutate(db, column.update({"c": 3})) == {"b": 2, "c": 3} diff --git a/tests/test_vectors.py b/tests/test_vectors.py new file mode 100644 index 0000000..3db6591 --- /dev/null +++ b/tests/test_vectors.py @@ -0,0 +1,580 @@ +import struct +import sys + +import pytest +import sqlalchemy as sa +from sqlalchemy.dialects import mysql, postgresql, sqlite +from sqlalchemy.orm import declared_attr + +from sqlargon import Base, Database, SQLAlchemyRepository +from sqlargon.dialects.mysql import MysqlQueryBuilder +from sqlargon.dialects.postgres import PostgresqlQueryBuilder +from sqlargon.dialects.sqlite import SQLiteQueryBuilder +from sqlargon.mixins import UUIDV7ModelMixin +from sqlargon.query_builder import Option, QueryBuilder +from sqlargon.types.vector import ( + cosine_distance, + distance_for, + l1_distance, + l2_distance, + max_inner_product, +) +from sqlargon.vectors import ( + AttributesMixin, + DistanceMetric, + EmbeddingBase, + EmbeddingMixin, + HybridVectorRepository, + TextBase, + TextEmbeddingBase, + TextMixin, + UnsupportedDialectError, + Vector, + VectorCollection, + VectorCollectionRepository, + VectorDocument, + VectorRepository, + init_vectors, + register_sqlite_vector, +) +from tests import MEMORY_URL + +# Models defined at module level to avoid re-registration with --count=3 + +VECTOR = [1.0, 2.0, 3.0] + + +class Note(UUIDV7ModelMixin, EmbeddingBase): + """An embedding and nothing else -- the minimum the extension supports.""" + + __tablename__ = "test_vectors_note" + __vector_dim__ = 3 + + @declared_attr.directive + def __table_args__(cls) -> tuple[sa.Index, ...]: + return (cls.embedding_index(),) + + +class Article(UUIDV7ModelMixin, TextBase): + """Text without an embedding.""" + + __tablename__ = "test_vectors_article" + __text_regconfig__ = "english" + + @declared_attr.directive + def __table_args__(cls) -> tuple[sa.Index, ...]: + return (cls.text_index(),) + + +class Chunk(UUIDV7ModelMixin, AttributesMixin, EmbeddingBase): + """Embedding plus attributes, composed from the mixins.""" + + __tablename__ = "test_vectors_chunk" + __vector_dim__ = 4 + __vector_distance__ = DistanceMetric.L2 + + @declared_attr.directive + def __table_args__(cls) -> tuple[sa.Index, ...]: + return (cls.embedding_index(), cls.attributes_index()) + + +class Document(VectorDocument): + """Everything: embedding, text, attributes and a collection.""" + + __tablename__ = "test_vectors_document" + __vector_dim__ = 3 + + @declared_attr.directive + def __table_args__(cls) -> tuple[sa.Index, ...]: + return (cls.embedding_index(), cls.attributes_index(), cls.text_index()) + + +class Plain(Base): + __tablename__ = "test_vectors_plain" + id = sa.Column(sa.Integer, primary_key=True) + + +class Composite(EmbeddingMixin, Base): + """A composite primary key, which the rank fusion cannot key rows by.""" + + __tablename__ = "test_vectors_composite" + __vector_dim__ = 3 + left = sa.Column(sa.Integer, primary_key=True) + right = sa.Column(sa.Integer, primary_key=True) + + +class NoteRepository(VectorRepository[Note]): + pass + + +class ArticleRepository(SQLAlchemyRepository[Article]): + pass + + +class ChunkRepository(VectorRepository[Chunk]): + pass + + +class DocumentRepository(HybridVectorRepository[Document]): + pass + + +class CompositeRepository(VectorRepository[Composite]): + pass + + +def _compile(expr, dialect, *, literal_binds=True) -> str: + """Compile an expression against ``dialect`` without executing it.""" + return str( + expr.compile(dialect=dialect, compile_kwargs={"literal_binds": literal_binds}) + ) + + +# --- Vector type --- + + +@pytest.mark.parametrize( + ("dialect", "expected"), + [ + (postgresql.dialect(), "VECTOR(3)"), + (sqlite.dialect(), "BLOB"), + (mysql.dialect(), "JSON"), + ], +) +def test_vector_column_type_per_dialect(dialect, expected): + assert expected in _compile( + sa.schema.CreateColumn(Note.__table__.c.embedding), dialect + ) + + +def test_vector_dimensions_are_configurable(): + assert "VECTOR(4)" in _compile( + sa.schema.CreateColumn(Chunk.__table__.c.embedding), postgresql.dialect() + ) + + +def test_vector_sqlite_round_trip(): + vector = Vector(3) + packed = vector.process_bind_param(VECTOR, sqlite.dialect()) + assert packed == struct.pack("<3f", *VECTOR) + assert vector.process_result_value(packed, sqlite.dialect()) == VECTOR + + +def test_vector_postgresql_bind_passes_the_list_through(): + assert Vector(3).process_bind_param(VECTOR, postgresql.dialect()) == VECTOR + + +def test_vector_normalises_non_list_results(): + assert ( + Vector(3).process_result_value((1.0, 2.0, 3.0), postgresql.dialect()) == VECTOR + ) + + +@pytest.mark.parametrize("dialect", [postgresql.dialect(), sqlite.dialect()]) +def test_vector_none_stays_none(dialect): + vector = Vector(3) + assert vector.process_bind_param(None, dialect) is None + assert vector.process_result_value(None, dialect) is None + + +# --- distance expressions --- + + +@pytest.mark.parametrize( + ("metric", "operator"), + [ + (DistanceMetric.COSINE, "<=>"), + (DistanceMetric.L2, "<->"), + (DistanceMetric.DOT, "<#>"), + (DistanceMetric.L1, "<+>"), + ], +) +def test_distance_compiles_to_the_pgvector_operator(metric, operator): + expression = distance_for(metric)(Note.embedding, VECTOR) + assert operator in _compile(expression, postgresql.dialect()) + + +@pytest.mark.parametrize( + "element", [cosine_distance, l2_distance, max_inner_product, l1_distance] +) +def test_distance_refuses_to_compile_on_sqlite(element): + """The regression guard: pgvector's own comparator would emit ``<=>`` here.""" + with pytest.raises(UnsupportedDialectError, match="sqlite"): + _compile(element(Note.embedding, VECTOR), sqlite.dialect()) + + +def test_comparator_exposes_the_distance_methods(): + assert "<=>" in _compile( + Note.embedding.cosine_distance(VECTOR), postgresql.dialect() + ) + assert "<->" in _compile(Note.embedding.l2_distance(VECTOR), postgresql.dialect()) + assert "<#>" in _compile( + Note.embedding.max_inner_product(VECTOR), postgresql.dialect() + ) + assert "<+>" in _compile(Note.embedding.l1_distance(VECTOR), postgresql.dialect()) + + +def test_metric_maps_to_pgvector_names(): + assert DistanceMetric.COSINE.pg_opclass == "vector_cosine_ops" + assert DistanceMetric.L2.pg_opclass == "vector_l2_ops" + assert DistanceMetric.DOT.pg_opclass == "vector_ip_ops" + assert DistanceMetric.L1.pg_opclass == "vector_l1_ops" + assert DistanceMetric.COSINE.sqlite_option == "COSINE" + + +# --- composability --- + + +def test_embedding_only_model_has_no_other_columns(): + assert sorted(c.name for c in Note.__table__.c) == ["embedding", "id"] + + +def test_mixins_compose_into_the_columns_they_add(): + assert sorted(c.name for c in Chunk.__table__.c) == [ + "attributes", + "embedding", + "id", + ] + assert sorted(c.name for c in Article.__table__.c) == ["id", "text"] + + +def test_vector_document_carries_every_column(): + assert sorted(c.name for c in Document.__table__.c) == [ + "attributes", + "collection_id", + "created_at", + "embedding", + "id", + "text", + "updated_at", + ] + + +def test_vector_collection_is_concrete(): + assert VectorCollection.__tablename__ == "vector_collection" + assert VectorCollectionRepository.model is VectorCollection + + +# --- indexes --- + + +def _index_ddl(model, name, dialect=postgresql.dialect()) -> str: + index = next(i for i in model.__table__.indexes if i.name == name) + return _compile(sa.schema.CreateIndex(index), dialect) + + +def test_embedding_index_uses_hnsw_with_the_metric_opclass(): + ddl = _index_ddl(Note, "ix_test_vectors_note__embedding") + assert "USING hnsw" in ddl + assert "vector_cosine_ops" in ddl + assert "m = 16" in ddl + assert "ef_construction = 64" in ddl + + +def test_embedding_index_follows_the_configured_metric(): + assert "vector_l2_ops" in _index_ddl(Chunk, "ix_test_vectors_chunk__embedding") + + +def test_attributes_index_uses_gin(): + ddl = _index_ddl(Chunk, "ix_test_vectors_chunk__attributes") + assert "USING gin" in ddl + assert "jsonb_path_ops" in ddl + + +def test_text_index_inlines_the_search_configuration(): + ddl = _index_ddl(Article, "ix_test_vectors_article__text") + assert "USING gin" in ddl + assert "to_tsvector('english', text)" in ddl + + +def test_custom_index_options(): + index = Note.embedding_index("custom_name", m=32, ef_construction=128) + options = index.dialect_options["postgresql"] + assert index.name == "custom_name" + assert options["using"] == "hnsw" + assert options["with"] == {"m": 32, "ef_construction": 128} + + +def test_invalid_regconfig_is_rejected(): + class Sneaky(TextMixin): + __text_regconfig__ = "english'; DROP TABLE users --" + + with pytest.raises(ValueError, match="text search configuration"): + Sneaky.text_document() + + +@pytest.mark.anyio +async def test_postgresql_only_indexes_are_skipped_on_sqlite(db: Database): + """``ddl_if`` keeps HNSW and GIN DDL out of the SQLite schema.""" + async with db.engine.begin() as conn: + await conn.run_sync(Document.__table__.create, checkfirst=True) + result = await conn.exec_driver_sql( + "SELECT name FROM sqlite_master WHERE type = 'index'" + ) + names = {row[0] for row in result} + await conn.run_sync(Document.__table__.drop, checkfirst=True) + assert "ix_test_vectors_document__embedding" not in names + assert "ix_test_vectors_document__text" not in names + assert "ix_test_vectors_document__collection_id" in names + + +# --- attributes filtering --- + + +def test_attributes_contain_builds_a_containment_predicate(): + assert "@>" in _compile( + Chunk.attributes_contain({"lang": "en"}), + postgresql.dialect(), + literal_binds=False, + ) + + +# --- repository model validation --- + + +@pytest.mark.parametrize( + ("repository", "model", "missing"), + [ + (VectorRepository, Article, "EmbeddingMixin"), + (HybridVectorRepository, Chunk, "TextMixin"), + (VectorRepository, Plain, "EmbeddingMixin"), + ], +) +def test_repository_rejects_a_model_without_the_mixin(repository, model, missing): + with pytest.raises(TypeError, match=missing): + + class Bad(repository[model]): + pass + + +def test_repository_accepts_a_model_carrying_the_mixin(): + assert NoteRepository.model is Note + assert DocumentRepository.model is Document + + +def test_hybrid_repository_requires_both_mixins(): + assert issubclass(Document, EmbeddingMixin) + assert issubclass(Document, TextMixin) + assert issubclass(TextEmbeddingBase, TextMixin) + + +# --- query builder capabilities --- + +PG = PostgresqlQueryBuilder() +SQLITE = SQLiteQueryBuilder() +MYSQL = MysqlQueryBuilder() + + +@pytest.mark.parametrize( + ("builder", "option", "supported"), + [ + (PG, Option.VECTORS, True), + (PG, Option.FULL_TEXT, True), + (SQLITE, Option.VECTORS, True), + (SQLITE, Option.FULL_TEXT, False), + (MYSQL, Option.VECTORS, False), + (MYSQL, Option.FULL_TEXT, False), + (QueryBuilder(), Option.VECTORS, False), + ], +) +def test_search_capability_claims(builder, option, supported): + assert builder.supports(option) is supported + + +def test_a_builder_without_vectors_refuses_to_build_one(): + builder = QueryBuilder() + with pytest.raises(UnsupportedDialectError, match="vector search"): + builder.vector_search(Note, VECTOR, limit=5) + with pytest.raises(UnsupportedDialectError, match="distance"): + builder.vector_distance(Note, VECTOR) + + +def test_a_builder_without_full_text_refuses_to_build_one(): + with pytest.raises(UnsupportedDialectError, match="full text search"): + SQLITE.text_search(Document, "hello", limit=5) + with pytest.raises(UnsupportedDialectError, match="reciprocal rank fusion"): + SQLITE.rrf_search(Document, VECTOR, "hello") + + +# --- search statements --- + + +def test_pg_search_orders_by_distance(): + sql = _compile( + PG.vector_search(Note, VECTOR, limit=5), + postgresql.dialect(), + literal_binds=False, + ) + assert "<=>" in sql + assert "ORDER BY" in sql + assert "LIMIT" in sql + + +def test_pg_search_applies_the_filters_it_is_given(): + sql = _compile( + PG.vector_search( + Chunk, + VECTOR, + Chunk.attributes_contain({"lang": "en"}), + Chunk.id.is_(None), + limit=5, + ), + postgresql.dialect(), + literal_binds=False, + ) + assert "@>" in sql + assert "IS NULL" in sql + + +def test_pg_search_honours_a_metric_override(): + sql = _compile( + PG.vector_search(Note, VECTOR, limit=5, metric=DistanceMetric.L2), + postgresql.dialect(), + literal_binds=False, + ) + assert "<->" in sql + + +def test_sqlite_search_joins_the_streaming_scan(): + sql = _compile( + SQLITE.vector_search(Note, VECTOR, limit=5), + sqlite.dialect(), + literal_binds=False, + ) + assert "vector_full_scan" in sql + assert "rowid" in sql + assert "ORDER BY" in sql + + +def test_sqlite_search_keeps_the_filters(): + sql = _compile( + SQLITE.vector_search(Chunk, VECTOR, Chunk.id.is_(None), limit=5), + sqlite.dialect(), + literal_binds=False, + ) + assert "vector_full_scan" in sql + assert "IS NULL" in sql + + +def test_sqlite_rejects_a_metric_the_column_was_not_built_for(): + with pytest.raises(UnsupportedDialectError, match="per column"): + SQLITE.vector_search(Note, VECTOR, limit=5, metric=DistanceMetric.L2) + + +def test_sqlite_accepts_the_metric_the_column_was_built_for(): + query = SQLITE.vector_search(Note, VECTOR, limit=5, metric=DistanceMetric.COSINE) + assert "vector_full_scan" in _compile(query, sqlite.dialect(), literal_binds=False) + + +def test_sqlite_declares_the_column_per_connection(): + sql = _compile(SQLITE.vector_init(Chunk), sqlite.dialect()) + assert "vector_init" in sql + assert "dimension=4" in sql + assert "distance=L2" in sql + assert "type=FLOAT32" in sql + + +@pytest.mark.parametrize("builder", [PG, MYSQL, QueryBuilder()]) +def test_only_sqlite_needs_a_declaration(builder): + assert builder.vector_init(Note) is None + + +def test_pg_text_search_ranks_by_ts_rank(): + sql = _compile( + PG.text_search(Document, "hello world", limit=5), + postgresql.dialect(), + literal_binds=False, + ) + assert "ts_rank" in sql + assert "websearch_to_tsquery" in sql + assert "@@" in sql + + +def test_rrf_query_fuses_both_rankings(): + sql = _compile( + PG.rrf_search(Document, VECTOR, "hello", k=60, limit=5, candidates=50), + postgresql.dialect(), + literal_binds=False, + ) + assert "vector_candidates" in sql + assert "text_candidates" in sql + assert "FULL OUTER JOIN" in sql + assert "ORDER BY rrf.score DESC" in sql + + +def test_identity_column_requires_a_single_primary_key(): + with pytest.raises(TypeError, match="single-column primary key"): + PG.identity_column(Composite) + + +# --- dialect guards --- + + +@pytest.mark.anyio +async def test_search_rejects_an_unsupported_dialect(): + repository = NoteRepository().using( + db=Database("mysql+asyncmy://user@localhost/db") + ) + with pytest.raises(UnsupportedDialectError, match="mysql"): + await repository.search(VECTOR) + + +@pytest.mark.anyio +async def test_text_search_is_postgresql_only(): + with pytest.raises(UnsupportedDialectError, match="text_search"): + await DocumentRepository().text_search("hello") + + +@pytest.mark.anyio +async def test_rrf_search_is_postgresql_only(): + with pytest.raises(UnsupportedDialectError, match="rrf_search"): + await DocumentRepository().rrf_search(VECTOR, "hello") + + +@pytest.mark.anyio +async def test_search_rejects_a_metric_override_on_sqlite(): + with pytest.raises(UnsupportedDialectError, match="per column"): + await NoteRepository().search(VECTOR, metric=DistanceMetric.L2) + + +# --- loader --- + + +@pytest.mark.anyio +async def test_init_vectors_rejects_an_unsupported_dialect(): + database = Database("mysql+asyncmy://user@localhost/db") + with pytest.raises(UnsupportedDialectError, match="mysql"): + await init_vectors(database) + + +@pytest.mark.anyio +async def test_register_sqlite_vector_reports_the_missing_package(monkeypatch, db): + monkeypatch.setitem(sys.modules, "sqlite_vector", None) + with pytest.raises(ImportError, match="vectors-sqlite"): + register_sqlite_vector(db.engine) + + +@pytest.mark.anyio +async def test_init_vectors_loads_the_sqlite_extension(): + """Proof the extension reaches the connection, not just the pool.""" + pytest.importorskip("sqlite_vector") + database = Database(MEMORY_URL) + try: + await init_vectors(database) + version = await database.execute(sa.select(sa.func.vector_version())) + assert version.scalar() + finally: + await database.dispose() + + +@pytest.mark.anyio +async def test_registering_the_same_engine_twice_is_a_no_op(): + pytest.importorskip("sqlite_vector") + database = Database(MEMORY_URL) + try: + register_sqlite_vector(database.engine) + register_sqlite_vector(database.engine) + version = await database.execute(sa.select(sa.func.vector_version())) + assert version.scalar() + finally: + await database.dispose() diff --git a/uv.lock b/uv.lock index e82e635..6f2db8b 100644 --- a/uv.lock +++ b/uv.lock @@ -175,15 +175,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/77/f5/21d2de20e8b8b0408f0681956ca2c69f1320a3848ac50e6e7f39c6159675/babel-2.18.0-py3-none-any.whl", hash = "sha256:e2b422b277c2b9a9630c1d7903c2a00d0830c409c59ac8cae9081c92f1aeba35", size = 10196845, upload-time = "2026-02-01T12:30:53.445Z" }, ] -[[package]] -name = "backports-asyncio-runner" -version = "1.2.0" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/8e/ff/70dca7d7cb1cbc0edb2c6cc0c38b65cba36cccc491eca64cabd5fe7f8670/backports_asyncio_runner-1.2.0.tar.gz", hash = "sha256:a5aa7b2b7d8f8bfcaa2b57313f70792df84e32a2a746f585213373f900b42162", size = 69893, upload-time = "2025-07-02T02:27:15.685Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/a0/59/76ab57e3fe74484f48a53f8e337171b4a2349e506eabe136d7e01d059086/backports_asyncio_runner-1.2.0-py3-none-any.whl", hash = "sha256:0da0a936a8aeb554eccb426dc55af3ba63bcdc69fa1a600b5bb305413a4477b5", size = 12313, upload-time = "2025-07-02T02:27:14.263Z" }, -] - [[package]] name = "backrefs" version = "6.2" @@ -689,6 +680,22 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/75/23/529140fe1aab80fc6992f93a706deec709140a6397439139a054e1515c45/docker-7.2.0-py3-none-any.whl", hash = "sha256:a3f45fdeb9165e2d25d9a1d02ddf3bc70fb572cf5ebbf9b58558c22caf29b71f", size = 148775, upload-time = "2026-07-09T14:53:45.224Z" }, ] +[[package]] +name = "eventiq" +version = "1.1.14" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio" }, + { name = "pydantic" }, + { name = "pydantic-asyncapi" }, + { name = "pydantic-settings" }, + { name = "typer" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/29/26/4dde2aaa36f3bb8d97b84b017aa8969c6ec344e7693fd86c570ffce8ad4d/eventiq-1.1.14.tar.gz", hash = "sha256:3c66c0516398bd2410d4f8983c575b4f28af66962e05eabff9dbbef5997aeb07", size = 38181, upload-time = "2026-04-13T07:26:43.609Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/99/22/4ca6e6428b2b87e944889bebcd722b0cfff25593e8515f8d8487df7e6c02/eventiq-1.1.14-py3-none-any.whl", hash = "sha256:66d4ed52de9687ec7097eb7664ba86beace88dcec3bab949c2de6816515f2943", size = 45574, upload-time = "2026-04-13T07:26:42.391Z" }, +] + [[package]] name = "exceptiongroup" version = "1.3.1" @@ -701,6 +708,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/8a/0e/97c33bf5009bdbac74fd2beace167cab3f978feb69cc36f1ef79360d6c4e/exceptiongroup-1.3.1-py3-none-any.whl", hash = "sha256:a7a39a3bd276781e98394987d3a5701d0c4edffb633bb7a5144577f82c773598", size = 16740, upload-time = "2025-11-21T23:01:53.443Z" }, ] +[[package]] +name = "execnet" +version = "2.1.2" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/bf/89/780e11f9588d9e7128a3f87788354c7946a9cbb1401ad38a48c4db9a4f07/execnet-2.1.2.tar.gz", hash = "sha256:63d83bfdd9a23e35b9c6a3261412324f964c2ec8dcd8d3c6916ee9373e0befcd", size = 166622, upload-time = "2025-11-12T09:56:37.75Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/ab/84/02fc1827e8cdded4aa65baef11296a9bbe595c474f0d6d758af082d849fd/execnet-2.1.2-py3-none-any.whl", hash = "sha256:67fba928dd5a544b783f6056f449e5e3931a5c378b128bc18501f7ea79e296ec", size = 40708, upload-time = "2025-11-12T09:56:36.333Z" }, +] + [[package]] name = "fastapi" version = "0.138.0" @@ -1469,6 +1485,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/ef/3c/2c197d226f9ea224a9ab8d197933f9da0ae0aac5b6e0f884e2b8d9c8e9f7/pathspec-1.0.4-py3-none-any.whl", hash = "sha256:fb6ae2fd4e7c921a165808a552060e722767cfa526f99ca5156ed2ce45a5c723", size = 55206, upload-time = "2026-01-27T03:59:45.137Z" }, ] +[[package]] +name = "pgvector" +version = "0.5.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/a7/ec/6eb80aebc728200f95229219882994c1b0585b956ca47da5edb9d062627a/pgvector-0.5.0.tar.gz", hash = "sha256:07a9dcf735696879406983afc6eba9a787cef7c0cf6c367ca1a5779f036dee74", size = 35170, upload-time = "2026-07-06T18:27:27.767Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/5c/e4/a5573f2c579ca9ad133293bfb624148ba0893674ca4a6eeec85ced9a6a09/pgvector-0.5.0-py3-none-any.whl", hash = "sha256:fedc9800894e6da2be51358d7b7c574bf34f247ca741a5a09513622135f5964f", size = 30958, upload-time = "2026-07-06T18:27:26.797Z" }, +] + [[package]] name = "platformdirs" version = "4.9.6" @@ -1534,6 +1559,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/5a/87/b70ad306ebb6f9b585f114d0ac2137d792b48be34d732d60e597c2f8465a/pydantic-2.12.5-py3-none-any.whl", hash = "sha256:e561593fccf61e8a20fc46dfc2dfe075b8be7d0188df33f221ad1f0139180f9d", size = 463580, upload-time = "2025-11-26T15:11:44.605Z" }, ] +[[package]] +name = "pydantic-asyncapi" +version = "0.2.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pydantic" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/2b/86/699fb6125a321c536f3a17d4623bd0e743fb41fe240ca842cacb9e898a2c/pydantic_asyncapi-0.2.1.tar.gz", hash = "sha256:d9894c09b5d5ad308d42da01c5f75f6dd0425d47744cd58e6abadc999b5edbe2", size = 11333, upload-time = "2024-07-26T11:32:55.185Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/20/a5/9648ef0235d78ca1220e8c8bddc9067092664336bc4bc4105d94d63667d6/pydantic_asyncapi-0.2.1-py3-none-any.whl", hash = "sha256:eefb5e37f27556cc42809b58481d6e3416cca3c2219b4155a1a67c4da747e137", size = 10783, upload-time = "2024-07-26T11:32:53.966Z" }, +] + [[package]] name = "pydantic-core" version = "2.41.5" @@ -1706,20 +1743,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d4/24/a372aaf5c9b7208e7112038812994107bc65a84cd00e0354a88c2c77a617/pytest-9.0.3-py3-none-any.whl", hash = "sha256:2c5efc453d45394fdd706ade797c0a81091eccd1d6e4bccfcd476e2b8e0ab5d9", size = 375249, upload-time = "2026-04-07T17:16:16.13Z" }, ] -[[package]] -name = "pytest-asyncio" -version = "1.3.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "backports-asyncio-runner", marker = "python_full_version < '3.11'" }, - { name = "pytest" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/90/2c/8af215c0f776415f3590cac4f9086ccefd6fd463befeae41cd4d3f193e5a/pytest_asyncio-1.3.0.tar.gz", hash = "sha256:d7f52f36d231b80ee124cd216ffb19369aa168fc10095013c6b014a34d3ee9e5", size = 50087, upload-time = "2025-11-10T16:07:47.256Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/e5/35/f8b19922b6a25bc0880171a2f1a003eaeb93657475193ab516fd87cac9da/pytest_asyncio-1.3.0-py3-none-any.whl", hash = "sha256:611e26147c7f77640e6d0a92a38ed17c3e9848063698d5c93d5aa7aa11cebff5", size = 15075, upload-time = "2025-11-10T16:07:45.537Z" }, -] - [[package]] name = "pytest-cov" version = "7.1.0" @@ -1771,6 +1794,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/fa/b6/3127540ecdf1464a00e5a01ee60a1b09175f6913f0644ac748494d9c4b21/pytest_timeout-2.4.0-py3-none-any.whl", hash = "sha256:c42667e5cdadb151aeb5b26d114aff6bdf5a907f176a007a30b940d3d865b5c2", size = 14382, upload-time = "2025-05-05T19:44:33.502Z" }, ] +[[package]] +name = "pytest-xdist" +version = "3.8.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "execnet" }, + { name = "pytest" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/78/b4/439b179d1ff526791eb921115fca8e44e596a13efeda518b9d845a619450/pytest_xdist-3.8.0.tar.gz", hash = "sha256:7e578125ec9bc6050861aa93f2d59f1d8d085595d6551c2c90b6f4fad8d3a9f1", size = 88069, upload-time = "2025-07-01T13:30:59.346Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/ca/31/d4e37e9e550c2b92a9cbc2e4d0b7420a27224968580b5a447f420847c975/pytest_xdist-3.8.0-py3-none-any.whl", hash = "sha256:202ca578cfeb7370784a8c33d6d05bc6e13b4f25b5053c30a152269fd10f0b88", size = 46396, upload-time = "2025-07-01T13:30:56.632Z" }, +] + [[package]] name = "python-dateutil" version = "2.9.0.post0" @@ -1958,6 +1994,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/03/36/76704c4f312257d6dbaae3c959add2a622f63fcca9d864659ce6d8d97d3d/ruff-0.15.9-py3-none-win_arm64.whl", hash = "sha256:0694e601c028fd97dc5c6ee244675bc241aeefced7ef80cd9c6935a871078f53", size = 11005870, upload-time = "2026-04-02T18:17:15.773Z" }, ] +[[package]] +name = "shellingham" +version = "1.5.4" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/58/15/8b3609fd3830ef7b27b655beb4b4e9c62313a4e8da8c676e142cc210d58e/shellingham-1.5.4.tar.gz", hash = "sha256:8dbca0739d487e5bd35ab3ca4b36e11c4078f3a234bfce294b0a0291363404de", size = 10310, upload-time = "2023-10-24T04:13:40.426Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/e0/f9/0595336914c5619e5f28a1fb793285925a8cd4b432c9da0a987836c7f822/shellingham-1.5.4-py2.py3-none-any.whl", hash = "sha256:7ecfff8f2fd72616f7481040475a65b2bf8af90a56c89140852d1120324e8686", size = 9755, upload-time = "2023-10-24T04:13:38.866Z" }, +] + [[package]] name = "six" version = "1.17.0" @@ -2060,6 +2105,9 @@ cron = [ { name = "anyio" }, { name = "croniter" }, ] +eventiq = [ + { name = "eventiq" }, +] fastapi = [ { name = "fastapi" }, ] @@ -2069,6 +2117,9 @@ mysql = [ opentelemetry = [ { name = "opentelemetry-instrumentation-sqlalchemy" }, ] +outbox = [ + { name = "anyio" }, +] pagination = [ { name = "sqlakeyset" }, ] @@ -2086,12 +2137,20 @@ standard = [ { name = "opentelemetry-instrumentation-sqlalchemy" }, { name = "sqlakeyset" }, ] +vectors = [ + { name = "pgvector" }, +] +vectors-sqlite = [ + { name = "sqliteai-vector" }, +] [package.dev-dependencies] dev = [ + { name = "anyio" }, { name = "bandit" }, { name = "cryptography" }, { name = "deptry" }, + { name = "eventiq" }, { name = "fastapi" }, { name = "greenlet" }, { name = "httpx" }, @@ -2102,12 +2161,13 @@ dev = [ { name = "mkdocstrings", extra = ["python"] }, { name = "mypy" }, { name = "pytest" }, - { name = "pytest-asyncio" }, { name = "pytest-cov" }, { name = "pytest-repeat" }, { name = "pytest-sugar" }, { name = "pytest-timeout" }, + { name = "pytest-xdist" }, { name = "ruff" }, + { name = "sqliteai-vector" }, { name = "testcontainers" }, { name = "watchdog" }, ] @@ -2121,6 +2181,7 @@ docs = [ ] e2e = [ { name = "cryptography" }, + { name = "sqliteai-vector" }, { name = "testcontainers" }, ] lint = [ @@ -2130,13 +2191,14 @@ lint = [ { name = "ruff" }, ] test = [ + { name = "anyio" }, { name = "httpx" }, { name = "pytest" }, - { name = "pytest-asyncio" }, { name = "pytest-cov" }, { name = "pytest-repeat" }, { name = "pytest-sugar" }, { name = "pytest-timeout" }, + { name = "pytest-xdist" }, ] [package.metadata] @@ -2145,30 +2207,36 @@ requires-dist = [ { name = "aiosqlite", marker = "extra == 'standard'", specifier = ">=0.19.0,<1" }, { name = "alembic", specifier = ">=1.13.1,<2" }, { name = "anyio", marker = "extra == 'cron'", specifier = ">=4.0,<5" }, + { name = "anyio", marker = "extra == 'outbox'", specifier = ">=4.0,<5" }, { name = "anyio", marker = "extra == 'standard'", specifier = ">=4.0,<5" }, { name = "asyncmy", marker = "extra == 'mysql'", specifier = ">=0.2.11" }, { name = "asyncpg", marker = "extra == 'postgres'", specifier = "<1.0" }, { name = "asyncpg", marker = "extra == 'standard'", specifier = "<1.0" }, { name = "croniter", marker = "extra == 'cron'", specifier = ">=2.0,<7" }, { name = "croniter", marker = "extra == 'standard'", specifier = ">=2.0,<7" }, + { name = "eventiq", marker = "extra == 'eventiq'", specifier = ">=1.1.14,<2" }, { name = "fastapi", marker = "extra == 'fastapi'", specifier = ">=0.115,<1" }, { name = "opentelemetry-instrumentation-sqlalchemy", marker = "extra == 'opentelemetry'" }, { name = "opentelemetry-instrumentation-sqlalchemy", marker = "extra == 'standard'" }, { name = "orjson", specifier = ">=3.11.9,<4" }, + { name = "pgvector", marker = "extra == 'vectors'", specifier = ">=0.5.0" }, { name = "pydantic", specifier = ">=2.0,<3" }, { name = "pydantic-settings", specifier = ">=2.1.0,<3" }, { name = "sqlakeyset", marker = "extra == 'pagination'", specifier = ">=2.0.1716332987,<3" }, { name = "sqlakeyset", marker = "extra == 'standard'", specifier = ">=2.0.1716332987,<3" }, { name = "sqlalchemy", specifier = ">2.0,<3" }, + { name = "sqliteai-vector", marker = "extra == 'vectors-sqlite'", specifier = ">=1.0.0,<2" }, { name = "uuid-utils", specifier = ">=0.16.0,<1" }, ] -provides-extras = ["fastapi", "postgres", "sqlite", "mysql", "pagination", "cron", "opentelemetry", "standard"] +provides-extras = ["fastapi", "postgres", "sqlite", "mysql", "pagination", "cron", "outbox", "eventiq", "opentelemetry", "standard", "vectors", "vectors-sqlite"] [package.metadata.requires-dev] dev = [ + { name = "anyio", specifier = ">=4.14.1" }, { name = "bandit" }, { name = "cryptography" }, { name = "deptry" }, + { name = "eventiq", specifier = ">=1.1.14" }, { name = "fastapi", specifier = ">=0.138.0" }, { name = "greenlet" }, { name = "httpx", specifier = ">=0.28.1" }, @@ -2179,12 +2247,13 @@ dev = [ { name = "mkdocstrings", extras = ["python"] }, { name = "mypy" }, { name = "pytest" }, - { name = "pytest-asyncio" }, { name = "pytest-cov" }, { name = "pytest-repeat" }, { name = "pytest-sugar" }, { name = "pytest-timeout", specifier = ">=2.4.0" }, + { name = "pytest-xdist", specifier = ">=3.6" }, { name = "ruff" }, + { name = "sqliteai-vector", specifier = ">=1.0.0,<2" }, { name = "testcontainers", specifier = ">=4.15.0" }, { name = "watchdog", specifier = ">=2.0,<4.0" }, ] @@ -2198,6 +2267,7 @@ docs = [ ] e2e = [ { name = "cryptography" }, + { name = "sqliteai-vector", specifier = ">=1.0.0,<2" }, { name = "testcontainers", specifier = ">=4.15.0" }, ] lint = [ @@ -2207,13 +2277,26 @@ lint = [ { name = "ruff" }, ] test = [ + { name = "anyio", specifier = ">=4.14.1" }, { name = "httpx", specifier = ">=0.28.1" }, { name = "pytest" }, - { name = "pytest-asyncio" }, { name = "pytest-cov" }, { name = "pytest-repeat" }, { name = "pytest-sugar" }, { name = "pytest-timeout", specifier = ">=2.4.0" }, + { name = "pytest-xdist", specifier = ">=3.6" }, +] + +[[package]] +name = "sqliteai-vector" +version = "1.0.0" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/2b/40/77ea9053e9538d097d71ce944e7fe9f356648e5acf39adc3983b01be9287/sqliteai_vector-1.0.0-py3-none-macosx_10_9_x86_64.whl", hash = "sha256:e6724447b32e00342cab3ee67734c918d87ab2caa336a458f48dd4fe05830f9d", size = 144141, upload-time = "2026-05-25T14:22:10.029Z" }, + { url = "https://files.pythonhosted.org/packages/0f/6c/800d73d53b77d6dae6b11796ae8ce6bce6d5745b851900157b56a85a3e89/sqliteai_vector-1.0.0-py3-none-macosx_11_0_arm64.whl", hash = "sha256:5f52dede7372e67b9f219bd2163e4fd246eb867debee5f4ad825bab998a08812", size = 144136, upload-time = "2026-05-25T14:22:13.585Z" }, + { url = "https://files.pythonhosted.org/packages/62/1e/ac1bbd2f444b8b3fb3dd72b798cd52003c3ab2b80420049868a5280496af/sqliteai_vector-1.0.0-py3-none-manylinux2014_aarch64.whl", hash = "sha256:f9428101b02f633a7e1dde766edcf979827a97ec4ff113d3cd060739297e68fc", size = 67543, upload-time = "2026-05-25T14:22:11.683Z" }, + { url = "https://files.pythonhosted.org/packages/10/fa/c8c067b3a430351e828e686afd70a1ab736ddb73a01236d3e3ab3566fa5d/sqliteai_vector-1.0.0-py3-none-manylinux2014_x86_64.whl", hash = "sha256:108c56a3b11d36fd5b807b9d08169f10ba763cbabc569a7f0029c09d43ca066b", size = 83527, upload-time = "2026-05-25T14:22:11.548Z" }, + { url = "https://files.pythonhosted.org/packages/bb/00/ddce58f23164d65fd166f60e545120a7124ea38be33c23e67589251e7b28/sqliteai_vector-1.0.0-py3-none-win_amd64.whl", hash = "sha256:80eeddcbd7e41bd99c9e267de8aee6fdc04452d49a7168537bc7416059747900", size = 85343, upload-time = "2026-05-25T14:22:20.418Z" }, ] [[package]] @@ -2317,6 +2400,21 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/7b/61/cceae43728b7de99d9b847560c262873a1f6c98202171fd5ed62640b494b/tomli-2.4.1-py3-none-any.whl", hash = "sha256:0d85819802132122da43cb86656f8d1f8c6587d54ae7dcaf30e90533028b49fe", size = 14583, upload-time = "2026-03-25T20:22:03.012Z" }, ] +[[package]] +name = "typer" +version = "0.27.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "annotated-doc" }, + { name = "colorama", marker = "sys_platform == 'win32'" }, + { name = "rich" }, + { name = "shellingham" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/ae/40/4a3db7990d1f62a53182aa96eaef57aeb2886a27f90a195bc66713565d31/typer-0.27.1.tar.gz", hash = "sha256:a79bef8469a79c45498e7b814ecf8d603cc7644e9acbd9e19cac0334240b18df", size = 203994, upload-time = "2026-08-03T14:41:03.438Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/43/89/9518bc0c3929bee36b3a4a8e3daddd6e03f92f9961c66d4983b837160543/typer-0.27.1-py3-none-any.whl", hash = "sha256:53150287edd11baeb4e4722c8e394fcdf8181c0ae89485cba8d25c778d5edd56", size = 122874, upload-time = "2026-08-03T14:41:04.391Z" }, +] + [[package]] name = "typing-extensions" version = "4.15.0" From feb10eb299c45890fb591e27c92ea24d560ed97f Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Tue, 25 Aug 2026 00:01:01 +0200 Subject: [PATCH 03/11] fix(i18n): make TranslatedRepository generic over its model TranslatedRepository was declared as a bare SQLAlchemyRepository, so TranslatedRepository[Article] raised "not a generic class" and the model could only be set by hand -- unlike every other repository in the library. Add TranslatableBase and the TranslatableModel type var, mirroring the SoftDeleteBase / SoftDeleteModel pair, and bind the repository to them. A model carrying no translation table now raises TypeError on subclassing, the way SoftDeleteRepository validates its own. Co-Authored-By: Claude Opus 5 (1M context) --- sqlargon/i18n/__init__.py | 4 ++++ sqlargon/i18n/repository.py | 30 ++++++++++++++++++++++-------- sqlargon/i18n/translatable.py | 21 ++++++++++++++++++++- 3 files changed, 46 insertions(+), 9 deletions(-) diff --git a/sqlargon/i18n/__init__.py b/sqlargon/i18n/__init__.py index 816b540..35f11d0 100644 --- a/sqlargon/i18n/__init__.py +++ b/sqlargon/i18n/__init__.py @@ -2,7 +2,9 @@ from .mixin import TranslationMixin from .repository import TranslatedRepository from .translatable import ( + TranslatableBase, TranslatableMixin, + TranslatableModel, TranslationBase, current_translation, translation_class, @@ -20,7 +22,9 @@ __all__ = [ "LocaleMap", + "TranslatableBase", "TranslatableMixin", + "TranslatableModel", "TranslatedRepository", "TranslatedString", "Translation", diff --git a/sqlargon/i18n/repository.py b/sqlargon/i18n/repository.py index 79fec33..796f804 100644 --- a/sqlargon/i18n/repository.py +++ b/sqlargon/i18n/repository.py @@ -4,13 +4,15 @@ from sqlargon.repository import SQLAlchemyRepository +from .translatable import TranslatableMixin, TranslatableModel, current_translation + if TYPE_CHECKING: from typing import Any from typing_extensions import Self -class TranslatedRepository(SQLAlchemyRepository, abstract=True): +class TranslatedRepository(SQLAlchemyRepository[TranslatableModel], abstract=True): """Repository for a model whose fields are backed by a translation table. Every ``select()`` outer-joins the active-locale translation row, so @@ -19,11 +21,26 @@ class TranslatedRepository(SQLAlchemyRepository, abstract=True): resolve to the translation table's columns -- works without an explicit join in the calling code. - The model must use :class:`TranslatableMixin`, whose - ``_current_translation`` relationship carries the join condition that - matches the model's primary key and the active locale. + The model type variable is bound to + :class:`~sqlargon.i18n.TranslatableBase`, so a type checker rejects a + model that carries no translation table. At runtime the looser + :class:`~sqlargon.i18n.TranslatableMixin` is enough -- its + ``_current_translation`` relationship carries the join condition matching + the model's primary key and the active locale -- and anything else raises + ``TypeError`` on subclassing. """ + def __init_subclass__(cls, *, abstract: bool = False, **kwargs: Any) -> None: + super().__init_subclass__(abstract=abstract, **kwargs) + if abstract: + return + if not issubclass(cls.model, TranslatableMixin): + msg = ( + f"{cls.model.__name__} must inherit from TranslatableMixin " + f"to be used with {cls.__name__}" + ) + raise TypeError(msg) + def select( self, *args: Any, @@ -32,8 +49,5 @@ def select( return ( super() .select(*args, **kwargs) - .join( - self.model._current_translation, # noqa: SLF001 - isouter=True, # type: ignore[union-attr] - ) + .join(current_translation(self.model), isouter=True) ) diff --git a/sqlargon/i18n/translatable.py b/sqlargon/i18n/translatable.py index b6ce34b..abd0a6b 100644 --- a/sqlargon/i18n/translatable.py +++ b/sqlargon/i18n/translatable.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, ClassVar, cast +from typing import TYPE_CHECKING, Any, ClassVar, TypeVar, cast from weakref import WeakKeyDictionary import sqlalchemy as sa @@ -14,6 +14,8 @@ ) from sqlalchemy.orm.collections import attribute_keyed_dict +from sqlargon.orm import Base + from .expression import current_locale from .mixin import TranslationMixin from .translation import LocaleMap, Translation, as_translation @@ -225,6 +227,23 @@ def clear_translations(self, field: str) -> None: setattr(row, field, None) +class TranslatableBase(TranslatableMixin, Base): + """Declarative base for models whose fields live in a translation table. + + Inherit it instead of combining :class:`TranslatableMixin` with + :class:`~sqlargon.orm.Base` by hand, so + :class:`~sqlargon.i18n.TranslatedRepository` can type its model:: + + class Article(UUIDModelMixin, TranslatableBase): + __translated_fields__ = ("title",) + """ + + __abstract__ = True + + +TranslatableModel = TypeVar("TranslatableModel", bound=TranslatableBase) + + def _translated_property(field: str) -> hybrid_property[Translation | None]: def getter(self: TranslatableMixin) -> Translation | None: data = self.get_translations(field) From 4d19f4e400e341d75b26142a1603a9ddb1d5163a Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Tue, 25 Aug 2026 00:01:01 +0200 Subject: [PATCH 04/11] fix(i18n): unquote translated values on mysql and mariadb JSON_EXTRACT hands back a quoted JSON scalar there, so LIKE and ordering matched against the quotes as well. An equality comparison coerces its operand to JSON and agrees either way, which is what made the rest of the surface look correct. Found by the new e2e suite: every LIKE query over a TranslatedString column returned nothing on both backends. Co-Authored-By: Claude Opus 5 (1M context) --- sqlargon/i18n/expression.py | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/sqlargon/i18n/expression.py b/sqlargon/i18n/expression.py index 4032163..f0529b5 100644 --- a/sqlargon/i18n/expression.py +++ b/sqlargon/i18n/expression.py @@ -91,6 +91,23 @@ def _compile_postgresql( ) +@compiles(translated_value, "mysql") +def _compile_mysql( + element: translated_value, compiler: SQLCompiler, **kwargs: Any +) -> str: + """MySQL and MariaDB read a JSON scalar back quoted. + + ``JSON_EXTRACT`` yields ``"text"`` rather than ``text``, which ``LIKE`` and + ordering then match against including the quotes -- an equality comparison + coerces its operand to JSON and happens to agree, which is what makes the + rest of the surface look correct. Unquoting brings all of them back in line. + """ + return compiler.process( + sa.func.json_unquote(sa.func.json_extract(_operand(element), _locale_path())), + **kwargs, + ) + + @compiles(translated_value) def _compile_default( element: translated_value, compiler: SQLCompiler, **kwargs: Any From 619f3158f97486ba36a9ae860df5b973109e1604 Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Tue, 25 Aug 2026 00:01:01 +0200 Subject: [PATCH 05/11] fix(versioned): bump integer version columns with a SQL expression SQLAlchemy never leaves version_id_generator as True -- it normalises the default into a callable incrementing the version it is handed -- so the branch meant to emit `version + 1` was unreachable and the callable one ran instead, calling it with None and yielding 1 every time. An integer version column therefore stayed pinned at 1 across every update, silently voiding the optimistic concurrency guard: update_if_match could never tell a stale row from a fresh one. Key the decision off the column type instead, and skip the bump in bulk_update, where an executemany cannot carry an expression. Co-Authored-By: Claude Opus 5 (1M context) --- sqlargon/repository/versioned.py | 41 +++++++++++++++++++++++--------- 1 file changed, 30 insertions(+), 11 deletions(-) diff --git a/sqlargon/repository/versioned.py b/sqlargon/repository/versioned.py index e216a2a..c1dc0b9 100644 --- a/sqlargon/repository/versioned.py +++ b/sqlargon/repository/versioned.py @@ -3,7 +3,7 @@ from collections.abc import Mapping as MappingABC from typing import TYPE_CHECKING, Any -from sqlalchemy import Text, cast +from sqlalchemy import Integer, Text, cast from sqlargon.mixins import VersionedMixin from sqlargon.orm import VersionedModel @@ -49,6 +49,12 @@ class UserRepository(VersionedRepository[User]): ... if updated is None: # someone else modified the row first + A model bringing an integer ``version_id_col`` of its own is versioned + by a counter: the bump is the SQL expression ``version + 1``, which + ``bulk_update`` cannot carry -- an executemany binds one set of + parameters per row, and a SQL expression is the same for all of them -- + so a counter is left alone there. + The version column is left out of the default ``ON CONFLICT DO UPDATE`` set, so an upsert cannot silently clobber a version. Regular ``update_one`` / ``update_many`` / ``bulk_update`` all auto-increment @@ -105,6 +111,19 @@ def _version_generator(cls) -> Any: def _is_server_versioned(cls) -> bool: return cls._version_generator() is False + @classmethod + def _is_counter(cls) -> bool: + """Whether the version column is an integer counter. + + SQLAlchemy normalises the default ``version_id_generator`` into a + callable incrementing the version it is handed, which a statement + level update has no way to call -- it never loads the row, so it has + no current version to pass, and every call would yield ``1``. Such a + column is bumped with a SQL expression instead. + """ + col = cls._version_col() + return col is not None and isinstance(col.type, Integer) + @classmethod def _get_default_set(cls) -> set[str]: col = cls._version_col() @@ -114,19 +133,19 @@ def _get_default_set(cls) -> set[str]: def _with_version_increment(self, values: Values) -> Values: """Return ``values`` with the version column bumped.""" col = self._version_col() - generator = self._version_generator() - if col is None or generator is False: + if col is None or self._is_server_versioned(): return values name = col.name - if callable(generator): + if self._is_counter(): if isinstance(values, MappingABC): - return {**values, name: generator(None)} - return [{**row, name: generator(None)} for row in values] - # generator is True — integer increment via SQL expression + return {**values, name: col + 1} + # a SQL expression is the same for every row, which is not what an + # executemany binding one set of parameters per row can carry + return values + generator = self._version_generator() if isinstance(values, MappingABC): - return {**values, name: col + 1} - # integer increment is not supported for executemany - return values + return {**values, name: generator(None)} + return [{**row, name: generator(None)} for row in values] def _version_filter(self, expected: Any) -> Any: """The guard matching ``expected`` against the version column. @@ -155,7 +174,7 @@ async def bulk_update( ) -> None: col = self._version_col() generator = self._version_generator() - if col is not None and callable(generator): + if col is not None and not self._is_counter() and callable(generator): name = col.name values = [{**row, name: generator(None)} for row in values] await super().bulk_update(values, *args, on_=on_, **kwargs) From 27ae4d49cc54594bff5840c24fcdfba23ba206fc Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Tue, 25 Aug 2026 00:01:13 +0200 Subject: [PATCH 06/11] test: cover i18n, and close the residual coverage gaps i18n shipped with tests for the two locale slots alone: translatable.py sat at 35%, translation.py at 48%, mixin.py at 44%. Cover both backends -- the TranslatedString column and the translation table -- across the Translation value type, the pydantic hooks, the multi-locale helpers, the per-dialect compilation and the repository join, plus an e2e suite running all of it against the five real backends. Also cover the integer version counter and the server managed guard, the repository's database attribute, JSON literal rendering and the vector type's non-list paths, and mark the three dependency-absent ImportError guards no cover -- the extras are installed in this environment, so they were never reachable. Unit coverage 95% -> 99%. Co-Authored-By: Claude Opus 5 (1M context) --- sqlargon/database.py | 2 +- sqlargon/pagination/cursor.py | 2 +- sqlargon/types/vector.py | 2 +- tests/e2e/conftest.py | 46 +- tests/e2e/models.py | 45 ++ tests/e2e/test_i18n.py | 156 +++++++ tests/test_i18n.py | 813 ++++++++++++++++++++++++++++++++-- tests/test_registry.py | 20 + tests/test_types.py | 17 + tests/test_vectors.py | 23 + tests/test_versioned.py | 156 ++++++- 11 files changed, 1229 insertions(+), 53 deletions(-) create mode 100644 tests/e2e/test_i18n.py diff --git a/sqlargon/database.py b/sqlargon/database.py index f37fadf..c3d7159 100644 --- a/sqlargon/database.py +++ b/sqlargon/database.py @@ -19,7 +19,7 @@ try: from opentelemetry.instrumentation.sqlalchemy import SQLAlchemyInstrumentor -except ImportError: +except ImportError: # pragma: no cover - the extra is installed here SQLAlchemyInstrumentor = None if TYPE_CHECKING: diff --git a/sqlargon/pagination/cursor.py b/sqlargon/pagination/cursor.py index 2de368d..71acb9a 100644 --- a/sqlargon/pagination/cursor.py +++ b/sqlargon/pagination/cursor.py @@ -9,7 +9,7 @@ try: from sqlakeyset import unserialize_bookmark from sqlakeyset.paging import core_page_from_rows, prepare_paging -except ImportError as e: +except ImportError as e: # pragma: no cover - the extra is installed here msg = ( "Cursor pagination requires the 'sqlakeyset' package; " "install 'sqlargon[pagination]'" diff --git a/sqlargon/types/vector.py b/sqlargon/types/vector.py index c17bfe7..46cc7cf 100644 --- a/sqlargon/types/vector.py +++ b/sqlargon/types/vector.py @@ -14,7 +14,7 @@ try: from pgvector.sqlalchemy import VECTOR -except ImportError as e: +except ImportError as e: # pragma: no cover - the extra is installed here msg = "Vector columns require the 'pgvector' package; install 'sqlargon[vectors]'" raise ImportError(msg) from e diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 0f81a03..2927e30 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -17,6 +17,7 @@ from sqlalchemy.ext.asyncio import create_async_engine from sqlargon import Base, Database +from sqlargon.i18n import expression, translation from sqlargon.vectors import init_vectors from .backends import Backend, parse_backends @@ -30,6 +31,8 @@ AuditCommentRepository, AuditFollowRepository, DocumentRepository, + I18nArticleRepository, + I18nPostRepository, OutboxUserRepository, RawAuditArticleRepository, SoftUserRepository, @@ -42,7 +45,7 @@ ) if TYPE_CHECKING: - from collections.abc import AsyncGenerator, Generator + from collections.abc import AsyncGenerator, Callable, Generator def pytest_generate_tests(metafunc: pytest.Metafunc) -> None: @@ -143,6 +146,47 @@ async def db( await database.dispose() +#: The chain each locale walks; a locale absent from it falls back to itself. +LOCALE_CHAINS = {"en": ("en",), "pl": ("pl", "en"), "de": ("de", "en")} + + +@pytest.fixture +def locales() -> Generator[Callable[[str], None]]: + """Install a locale getter and fallback chain, and hand back the setter. + + Both are process-global slots, so whatever was in them is put back -- + xdist keeps a backend's whole test set on one worker, and a leaked getter + would follow the suite into the next test. + """ + active = "en" + + def switch(value: str) -> None: + nonlocal active + active = value + + previous_locale = expression._get_locale + previous_fallback = translation._get_fallback + expression.set_locale_getter(lambda: active) + translation.set_fallback_chain( + lambda requested: LOCALE_CHAINS.get(requested or active, (requested or active,)) + ) + try: + yield switch + finally: + expression._get_locale = previous_locale + translation._get_fallback = previous_fallback + + +@pytest.fixture +def i18n_posts() -> I18nPostRepository: + return I18nPostRepository() + + +@pytest.fixture +def i18n_articles() -> I18nArticleRepository: + return I18nArticleRepository() + + @pytest.fixture def needs_partial_upsert(backend: Backend) -> None: if not backend.partial_upsert: diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 137a571..2cda746 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -28,6 +28,14 @@ version_foreign_key, version_mapped_column, ) +from sqlargon.i18n import ( + TranslatableBase, + TranslatedRepository, + TranslatedString, + Translation, + TranslationMixin, + translation_table, +) from sqlargon.mixins import CreatedUpdatedMixin, UUIDModelMixin from sqlargon.outbox import OutboxConfig, OutboxEvent, OutboxRepository from sqlargon.types import GUID, JSON, GenerateUUID, GenerateUUIDV7, Timestamp, now @@ -288,6 +296,40 @@ class VectorDocRepository(HybridVectorRepository[VectorDoc]): default_order_by = VectorDoc.created_at +class I18nPost(UUIDModelMixin, TranslationMixin, Base): + """The JSON column backend: every locale lives in one column.""" + + __tablename__ = "e2e_i18n_post" + + title: Mapped[Translation] = mapped_column(TranslatedString(), nullable=True) + + +class I18nArticle(UUIDModelMixin, TranslatableBase): + """The translation table backend, with a plain column alongside.""" + + __tablename__ = "e2e_i18n_article" + __translated_fields__ = ("title", "body") + + slug: Mapped[str | None] = mapped_column(sa.Unicode(64), nullable=True) + + +class I18nArticleTranslation(translation_table(I18nArticle)): # type: ignore[misc] + __tablename__ = "e2e_i18n_article_translation" + + # declared NOT NULL on purpose: a locale row carries only the fields + # translated to that locale, so `translation_table` overrides it + title: Mapped[str] = mapped_column(sa.Unicode(64), nullable=False) + body: Mapped[str] = mapped_column(sa.UnicodeText(), nullable=False) + + +class I18nPostRepository(SQLAlchemyRepository[I18nPost]): + default_order_by = I18nPost.id + + +class I18nArticleRepository(TranslatedRepository[I18nArticle]): + default_order_by = I18nArticle.id + + def _tables(*models: type[Base]) -> tuple[sa.Table, ...]: return tuple(Base.metadata.tables[model.__tablename__] for model in models) @@ -307,6 +349,9 @@ def _tables(*models: type[Base]) -> tuple[sa.Table, ...]: AuditComment, AuditFollow, AuditArticle, + I18nPost, + I18nArticleTranslation, + I18nArticle, ) #: Tables the vector suite needs, which only a backend that can search vectors diff --git a/tests/e2e/test_i18n.py b/tests/e2e/test_i18n.py new file mode 100644 index 0000000..676a319 --- /dev/null +++ b/tests/e2e/test_i18n.py @@ -0,0 +1,156 @@ +"""E2E tests for both i18n backends across real database backends. + +The JSON column backend reads a locale out of a document, which each dialect +spells its own way -- the point of running it here rather than on SQLite alone. +""" + +from __future__ import annotations + +import pytest + +from .models import I18nArticle, I18nPost + +TITLES = {"en": "Hello", "pl": "Czesc"} + + +async def seed_posts(i18n_posts, *documents: dict[str, str]) -> None: + for document in documents: + await i18n_posts.create(title=document) + + +async def seed_articles(i18n_articles, *documents: dict[str, str]) -> None: + async with i18n_articles.session() as session: + for document in documents: + article = I18nArticle() + article.title = document + session.add(article) + await session.flush() + + +# --- the JSON column backend --- + + +@pytest.mark.usefixtures("locales") +async def test_json_backend_round_trips_every_locale(i18n_posts): + await seed_posts(i18n_posts, TITLES) + + stored = await i18n_posts.select().one() + + assert stored.title.data == TITLES + + +async def test_json_backend_reads_the_active_locale(i18n_posts, locales): + await seed_posts(i18n_posts, TITLES) + + locales("pl") + assert str((await i18n_posts.select().one()).title) == "Czesc" + + locales("en") + assert str((await i18n_posts.select().one()).title) == "Hello" + + +async def test_json_backend_falls_back_across_locales(i18n_posts, locales): + await seed_posts(i18n_posts, {"en": "Hello"}) + + locales("de") + + assert str((await i18n_posts.select().one()).title) == "Hello" + + +async def test_json_backend_filters_by_the_active_locale(i18n_posts, locales): + await seed_posts(i18n_posts, TITLES, {"en": "World", "pl": "Swiat"}) + + locales("pl") + + assert await i18n_posts.count(I18nPost.title == "Czesc") == 1 + assert await i18n_posts.count(I18nPost.title == "Hello") == 0 + + +async def test_json_backend_matches_with_like(i18n_posts, locales): + await seed_posts(i18n_posts, {"en": "Hello world", "pl": "Czesc swiecie"}) + + locales("pl") + + assert await i18n_posts.count(I18nPost.title.like("Czesc%")) == 1 + assert await i18n_posts.count(I18nPost.title.like("Hello%")) == 0 + + +async def test_json_backend_orders_by_the_active_locale(i18n_posts, locales): + await seed_posts( + i18n_posts, {"en": "Zulu", "pl": "Alfa"}, {"en": "Alpha", "pl": "Zeta"} + ) + + locales("pl") + ordered = await i18n_posts.select().order_by(I18nPost.title.asc()).all() + + assert [str(row.title) for row in ordered] == ["Alfa", "Zeta"] + + +# --- the translation table backend --- + + +@pytest.mark.usefixtures("locales") +async def test_translation_table_round_trips_every_locale(i18n_articles): + await seed_articles(i18n_articles, TITLES) + + stored = await i18n_articles.select().one() + + assert stored.get_translations("title") == TITLES + + +async def test_translation_table_reads_the_active_locale(i18n_articles, locales): + await seed_articles(i18n_articles, TITLES) + + locales("pl") + assert str((await i18n_articles.select().one()).title) == "Czesc" + + locales("en") + assert str((await i18n_articles.select().one()).title) == "Hello" + + +async def test_translation_table_leaves_an_untranslated_field_null( + i18n_articles, locales +): + async with i18n_articles.session() as session: + article = I18nArticle() + article.title = TITLES + article.body = {"en": "Body"} + session.add(article) + await session.flush() + + locales("en") + stored = await i18n_articles.select().one() + + assert stored.get_translations("body") == {"en": "Body"} + assert stored.get_translations("title") == TITLES + + +async def test_translation_table_filters_by_a_translated_field(i18n_articles, locales): + await seed_articles(i18n_articles, TITLES, {"en": "World", "pl": "Swiat"}) + + locales("pl") + matched = await i18n_articles.select().filter(I18nArticle.title == "Czesc").all() + + assert [str(row.title) for row in matched] == ["Czesc"] + + +async def test_translation_table_orders_by_a_translated_field(i18n_articles, locales): + await seed_articles( + i18n_articles, {"en": "Zulu", "pl": "Alfa"}, {"en": "Alpha", "pl": "Zeta"} + ) + + locales("pl") + ordered = await i18n_articles.select().order_by(I18nArticle.title.asc()).all() + + assert [str(row.title) for row in ordered] == ["Alfa", "Zeta"] + + +@pytest.mark.usefixtures("needs_foreign_keys") +async def test_translation_table_cascades_from_the_parent(i18n_articles, locales): + await seed_articles(i18n_articles, TITLES) + locales("en") + article = await i18n_articles.select().one() + + await i18n_articles.delete_one(I18nArticle.id == article.id) + + assert await i18n_articles.count() == 0 diff --git a/tests/test_i18n.py b/tests/test_i18n.py index 33b6bfe..d7c1c6a 100644 --- a/tests/test_i18n.py +++ b/tests/test_i18n.py @@ -1,64 +1,107 @@ -"""Unit tests for the locale getter / fallback-chain callable slots -and for ``TranslatedRepository``. +"""Unit tests for the two i18n backends, the locale slots and the repository. Mirrors ``test_outbox.py``: in-memory SQLite via the shared ``db`` fixture, -module-level models to survive ``--count=3`` re-registration. +module-level models to survive ``--count=3`` re-registration. Tests that need +models of their own build a throwaway declarative base -- see `local_base`. """ +from contextvars import ContextVar + import pytest import sqlalchemy as sa -from sqlalchemy.orm import relationship +from pydantic import BaseModel, TypeAdapter +from sqlalchemy.dialects import mysql, postgresql, sqlite +from sqlalchemy.orm import DeclarativeBase import sqlargon.i18n.expression as _expr import sqlargon.i18n.translation as _trans -from sqlargon import Base, Database +from sqlargon import Base, Database, SQLAlchemyRepository from sqlargon.i18n import ( + TranslatableBase, + TranslatableMixin, TranslatedRepository, + TranslatedString, + Translation, + TranslationMixin, + as_translation, + current_translation, fallback_chain, get_locale, + select_current, set_fallback_chain, set_locale_getter, + translation_class, + translation_table, ) +from sqlargon.i18n.translatable import LOCALE_LENGTH, _declarative_base -# ── models ────────────────────────────────────────────────────────────────── +locale: ContextVar[str] = ContextVar("locale", default="en") +#: The chain each locale walks. A locale absent from the map falls back to +#: itself alone, which is what `select_current` then has to give up on. +CHAINS = { + "en": ("en",), + "pl": ("pl", "en"), + "de": ("de", "en"), +} -class RepoModel(Base): - __tablename__ = "test_i18n_repo_model" - id = sa.Column(sa.Integer, primary_key=True) +def _chain(requested: str | None) -> tuple[str, ...]: + key = requested or locale.get() + return CHAINS.get(key, (key,)) -class RepoTranslation(Base): - __tablename__ = "test_i18n_repo_trans" +def _compile(expr, dialect) -> str: + """Compile an expression against ``dialect`` without executing it.""" + return str(expr.compile(dialect=dialect, compile_kwargs={"literal_binds": True})) - id = sa.Column( - sa.Integer, - sa.ForeignKey("test_i18n_repo_model.id"), - primary_key=True, - ) +# --- models --- -RepoModel._current_translation = relationship( - RepoTranslation, - primaryjoin=RepoModel.id == RepoTranslation.id, - uselist=False, - viewonly=True, - lazy="raise", -) +class Post(TranslationMixin, Base): + """The JSON column backend: every locale lives in one ``TranslatedString``.""" + + __tablename__ = "test_i18n_post" + + id = sa.Column(sa.Integer, primary_key=True, autoincrement=True) + title = sa.Column(TranslatedString(), nullable=True) + + +class Article(TranslatableBase): + """The translation table backend, with a plain column alongside.""" + + __tablename__ = "test_i18n_article" + __translated_fields__ = ("title", "body") + + id = sa.Column(sa.Integer, primary_key=True, autoincrement=True) + slug = sa.Column(sa.Unicode(255), nullable=True) + + +class ArticleTranslation(translation_table(Article)): # type: ignore[misc] + __tablename__ = "test_i18n_article_translation" + + # declared NOT NULL on purpose: `_allow_untranslated_fields` overrides it, + # because a locale row carries only the fields translated to that locale + title = sa.Column(sa.Unicode(255), nullable=False) + body = sa.Column(sa.UnicodeText(), nullable=False) + + +class Untranslated(Base): + __tablename__ = "test_i18n_untranslated" + + id = sa.Column(sa.Integer, primary_key=True, autoincrement=True) + + +class PostRepository(SQLAlchemyRepository[Post]): + pass -class TestRepo(TranslatedRepository): - """Concrete repository on a model whose ``_current_translation`` is a - plain relationship -- not one created by ``TranslatableMixin`` -- which - is enough to verify the join behaviour. - """ - __test__ = False - model = RepoModel +class ArticleRepository(TranslatedRepository[Article]): + pass -# ── fixtures ──────────────────────────────────────────────────────────────── +# --- fixtures --- @pytest.fixture(autouse=True) @@ -66,11 +109,25 @@ def _reset_locale_slots(): """Reset the locale and fallback callable slots before every test.""" _expr._get_locale = None _trans._get_fallback = None + locale.set("en") + + +@pytest.fixture +def locales(): + """Configure the slots and hand back the setter switching locale.""" + set_locale_getter(locale.get) + set_fallback_chain(_chain) + return locale.set @pytest.fixture(autouse=True) -async def _create_tables(db: Database): - created = (RepoModel.__table__, RepoTranslation.__table__) +async def tables(db: Database): + created = ( + Post.__table__, + Article.__table__, + ArticleTranslation.__table__, + Untranslated.__table__, + ) async with db.engine.begin() as conn: for table in created: await conn.run_sync(table.create, checkfirst=True) @@ -80,7 +137,30 @@ async def _create_tables(db: Database): await conn.run_sync(table.drop, checkfirst=True) -# --- locale getter ----------------------------------------------------------- +@pytest.fixture +def posts(): + return PostRepository() + + +@pytest.fixture +def articles(): + return ArticleRepository() + + +@pytest.fixture +def local_base(): + """A throwaway declarative base, so a test can declare models of its own + without colliding with the shared metadata under ``--count=3``. + """ + + class LocalBase(DeclarativeBase): + pass + + yield LocalBase + LocalBase.registry.dispose() + + +# --- locale getter --- def test_locale_getter_raises_before_configuration(): @@ -110,7 +190,7 @@ def test_locale_getter_slot_is_cleared_between_tests(): assert get_locale() == "fr" -# --- fallback chain ---------------------------------------------------------- +# --- fallback chain --- def test_fallback_chain_raises_before_configuration(): @@ -167,24 +247,663 @@ def test_fallback_chain_slot_is_cleared_between_tests(): assert fallback_chain() == ("fr", "en") -# --- TranslatedRepository ---------------------------------------------------- +# --- select_current --- -def test_translated_repository_select_includes_the_outer_join(): - repo = TestRepo() - stmt = repo.select().query +@pytest.mark.parametrize( + ("active", "data", "expected"), + [ + ("en", {"en": "Hello", "pl": "Czesc"}, "Hello"), + ("pl", {"en": "Hello", "pl": "Czesc"}, "Czesc"), + # "de" is missing, so its chain walks down to "en" + ("de", {"en": "Hello", "pl": "Czesc"}, "Hello"), + # nothing in the chain matches: the first known translation is taken + ("en", {"pl": "Czesc"}, "Czesc"), + ], +) +@pytest.mark.usefixtures("locales") +def test_select_current_walks_the_chain(active, data, expected): + locale.set(active) + + assert select_current(data) == expected + + +@pytest.mark.usefixtures("locales") +def test_select_current_honours_an_explicit_locale(): + locale.set("en") + + assert select_current({"en": "Hello", "pl": "Czesc"}, "pl") == "Czesc" + + +@pytest.mark.usefixtures("locales") +def test_select_current_of_an_empty_mapping_is_none(): + assert select_current({}) is None + + +# --- Translation --- + + +@pytest.mark.usefixtures("locales") +def test_translation_is_the_active_locale_text(): + translation = Translation("Hello", {"en": "Hello", "pl": "Czesc"}) + + assert translation == "Hello" + assert isinstance(translation, str) + + +@pytest.mark.usefixtures("locales") +def test_translation_without_data_keys_the_active_locale(): + locale.set("pl") + + assert Translation("Czesc").data == {"pl": "Czesc"} + + +@pytest.mark.usefixtures("locales") +def test_translation_data_is_a_copy(): + translation = Translation("Hello", {"en": "Hello"}) + + translation.data["pl"] = "Czesc" + + assert translation.data == {"en": "Hello"} + + +@pytest.mark.usefixtures("locales") +def test_translation_get_returns_none_for_an_unknown_locale(): + translation = Translation("Hello", {"en": "Hello"}) + + assert translation.get("en") == "Hello" + assert translation.get("pl") is None + + +@pytest.mark.usefixtures("locales") +def test_translation_update_returns_a_new_instance(): + original = Translation("Hello", {"en": "Hello"}) + + updated = original.update("Czesc", "pl") + + assert updated.data == {"en": "Hello", "pl": "Czesc"} + assert original.data == {"en": "Hello"} + + +@pytest.mark.usefixtures("locales") +def test_translation_update_defaults_to_the_active_locale(): + locale.set("pl") + + updated = Translation("Hello", {"en": "Hello"}).update("Czesc") + + assert updated.data == {"en": "Hello", "pl": "Czesc"} + assert updated == "Czesc" + + +@pytest.mark.usefixtures("locales") +def test_translation_update_of_an_unreachable_locale_falls_back_to_the_value(): + """No chain reaches "es", so the new text is what the copy reads as.""" + updated = Translation("Hello", {}).update("Hola", "es") + + assert updated == "Hola" + + +@pytest.mark.usefixtures("locales") +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("Hello", {"en": "Hello"}), + ({"en": "Hello", "pl": "Czesc"}, {"en": "Hello", "pl": "Czesc"}), + ], +) +def test_translation_validate_accepts_text_and_mappings(value, expected): + assert Translation._validate(value).data == expected + + +@pytest.mark.usefixtures("locales") +def test_translation_validate_rejects_none(): + with pytest.raises(ValueError, match="string or a mapping"): + Translation._validate(None) + + +# --- pydantic integration --- + + +@pytest.mark.usefixtures("locales") +def test_translation_validates_from_a_string(): + result = TypeAdapter(Translation).validate_python("Hello") + + assert isinstance(result, Translation) + assert result.data == {"en": "Hello"} + + +@pytest.mark.usefixtures("locales") +def test_translation_validates_from_a_mapping(): + locale.set("pl") + + result = TypeAdapter(Translation).validate_python({"en": "Hello", "pl": "Czesc"}) + + assert result == "Czesc" + assert result.data == {"en": "Hello", "pl": "Czesc"} + + +@pytest.mark.usefixtures("locales") +def test_translation_rejects_an_unusable_type_as_a_validation_error(): + """A `ValueError` is what pydantic turns into a validation error.""" + + class Model(BaseModel): + title: Translation + + with pytest.raises(ValueError, match="Cannot build a Translation from int"): + Model(title=1) + + +@pytest.mark.usefixtures("locales") +def test_translation_serializes_to_the_active_locale_text(): + locale.set("pl") + translation = Translation("Czesc", {"en": "Hello", "pl": "Czesc"}) + + adapter = TypeAdapter(Translation) + + assert adapter.dump_python(translation) == "Czesc" + assert adapter.dump_json(translation) == b'"Czesc"' + + +def test_translation_json_schema_in_validation_mode_accepts_both_shapes(): + schema = TypeAdapter(Translation).json_schema() + + assert "anyOf" in schema + assert {"type": "string"} in schema["anyOf"] + + +def test_translation_json_schema_in_serialization_mode_is_a_string(): + schema = TypeAdapter(Translation).json_schema(mode="serialization") + + assert schema == {"type": "string"} + + +# --- as_translation --- + + +@pytest.mark.usefixtures("locales") +def test_as_translation_passes_none_through(): + assert as_translation(None) is None + + +@pytest.mark.usefixtures("locales") +def test_as_translation_passes_a_translation_through_unchanged(): + translation = Translation("Hello", {"en": "Hello"}) + + assert as_translation(translation) is translation + + +@pytest.mark.usefixtures("locales") +def test_as_translation_keys_a_string_by_the_active_locale(): + locale.set("pl") + + assert as_translation("Czesc").data == {"pl": "Czesc"} + + +@pytest.mark.usefixtures("locales") +def test_as_translation_honours_an_explicit_locale(): + assert as_translation("Czesc", "pl").data == {"pl": "Czesc"} + + +@pytest.mark.usefixtures("locales") +def test_as_translation_of_a_mapping_selects_the_current_text(): + locale.set("pl") + + translation = as_translation({"en": "Hello", "pl": "Czesc"}) + + assert translation == "Czesc" + + +@pytest.mark.usefixtures("locales") +def test_as_translation_of_an_empty_mapping_is_empty_text(): + assert as_translation({}) == "" + + +@pytest.mark.usefixtures("locales") +def test_as_translation_stringifies_mapping_keys(): + assert as_translation({1: "One"}).data == {"1": "One"} + + +@pytest.mark.usefixtures("locales") +def test_as_translation_rejects_a_non_string_text(): + with pytest.raises(ValueError, match="locale 'en' must be a string, got int"): + as_translation({"en": 1}) + + +@pytest.mark.usefixtures("locales") +def test_as_translation_rejects_an_unsupported_type(): + with pytest.raises(ValueError, match="Cannot build a Translation from float"): + as_translation(1.5) + + +# --- TranslationMixin --- + + +@pytest.mark.usefixtures("locales") +def test_get_translations_returns_every_locale(): + post = Post(title={"en": "Hello", "pl": "Czesc"}) + + assert post.get_translations("title") == {"en": "Hello", "pl": "Czesc"} + + +@pytest.mark.usefixtures("locales") +def test_get_translations_of_an_unset_field_is_empty(): + assert Post().get_translations("title") == {} + + +@pytest.mark.usefixtures("locales") +def test_get_translation_of_the_active_locale_walks_the_chain(): + locale.set("de") + post = Post(title={"en": "Hello"}) + + assert post.get_translation("title") == "Hello" + + +@pytest.mark.usefixtures("locales") +def test_get_translation_of_an_explicit_locale_does_not_walk_the_chain(): + """An explicit locale is looked up as given -- "de" is simply missing.""" + post = Post(title={"en": "Hello"}) + + assert post.get_translation("title", "de") is None + + +@pytest.mark.usefixtures("locales") +def test_set_translation_merges_into_the_existing_locales(): + post = Post(title={"en": "Hello"}) + + post.set_translation("title", "Czesc", "pl") + + assert post.get_translations("title") == {"en": "Hello", "pl": "Czesc"} + + +@pytest.mark.usefixtures("locales") +def test_set_translation_defaults_to_the_active_locale(): + locale.set("pl") + post = Post(title={"en": "Hello"}) + + post.set_translation("title", "Czesc") + + assert post.get_translation("title", "pl") == "Czesc" - compiled = str(stmt.compile(compile_kwargs={"literal_binds": True})) - assert "LEFT OUTER JOIN" in compiled - assert "test_i18n_repo_trans" in compiled +@pytest.mark.usefixtures("locales") +def test_clear_translations_leaves_an_empty_translation(): + """The column keeps holding a translation, so it may be non-nullable.""" + post = Post(title={"en": "Hello"}) + + post.clear_translations("title") + + assert post.title == "" + assert post.get_translations("title") == {} + + +# --- TranslatedString --- + + +@pytest.mark.usefixtures("locales") +@pytest.mark.parametrize( + ("value", "expected"), + [ + (None, None), + ("Hello", {"en": "Hello"}), + ({"en": "Hello", "pl": "Czesc"}, {"en": "Hello", "pl": "Czesc"}), + ], +) +def test_translated_string_binds_the_whole_locale_map(value, expected): + assert TranslatedString().process_bind_param(value, sqlite.dialect()) == expected + + +@pytest.mark.usefixtures("locales") +def test_translated_string_reads_back_as_a_translation(): + result = TranslatedString().process_result_value( + {"en": "Hello", "pl": "Czesc"}, sqlite.dialect() + ) + + assert isinstance(result, Translation) + assert result == "Hello" + assert result.data == {"en": "Hello", "pl": "Czesc"} + + +def test_translated_string_reads_none_as_none(): + assert TranslatedString().process_result_value(None, sqlite.dialect()) is None + + +@pytest.mark.usefixtures("locales") +def test_translated_string_skips_non_string_values_on_read(): + result = TranslatedString().process_result_value( + {"en": "Hello", "count": 3}, sqlite.dialect() + ) + + assert result.data == {"en": "Hello"} + + +@pytest.mark.usefixtures("locales") +def test_translated_string_compares_whole_locale_maps(): + """Two values reading the same in the active locale still differ.""" + kind = TranslatedString() + one = {"en": "Hello", "pl": "Czesc"} + other = {"en": "Hello"} + + assert kind.compare_values(one, dict(one)) + assert not kind.compare_values(one, other) + + +@pytest.mark.usefixtures("locales") +def test_translated_string_compares_none(): + assert TranslatedString().compare_values(None, None) + + +@pytest.mark.usefixtures("locales") +def test_translated_string_filter_targets_the_active_locale_on_postgresql(): + sql = _compile(sa.select(Post).where(Post.title == "Hello"), postgresql.dialect()) + + assert "->>" in sql + + +@pytest.mark.usefixtures("locales") +def test_translated_string_filter_falls_back_to_json_extract(): + sql = _compile(sa.select(Post).where(Post.title == "Hello"), sqlite.dialect()) + + assert "json_extract(test_i18n_post.title, '$.\"en\"')" in sql + + +@pytest.mark.usefixtures("locales") +def test_translated_string_filter_unquotes_on_mysql(): + """``JSON_EXTRACT`` yields a quoted scalar there, which ``LIKE`` would + match against including the quotes. + """ + sql = _compile(sa.select(Post).where(Post.title.like("Hello%")), mysql.dialect()) + + assert "json_unquote(json_extract(test_i18n_post.title, '$.\"en\"'))" in sql + + +@pytest.mark.usefixtures("locales") +def test_translated_string_ordering_targets_the_active_locale(): + ascending = _compile(sa.select(Post).order_by(Post.title.asc()), sqlite.dialect()) + descending = _compile(sa.select(Post).order_by(Post.title.desc()), sqlite.dialect()) + ordered = "json_extract(test_i18n_post.title, '$.\"en\"')" + + assert ascending.endswith(f"ORDER BY {ordered} ASC") + assert descending.endswith(f"ORDER BY {ordered} DESC") + + +@pytest.mark.usefixtures("locales") +def test_translated_string_reverse_operations_target_the_active_locale(): + """Taken off the table: an ORM attribute routes its reflected operators + through SQLAlchemy's own ``ColumnProperty.Comparator`` instead. + """ + column = Post.__table__.c.title + + sql = _compile(sa.select("greeting" + column), sqlite.dialect()) + + assert sql.startswith( + "SELECT 'greeting' || json_extract(test_i18n_post.title, '$.\"en\"')" + ) + + +@pytest.mark.usefixtures("locales") +def test_translated_string_json_operators_still_address_the_locale_map(): + """The inherited JSON methods key by locale name, not by locale text.""" + sql = _compile(sa.select(Post).where(Post.title.has_key("pl")), sqlite.dialect()) + + assert '$."pl"' in sql + assert '$."en"' not in sql + + +async def test_translated_string_round_trips_through_the_database(posts, locales): + await posts.create(title={"en": "Hello", "pl": "Czesc"}) + + locales("pl") + stored = await posts.select().one() + + assert stored.title == "Czesc" + assert stored.title.data == {"en": "Hello", "pl": "Czesc"} + + +async def test_translated_string_filters_by_the_active_locale(posts, locales): + await posts.create(title={"en": "Hello", "pl": "Czesc"}) + await posts.create(title={"en": "World", "pl": "Swiat"}) + + locales("pl") + matched = await posts.select().filter(Post.title == "Czesc").all() + + assert [row.title for row in matched] == ["Czesc"] + assert await posts.count(Post.title == "Hello") == 0 + + +async def test_translated_string_orders_by_the_active_locale(posts, locales): + await posts.create(title={"en": "Zulu", "pl": "Alfa"}) + await posts.create(title={"en": "Alpha", "pl": "Zeta"}) + + locales("pl") + ordered = await posts.select().order_by(Post.title.asc()).all() + + assert [str(row.title) for row in ordered] == ["Alfa", "Zeta"] + + +@pytest.mark.usefixtures("locales") +async def test_translated_string_matches_with_like(posts): + await posts.create(title={"en": "Hello world"}) + + matched = await posts.select().filter(Post.title.like("Hello%")).all() + + assert len(matched) == 1 + + +# --- translation_table --- + + +def test_translation_table_keys_on_the_parent_key_and_the_locale(): + table = ArticleTranslation.__table__ + + assert set(table.primary_key.columns.keys()) == {"id", "locale"} + + +def test_translation_table_cascades_from_the_parent(): + (foreign_key,) = ArticleTranslation.__table__.c.id.foreign_keys + + assert foreign_key.column is Article.__table__.c.id + assert foreign_key.ondelete == "CASCADE" + + +def test_translation_table_sizes_the_locale_column(): + assert ArticleTranslation.__table__.c.locale.type.length == LOCALE_LENGTH + + +def test_translation_table_makes_translated_columns_nullable(): + """Declared NOT NULL, but a locale row holds only what it translates.""" + assert ArticleTranslation.__table__.c.title.nullable + assert ArticleTranslation.__table__.c.body.nullable + + +def test_translation_table_rejects_a_field_it_has_no_column_for(local_base): + class Doc(TranslatableMixin, local_base): + __tablename__ = "doc" + __translated_fields__ = ("headline",) + + id = sa.Column(sa.Integer, primary_key=True) + + with pytest.raises(TypeError, match="has no column 'headline'"): + + class DocTranslation(translation_table(Doc)): + __tablename__ = "doc_translation" + + body = sa.Column(sa.Unicode(50)) + + +def test_translation_table_rejects_a_non_declarative_parent(): + with pytest.raises(TypeError, match="is not a declarative model"): + _declarative_base(object) + + +def test_translation_class_resolves_the_registered_model(): + assert translation_class(Article) is ArticleTranslation + + +def test_translation_class_resolves_through_the_mro(local_base): + class Doc(TranslatableMixin, local_base): + __tablename__ = "doc" + __translated_fields__ = ("title",) + + id = sa.Column(sa.Integer, primary_key=True) + + class DocTranslation(translation_table(Doc)): + __tablename__ = "doc_translation" + + title = sa.Column(sa.Unicode(50)) + + class Guide(Doc): + """A subclass inherits the translation table its parent declared.""" + + assert translation_class(Guide) is DocTranslation + + +def test_translation_class_raises_for_a_model_without_one(): + with pytest.raises(LookupError, match="No translation table declared"): + translation_class(Untranslated) + + +def test_current_translation_is_the_join_target_relationship(): + assert current_translation(Article) is Article._current_translation + + +# --- TranslatableMixin --- + + +@pytest.mark.usefixtures("locales") +async def test_translatable_writes_and_reads_every_locale(articles): + async with articles.session() as session: + article = Article() + article.title = {"en": "Hello", "pl": "Czesc"} + session.add(article) + await session.flush() + + assert article.get_translations("title") == {"en": "Hello", "pl": "Czesc"} + locale.set("pl") + assert article.title == "Czesc" + + +@pytest.mark.usefixtures("locales") +async def test_translatable_leaves_an_untranslated_field_null(articles): + async with articles.session() as session: + article = Article() + article.title = {"en": "Hello", "pl": "Czesc"} + article.body = {"en": "Body"} + session.add(article) + await session.flush() + + assert article._translations["pl"].body is None + assert article.get_translations("body") == {"en": "Body"} + + +@pytest.mark.usefixtures("locales") +async def test_translatable_assigning_a_string_only_touches_the_active_locale(articles): + async with articles.session() as session: + article = Article() + article.title = {"en": "Hello", "pl": "Czesc"} + session.add(article) + await session.flush() + + locale.set("pl") + article.title = "Witaj" + + assert article.get_translations("title") == {"en": "Hello", "pl": "Witaj"} + + +@pytest.mark.usefixtures("locales") +async def test_translatable_assigning_none_clears_every_locale(articles): + async with articles.session() as session: + article = Article() + article.title = {"en": "Hello", "pl": "Czesc"} + session.add(article) + await session.flush() + + article.title = None + + assert article.get_translations("title") == {} + assert article.title is None + + +@pytest.mark.usefixtures("locales") +async def test_translatable_falls_back_across_locales(articles): + async with articles.session() as session: + article = Article() + article.title = {"en": "Hello"} + session.add(article) + await session.flush() + + locale.set("de") + + assert article.title == "Hello" + + +@pytest.mark.usefixtures("locales") +def test_translatable_delegates_an_untranslated_field_to_the_mixin(): + article = Article(slug="hello") + + assert article.get_translations("slug") == {"en": "hello"} + + +# --- TranslatedRepository --- + + +def test_translated_repository_resolves_the_model_from_the_subscript(): + assert ArticleRepository.model is Article + + +def test_translated_repository_select_includes_the_outer_join(): + sql = str(ArticleRepository().select().query) + + assert "LEFT OUTER JOIN" in sql + assert ArticleTranslation.__tablename__ in sql def test_translated_repository_select_accepts_column_args(): - repo = TestRepo() - stmt = repo.select(RepoModel.id).query + sql = str(ArticleRepository().select(Article.id).query) + + assert "LEFT OUTER JOIN" in sql + assert Article.__tablename__ in sql + + +def test_translated_repository_rejects_a_model_without_the_mixin(): + with pytest.raises(TypeError, match="must inherit from TranslatableMixin"): + + class Repository(TranslatedRepository[Untranslated]): + pass + + +def test_translated_repository_allows_an_abstract_subclass(): + class Abstract(TranslatedRepository, abstract=True): + """An abstract subclass declares no model of its own.""" + + assert not hasattr(Abstract, "model") + + +@pytest.mark.usefixtures("locales") +async def test_translated_repository_filters_by_a_translated_field(articles): + async with articles.session() as session: + for titles in ({"en": "Hello", "pl": "Czesc"}, {"en": "World", "pl": "Swiat"}): + article = Article() + article.title = titles + session.add(article) + await session.flush() + + locale.set("pl") + matched = await articles.select().filter(Article.title == "Czesc").all() + + assert len(matched) == 1 + assert matched[0].title == "Czesc" + + +@pytest.mark.usefixtures("locales") +async def test_translated_repository_orders_by_a_translated_field(articles): + async with articles.session() as session: + for titles in ({"en": "Zulu", "pl": "Alfa"}, {"en": "Alpha", "pl": "Zeta"}): + article = Article() + article.title = titles + session.add(article) + await session.flush() - compiled = str(stmt.compile(compile_kwargs={"literal_binds": True})) + locale.set("pl") + ordered = await articles.select().order_by(Article.title.asc()).all() - assert "LEFT OUTER JOIN" in compiled - assert "test_i18n_repo_model" in compiled + assert [str(row.title) for row in ordered] == ["Alfa", "Zeta"] diff --git a/tests/test_registry.py b/tests/test_registry.py index 215861c..10d93c8 100644 --- a/tests/test_registry.py +++ b/tests/test_registry.py @@ -88,6 +88,26 @@ class Repository(SQLAlchemyRepository[user_model]): await db.drop_all() +def test_repository_database_attribute_wins_over_default(db, user_model): + set_default_database(db) + bound = Database(MEMORY_URL) + + class Repository(SQLAlchemyRepository[user_model]): + database = bound + + assert Repository().db is bound + + +def test_repository_using_db_wins_over_the_database_attribute(user_model): + bound = Database(MEMORY_URL) + other = Database(MEMORY_URL) + + class Repository(SQLAlchemyRepository[user_model]): + database = bound + + assert Repository().using(db=other).db is other + + async def test_repository_using_db_wins_over_default(db, user_model): set_default_database(db) other = Database(MEMORY_URL) diff --git a/tests/test_types.py b/tests/test_types.py index 4fcbf0e..c99295b 100644 --- a/tests/test_types.py +++ b/tests/test_types.py @@ -257,6 +257,23 @@ def test_json_dialect_impl_fallback(dialect): assert impl.none_as_null is True +# --- JSON literal rendering --- + + +@pytest.mark.parametrize( + "dialect", [postgresql.dialect(), sqlite.dialect(), mysql.dialect()] +) +def test_json_literal_processor_renders_null(dialect): + assert JSON().literal_processor(dialect)(None) == "NULL" + + +@pytest.mark.parametrize( + "dialect", [postgresql.dialect(), sqlite.dialect(), mysql.dialect()] +) +def test_json_literal_processor_serializes_and_quotes_a_document(dialect): + assert JSON().literal_processor(dialect)({"a": 1}) == "'{\"a\":1}'" + + # --- JSON per-dialect compilation --- _json_col = sa.column("data", JSON()) diff --git a/tests/test_vectors.py b/tests/test_vectors.py index 731096d..c03588c 100644 --- a/tests/test_vectors.py +++ b/tests/test_vectors.py @@ -170,6 +170,15 @@ def test_vector_normalises_non_list_results(): ) +def test_vector_normalises_a_non_list_bind(): + assert Vector(3).process_bind_param((1, 2, 3), postgresql.dialect()) == VECTOR + + +def test_vector_postgresql_result_passes_the_list_through(): + stored = [1.0, 2.0, 3.0] + assert Vector(3).process_result_value(stored, postgresql.dialect()) is stored + + @pytest.mark.parametrize("dialect", [postgresql.dialect(), sqlite.dialect()]) def test_vector_none_stays_none(dialect): vector = Vector(3) @@ -214,6 +223,20 @@ def test_comparator_exposes_the_distance_methods(): assert "<+>" in _compile(Note.embedding.l1_distance(VECTOR), postgresql.dialect()) +@pytest.mark.parametrize( + ("metric", "operator"), + [ + (DistanceMetric.COSINE, "<=>"), + (DistanceMetric.L2, "<->"), + (DistanceMetric.DOT, "<#>"), + (DistanceMetric.L1, "<+>"), + ], +) +def test_comparator_distance_selects_the_metric(metric, operator): + expression = Note.embedding.distance(VECTOR, metric) + assert operator in _compile(expression, postgresql.dialect()) + + def test_metric_maps_to_pgvector_names(): assert DistanceMetric.COSINE.pg_opclass == "vector_cosine_ops" assert DistanceMetric.L2.pg_opclass == "vector_l2_ops" diff --git a/tests/test_versioned.py b/tests/test_versioned.py index 7d7ccf6..ff15f1f 100644 --- a/tests/test_versioned.py +++ b/tests/test_versioned.py @@ -4,10 +4,12 @@ fixture, module-level models to survive ``--count=3`` re-registration. """ +from typing import Any from uuid import UUID import pytest import sqlalchemy as sa +from sqlalchemy.orm import declared_attr from sqlargon import ( Base, @@ -18,6 +20,7 @@ VersionedBase, VersionedMixin, VersionedRepository, + XminVersionedBase, ) from sqlargon.mixins import UUIDModelMixin from sqlargon.types import GUID, GenerateUUID @@ -50,6 +53,29 @@ class MarkerOnly(VersionedMixin, Base): id = sa.Column(sa.Integer, primary_key=True) +class CounterVersioned(VersionedMixin, Base): + """An integer version column, the counter SQLAlchemy versions by default.""" + + __tablename__ = "test_versioned_counter" + + id = sa.Column(sa.Integer, primary_key=True, autoincrement=True) + name = sa.Column(sa.Unicode(255), nullable=True) + version_id = sa.Column(sa.Integer, nullable=False, default=1) + + @declared_attr.directive + def __mapper_args__(cls) -> dict[str, Any]: + return {"eager_defaults": True, "version_id_col": cls.version_id} + + +class ServerVersioned(XminVersionedBase): + """A server managed version, which the repository must never set itself.""" + + __tablename__ = "test_versioned_server" + + id = sa.Column(sa.Integer, primary_key=True) + name = sa.Column(sa.Unicode(255), nullable=True) + + class VersionedArticleRepository(VersionedRepository[VersionedArticle]): default_order_by = VersionedArticle.id @@ -64,16 +90,32 @@ class HandVersionedRepository(VersionedRepository[HandVersioned]): # type: igno pass +class CounterRepository(VersionedRepository[CounterVersioned]): # type: ignore[type-var] + default_order_by = CounterVersioned.id + + +class ServerVersionedRepository(VersionedRepository[ServerVersioned]): + pass + + # --- fixtures --- @pytest.fixture(autouse=True) async def tables(db: Database): + created = (VersionedArticle.__table__, CounterVersioned.__table__) async with db.engine.begin() as conn: - await conn.run_sync(VersionedArticle.__table__.create, checkfirst=True) + for table in created: + await conn.run_sync(table.create, checkfirst=True) yield async with db.engine.begin() as conn: - await conn.run_sync(VersionedArticle.__table__.drop, checkfirst=True) + for table in reversed(created): + await conn.run_sync(table.drop, checkfirst=True) + + +@pytest.fixture +def counters(): + return CounterRepository() @pytest.fixture @@ -384,3 +426,113 @@ async def test_update_one_with_manual_version_filter_works_on_match( assert updated is not None assert updated.name == "jane" + + +# --- integer counter versions --- + + +def test_counter_version_is_recognised(): + assert CounterRepository._is_counter() + assert not VersionedArticleRepository._is_counter() + + +def test_counter_update_bumps_with_a_sql_expression(counters): + """A counter has no fresh value to bind: it is bumped relative to the row.""" + query = counters.update({"name": "jane"}).query + + assert "version_id + " in str(query.compile(compile_kwargs={"literal_binds": True})) + + +async def test_counter_version_starts_at_one(counters): + row = await counters.create(name="john") + + assert row.version_id == 1 + + +async def test_counter_version_increments_on_every_update(counters): + row = await counters.create(name="john") + + for expected in (2, 3, 4): + row = await counters.update_one({"name": "jane"}, CounterVersioned.id == row.id) + assert row.version_id == expected + + +async def test_counter_update_if_match_rejects_a_stale_version(counters): + row = await counters.create(name="john") + await counters.update_one({"name": "jane"}, CounterVersioned.id == row.id) + + stale = await counters.update_if_match( + {"name": "joan"}, CounterVersioned.id == row.id, expected_version=1 + ) + + assert stale is None + + +async def test_counter_update_if_match_accepts_the_current_version(counters): + row = await counters.create(name="john") + + updated = await counters.update_if_match( + {"name": "jane"}, + CounterVersioned.id == row.id, + expected_version=row.version_id, + ) + + assert updated is not None + assert updated.version_id == 2 + + +async def test_counter_bulk_update_leaves_the_version_alone(counters): + """An executemany binds one set of parameters per row, which a SQL + expression -- the same for every row -- is not. + """ + first = await counters.create(name="a") + second = await counters.create(name="b") + + await counters.bulk_update( + [{"id": first.id, "name": "a2"}, {"id": second.id, "name": "b2"}] + ) + + rows = await counters.select().order_by(CounterVersioned.id).all() + assert [row.name for row in rows] == ["a2", "b2"] + assert [row.version_id for row in rows] == [1, 1] + + +# --- server managed versions --- + + +def test_server_managed_version_is_never_set(): + query = ServerVersionedRepository().update({"name": "jane"}).query + + assert "xmin" not in str(query.compile(compile_kwargs={"literal_binds": True})) + + +def test_server_managed_version_is_recognised(): + assert ServerVersionedRepository._is_server_versioned() + assert not VersionedArticleRepository._is_server_versioned() + + +# --- multi row updates --- + + +def test_update_of_multiple_rows_versions_each_of_them(repository): + versioned = repository._with_version_increment( + [{"id": 1, "name": "a"}, {"id": 2, "name": "b"}] + ) + + versions = [row["version_id"] for row in versioned] + assert all(isinstance(version, UUID) for version in versions) + assert versions[0] != versions[1] + + +def test_counter_update_of_multiple_rows_leaves_the_version_alone(counters): + values = [{"id": 1, "name": "a"}, {"id": 2, "name": "b"}] + + assert counters._with_version_increment(values) == values + + +def test_server_managed_guard_compares_the_column_as_text(): + """``xid = varchar`` is not an operator PostgreSQL has.""" + guard = ServerVersionedRepository()._version_filter(999) + + sql = str(guard.compile(compile_kwargs={"literal_binds": True})) + assert sql == "CAST(test_versioned_server.xmin AS TEXT) = '999'" From 6b21e341e37049ce883556ce80dcab1e4fc7cd1a Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Tue, 25 Aug 2026 00:01:13 +0200 Subject: [PATCH 07/11] docs: document i18n and fill the api reference gaps i18n was not mentioned anywhere in docs/ or the README despite shipping two storage backends. Add a narrative page covering both, the startup callables, the multi-locale helpers and per-backend behaviour, and note that mypy rejects translation_table() as a dynamic base. The api reference had no i18n and no vectors section either, though vectors ships a full page of its own; add both. Co-Authored-By: Claude Opus 5 (1M context) --- README.md | 33 +++++++ docs/i18n.md | 204 ++++++++++++++++++++++++++++++++++++++++++ docs/index.md | 3 + docs/reference/api.md | 58 ++++++++++++ mkdocs.yaml | 1 + 5 files changed, 299 insertions(+) create mode 100644 docs/i18n.md diff --git a/README.md b/README.md index a3943d7..9c27e0e 100644 --- a/README.md +++ b/README.md @@ -51,6 +51,8 @@ Repository: https://github.com/asynq-io/sqlargon - **Auditable models** — append-only versioned history with point-in-time reads and restore - **Vector search** — embeddings with cosine, L2, dot and L1 similarity, full-text and hybrid reciprocal-rank-fusion search on PostgreSQL and SQLite +- **Internationalization** — multi-locale text in a JSON column or a translation table, + with per-request locales and fallback chains - **FastAPI-ready** — repositories and units of work work directly as dependencies - **Alembic migrations** — async-first migration setup - **OpenTelemetry** — optional SQLAlchemy instrumentation @@ -410,6 +412,37 @@ a human-readable counter (`AuditableBase`) or a sortable UUIDv7 (`UUIDAuditableB `sqlargon.audit` relates other tables to one exact version or to whichever is newest. See the [documentation](https://asynq-io.github.io/sqlargon/auditable/) for the full picture. +## Internationalization + +`sqlargon.i18n` keeps text in more than one locale and reads back whichever the current +request wants. Register a locale getter and a fallback chain at startup, and attribute +access stays a plain string in the active locale: + +```python +from sqlargon.i18n import TranslatedString, Translation, set_fallback_chain, set_locale_getter + +set_locale_getter(lambda: request_locale.get()) +set_fallback_chain(lambda locale: (locale or "en", "en")) + + +class Post(TranslationMixin, Base): + title: Mapped[Translation] = mapped_column(TranslatedString()) + + +await PostRepository().create(title={"en": "Hello", "pl": "Czesc"}) + +post = await PostRepository().select().one() +str(post.title) # "Czesc" under a "pl" locale +post.title.data # {"en": "Hello", "pl": "Czesc"} +``` + +Every column operator is rewritten onto the active locale's text, so `Post.title == "Czesc"`, +`.like(...)` and `order_by` need no join and no special syntax — the dialect-specific JSON +read is handled per backend. Long text or many locales are better served by the second +backend, `translation_table`, which keeps one row per locale in a side table and joins it +through `TranslatedRepository`. See the +[documentation](https://asynq-io.github.io/sqlargon/i18n/) for the full picture. + ## FastAPI Repository and unit-of-work `__init__` take no arguments, so subclasses work directly as diff --git a/docs/i18n.md b/docs/i18n.md new file mode 100644 index 0000000..767c1a5 --- /dev/null +++ b/docs/i18n.md @@ -0,0 +1,204 @@ +# Internationalization + +`sqlargon.i18n` keeps a model's text in more than one locale and reads back +whichever one the current request wants. Two backends store it, both reached +through the same attribute access — `article.title` is a string, in the active +locale — so the choice is about storage, not about calling code. + +Nothing is configured by default. Register the two callables at startup: + +```python +from sqlargon.i18n import set_fallback_chain, set_locale_getter + +set_locale_getter(lambda: request_locale.get()) +set_fallback_chain(lambda locale: (locale or "en", "en")) +``` + +`set_locale_getter` returns the locale of the request being served; it is read +at SQL execution time through a late-binding bind parameter, so one cached +statement serves every locale. `set_fallback_chain` turns a locale into the +ordered list to try — `("de", "en")` reads German where it exists and English +where it does not. Both raise `RuntimeError` until they are set. + +## Choosing a backend + +| | `TranslatedString` column | `translation_table` | +| --- | --- | --- | +| Storage | one JSON column, `{locale: text}` | one row per locale in a side table | +| Reads | no join | outer join, added by `TranslatedRepository` | +| Adding a locale | no DDL | no DDL | +| Per-locale indexes | expression index only | ordinary column index | +| Per-locale constraints | no | yes | +| Row size | grows with every locale | constant | + +The JSON column suits a handful of short fields and a handful of locales. The +translation table suits long text, many locales, or a locale that needs its own +index or unique constraint. + +## The JSON column + +`TranslatedString` is a JSON column holding `{locale: text}` that reads back as +a `Translation`: + +```python +import sqlalchemy as sa +from sqlalchemy.orm import Mapped, mapped_column + +from sqlargon import Base, SQLAlchemyRepository +from sqlargon.i18n import TranslatedString, Translation, TranslationMixin + + +class Post(TranslationMixin, Base): + id: Mapped[int] = mapped_column(sa.Integer, primary_key=True) + title: Mapped[Translation] = mapped_column(TranslatedString()) + + +class PostRepository(SQLAlchemyRepository[Post]): + pass + + +await PostRepository().create(title={"en": "Hello", "pl": "Czesc"}) + +post = await PostRepository().select().one() +str(post.title) # "Hello" under an "en" locale, "Czesc" under "pl" +post.title.data # {"en": "Hello", "pl": "Czesc"} +``` + +Every column operator is rewritten to act on the active locale's text, so no +join and no special syntax are needed: + +```python +await PostRepository().count(Post.title == "Czesc") # under a "pl" locale +await PostRepository().select().order_by(Post.title.asc()).all() +await PostRepository().count(Post.title.like("Hello%")) +``` + +The JSON methods inherited from [`JSON`](reference/types.md#json) — +`has_key`, `get`, `contains`, `set_key`, and the rest — still address the whole +locale map, keyed by locale *name*: `Post.title.has_key("pl")` asks whether a +Polish translation exists. + +`Translation` is a `str` subclass carrying the other locales alongside the +active text, so it drops into pydantic models and templates unchanged: + +| | | +| --- | --- | +| `str(value)` | the active locale's text, after walking the fallback chain | +| `value.data` | a copy of every known translation | +| `value.get("pl")` | one locale, or `None` | +| `value.update(text, "pl")` | a **new** `Translation` with that locale replaced | + +It validates from either a string or a `{locale: text}` mapping and serializes +to the active locale's text, so a pydantic response model needs no adapter. + +## The translation table + +`translation_table(Parent)` builds the declarative base of a side table: a copy +of the parent's primary key cascading back to it, plus a `locale` column, all +part of the translation table's own key. + +```python +from sqlargon.i18n import TranslatableBase, TranslatedRepository, translation_table + + +class Article(TranslatableBase): + __translated_fields__ = ("title", "body") + + id: Mapped[int] = mapped_column(sa.Integer, primary_key=True) + slug: Mapped[str] = mapped_column(sa.Unicode(64)) + + +class ArticleTranslation(translation_table(Article)): + title: Mapped[str] = mapped_column(sa.Unicode(255)) + body: Mapped[str] = mapped_column(sa.UnicodeText()) + + +class ArticleRepository(TranslatedRepository[Article]): + pass +``` + +!!! note + + `translation_table(Article)` is a call in a base class list, which mypy + rejects outright as a dynamic base — no return annotation avoids it. Add + `# type: ignore[misc]` to the class statement in a type-checked codebase: + + ```python + class ArticleTranslation(translation_table(Article)): # type: ignore[misc] + ... + ``` + +Each name in `__translated_fields__` becomes a hybrid property reading the +active locale and writing to it, creating the locale row on demand. A field the +translation table has no column for is a typo, raised where it is declared +rather than at flush time. + +The translated columns are made nullable whatever their annotation says: a +locale row carries only the fields translated to that locale, so `NULL` is how +an untranslated field is stored, and reads skip it. + +```python +article = Article(slug="hello") +article.title = {"en": "Hello", "pl": "Czesc"} +article.body = {"en": "Body"} # no Polish body yet + +str(article.title) # "Czesc" under a "pl" locale +article.body.data # {"en": "Body"} +``` + +Writes merge — assigning a plain string only touches the active locale, and no +write drops a locale it does not name. Assigning `None` clears the field in +every locale. + +### Querying it + +`TranslatedRepository.select()` outer-joins the active locale's row, so class +level expressions resolve to the translation table's columns and filtering +needs no explicit join: + +```python +await ArticleRepository().select().filter(Article.title == "Czesc").all() +await ArticleRepository().select().order_by(Article.title.asc()).all() +``` + +The join is the point of the repository: `_current_translation` is declared +`lazy="raise"` so it is never loaded on its own, which keeps a joined query at +one statement and stops an async session tripping over a lazy load. Reads of an +already-loaded row go through `_translations`, eagerly loaded with the row. + +A model must inherit `TranslatableBase` — or, at runtime, the looser +`TranslatableMixin` — or `TranslatedRepository` raises `TypeError` on +subclassing, the way `SoftDeleteRepository` validates its own. + +## Multi-locale helpers + +Both backends share `TranslationMixin`, for the cases that reach past the +active locale — an editing form showing every translation, say: + +```python +article.get_translations("title") # {"en": "Hello", "pl": "Czesc"} +article.get_translation("title", "pl") # "Czesc", or None +article.set_translation("title", "Witaj", "pl") # merges +article.clear_translations("title") # drops every locale +``` + +`get_translation` with an explicit locale looks it up as given; only the active +locale walks its fallback chain. + +## Backend support + +| | PostgreSQL | SQLite | MySQL / MariaDB | +| --- | --- | --- | --- | +| `translation_table` | yes | yes | yes | +| `TranslatedString` storage | `JSONB` | `JSON` | `JSON` | +| `TranslatedString` reads | `->>` | `json_extract` | `json_unquote(json_extract(...))` | + +The translation table is ordinary columns and an ordinary join, so it behaves +identically everywhere. The JSON column reads a locale out of a document, which +each dialect spells its own way — `sqlargon.i18n.expression` registers a +compilation hook per dialect and a default that covers the rest. + +MySQL and MariaDB need the extra `json_unquote`: `JSON_EXTRACT` hands back a +quoted JSON scalar, so `LIKE` and ordering would match against the quotes too. +An equality comparison coerces its operand to JSON and agrees either way, which +is what makes the difference easy to miss. diff --git a/docs/index.md b/docs/index.md index 7a850b6..a7fb211 100644 --- a/docs/index.md +++ b/docs/index.md @@ -53,6 +53,8 @@ Repository: [https://github.com/asynq-io/sqlargon](https://github.com/asynq-io/s reads and restore - **Vector search** — [embeddings with similarity, full-text and hybrid reciprocal-rank-fusion search](vectors.md) on PostgreSQL and SQLite +- **Internationalization** — [multi-locale text](i18n.md) in a JSON column or a + translation table, with per-request locales and fallback chains - **FastAPI-ready** — repositories and units of work work directly as dependencies - **Alembic migrations** — async-first [migration setup](migrations.md) - **OpenTelemetry** — optional SQLAlchemy instrumentation @@ -135,6 +137,7 @@ or from `DATABASE_*` environment variables. - **[Cron](cron.md)** — database-backed scheduling with namespaces and multi-instance safety. - **[Outbox](outbox.md)** — the transactional outbox pattern and its relay. - **[Vector Search](vectors.md)** — embeddings, similarity and hybrid search. +- **[Internationalization](i18n.md)** — multi-locale text and per-request locales. - **[Auditable Models](auditable.md)** — append-only versioned history. - **[Examples](examples.md)** — end-to-end recipes: a FastAPI service, batch workers, multi-tenant sharding, testing. diff --git a/docs/reference/api.md b/docs/reference/api.md index bd25dbc..a2b2ac4 100644 --- a/docs/reference/api.md +++ b/docs/reference/api.md @@ -132,6 +132,64 @@ Generated from the source. See [Usage](../usage.md) for a narrative introduction ::: sqlargon.integrations.eventiq.eventiq_publisher +## Vector search + +::: sqlargon.vectors.VectorRepository + +::: sqlargon.vectors.TextSearchRepository + +::: sqlargon.vectors.HybridVectorRepository + +::: sqlargon.vectors.VectorCollectionRepository + +::: sqlargon.vectors.EmbeddingMixin + +::: sqlargon.vectors.TextMixin + +::: sqlargon.vectors.AttributesMixin + +::: sqlargon.vectors.VectorCollectionMixin + +::: sqlargon.types.vector.Vector + +::: sqlargon.types.vector.DistanceMetric + +::: sqlargon.vectors.init_vectors + +::: sqlargon.vectors.register_sqlite_vector + +## Internationalization + +::: sqlargon.i18n.TranslatedRepository + +::: sqlargon.i18n.TranslatableMixin + +::: sqlargon.i18n.TranslatableBase + +::: sqlargon.i18n.TranslationMixin + +::: sqlargon.i18n.TranslatedString + +::: sqlargon.i18n.Translation + +::: sqlargon.i18n.translation_table + +::: sqlargon.i18n.translation_class + +::: sqlargon.i18n.current_translation + +::: sqlargon.i18n.set_locale_getter + +::: sqlargon.i18n.get_locale + +::: sqlargon.i18n.set_fallback_chain + +::: sqlargon.i18n.fallback_chain + +::: sqlargon.i18n.select_current + +::: sqlargon.i18n.as_translation + ## ORM and types ::: sqlargon.orm.Base diff --git a/mkdocs.yaml b/mkdocs.yaml index f5d93fd..a3b83b2 100644 --- a/mkdocs.yaml +++ b/mkdocs.yaml @@ -27,6 +27,7 @@ nav: - "Outbox": outbox.md - "Vector Search": vectors.md - "Auditable Models": auditable.md + - "Internationalization": i18n.md - "Examples": examples.md - "Reference": - "Column Types & Mixins": reference/types.md From b0501dec65f876ae7aa34f8cf159c0cc38b0dd19 Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Tue, 25 Aug 2026 12:27:03 +0200 Subject: [PATCH 08/11] f --- docs/outbox.md | 35 ++++++++++++++++++++++++------ sqlargon/i18n/translatable.py | 15 ++++++++++--- sqlargon/outbox/config.py | 12 +++++++++-- sqlargon/outbox/models.py | 8 +++---- sqlargon/outbox/repository.py | 7 +++--- tests/e2e/test_outbox.py | 5 ++--- tests/test_outbox.py | 40 +++++++++++++++++++++++++++++++---- 7 files changed, 96 insertions(+), 26 deletions(-) diff --git a/docs/outbox.md b/docs/outbox.md index 2b8a8b7..cebe12b 100644 --- a/docs/outbox.md +++ b/docs/outbox.md @@ -105,6 +105,29 @@ the row — and its value is stringified as it is: a `UUID` becomes its canonical string form. `{id}` reads the row's `id`, `{organization_id}` its `organization_id`, and so on. +`{operation}` is the exception: it names the write rather than the row, and is +filled with `created`, `updated` or `deleted` — the same word the event `type` +ends in. One template then serves every operation: + +```python +class UserRepository(OutboxRepository[User]): + outbox = OutboxConfig( + topic="events.organizations.{organization_id}.users.{id}.{operation}", + type_prefix="user", + ) + + +user = await UserRepository().create(name="John", organization_id=21) +# -> topic "events.organizations.21.users..created" +await UserRepository().update({"name": "Jane"}, id=user.id) +# -> topic "events.organizations.21.users..updated" +await UserRepository().remove(id=user.id) +# -> topic "events.organizations.21.users..deleted" +``` + +A column named `operation` is shadowed by it; name the placeholder for the +write and read the column from the payload. + The value is read the same way the payload and the extra attributes are: from the row the write produced (before it, for a delete). It is read **when the write happens**, not when the relay publishes the event, so a templated topic @@ -112,8 +135,8 @@ always reflects the state at write time. A repository serves one event per written row, so a bulk write of rows from different organizations lands on their own topics. -`format_topic(topic, row)` does the substitution on its own and is exported -from `sqlargon.outbox`. +`format_topic(topic, row, operation)` does the substitution on its own and is +exported from `sqlargon.outbox`. ### Extra attributes @@ -230,8 +253,8 @@ an existing task group with `await tg.start(relay.run)`. Delivery semantics: - Events are published **one at a time, in the order they were written** - (`created_at`, then `id`). An outbox that reorders its events is not much of - an outbox. + (by `id`, a UUIDv7 and so sortable by write time). An outbox that reorders + its events is not much of an outbox. - A publisher that raises stops the batch, so a broker outage delays the events behind the failed one rather than letting them overtake it. The failure is logged (logger `sqlargon.outbox.relay`), recorded in @@ -356,8 +379,8 @@ it, and an Alembic autogenerate pass picks it up (see | Column | Purpose | | --- | --- | -| `id` | The CloudEvents `id`. | -| `created_at` | The CloudEvents `time`, and the dispatch order. | +| `id` | The CloudEvents `id`, a UUIDv7 -- and the dispatch order. | +| `created_at` | The CloudEvents `time`. | | `topic`, `type`, `source` | The CloudEvents routing attributes. | | `data` | The row snapshot, as JSON. | | `attributes` | The extra CloudEvents attributes, as JSON. | diff --git a/sqlargon/i18n/translatable.py b/sqlargon/i18n/translatable.py index abd0a6b..26e9ed6 100644 --- a/sqlargon/i18n/translatable.py +++ b/sqlargon/i18n/translatable.py @@ -178,8 +178,11 @@ def __init_subclass__(cls, **kwargs: Any) -> None: @declared_attr @classmethod def _translations(cls) -> Mapped[dict[str, Any]]: + def target() -> type[Any]: + return translation_class(cls) + return relationship( - lambda: translation_class(cls), + target, collection_class=attribute_keyed_dict("locale"), cascade="all, delete-orphan", lazy="selectin", @@ -188,9 +191,15 @@ def _translations(cls) -> Mapped[dict[str, Any]]: @declared_attr @classmethod def _current_translation(cls) -> Mapped[Any]: + def target() -> type[Any]: + return translation_class(cls) + + def primaryjoin() -> Any: + return _locale_join(cls, current_locale()) + return relationship( - lambda: translation_class(cls), - primaryjoin=lambda: _locale_join(cls, current_locale()), + target, + primaryjoin=primaryjoin, uselist=False, viewonly=True, lazy="raise", diff --git a/sqlargon/outbox/config.py b/sqlargon/outbox/config.py index 790185c..5f5f169 100644 --- a/sqlargon/outbox/config.py +++ b/sqlargon/outbox/config.py @@ -44,10 +44,18 @@ class Operation(str, Enum): ALL_OPERATIONS = frozenset(Operation) -def format_topic(topic: str, row: Any) -> str: +def format_topic(topic: str, row: Any, operation: Operation | None = None) -> str: + """Fill a topic's placeholders from the row the event was written from. + + ``{operation}`` names the write itself rather than an attribute of the + row, and wins over a column of that name. + """ if "{" not in topic: return topic - return topic.format(**vars(row)) + values = dict(vars(row)) + if operation is not None: + values["operation"] = operation.value + return topic.format(**values) @dataclass(frozen=True, slots=True) diff --git a/sqlargon/outbox/models.py b/sqlargon/outbox/models.py index 5ee2442..5361952 100644 --- a/sqlargon/outbox/models.py +++ b/sqlargon/outbox/models.py @@ -4,13 +4,13 @@ import sqlalchemy as sa from sqlalchemy.orm import Mapped, mapped_column -from sqlargon.mixins import CreatedUpdatedMixin, UUIDModelMixin +from sqlargon.mixins import CreatedUpdatedMixin, UUIDV7ModelMixin from sqlargon.orm import Base from sqlargon.types import JSON, Timestamp, now from sqlargon.utils import utc_now -class OutboxEvent(UUIDModelMixin, CreatedUpdatedMixin, Base): +class OutboxEvent(UUIDV7ModelMixin, CreatedUpdatedMixin, Base): """A CloudEvent awaiting publication. ``id`` is the CloudEvent ``id`` and ``created_at`` its ``time``; the @@ -24,9 +24,7 @@ class OutboxEvent(UUIDModelMixin, CreatedUpdatedMixin, Base): __tablename__ = "outbox_events" __table_args__ = ( - sa.Index( - "idx_outbox_events_pending", "published_at", "available_at", "created_at" - ), + sa.Index("idx_outbox_events_pending", "published_at", "id", "available_at"), ) topic: Mapped[str] = mapped_column(sa.String(255), nullable=False) diff --git a/sqlargon/outbox/repository.py b/sqlargon/outbox/repository.py index 86a9731..4f2fe4b 100644 --- a/sqlargon/outbox/repository.py +++ b/sqlargon/outbox/repository.py @@ -58,7 +58,7 @@ async def claim_pending( events = await ( self.select(with_for_update={"skip_locked": True}) .filter(*filters) - .order_by(OutboxEvent.created_at, OutboxEvent.id) + .order_by(OutboxEvent.id) .limit(limit) .all() ) @@ -164,7 +164,8 @@ def topic(self) -> str: The configured topic, or the table name when none is configured. A topic with ``{placeholders}`` is a template: each event's topic is - filled from the row it was written from. + filled from the row it was written from, and ``{operation}`` from the + write that produced it. """ return self.outbox.topic or self.model.__tablename__ @@ -232,7 +233,7 @@ def _build_events( topic = self.topic return [ { - "topic": format_topic(topic, row), + "topic": format_topic(topic, row, operation), "type": event_type, "source": self.outbox.source, "data": to_jsonable_python( diff --git a/tests/e2e/test_outbox.py b/tests/e2e/test_outbox.py index 9677de4..834664a 100644 --- a/tests/e2e/test_outbox.py +++ b/tests/e2e/test_outbox.py @@ -68,8 +68,7 @@ async def test_update_and_delete_are_recorded( await outbox_users.remove(OutboxUser.id == user.id) types = [ - event.type - for event in await events.select().order_by(OutboxEvent.created_at).all() + event.type for event in await events.select().order_by(OutboxEvent.id).all() ] assert types == ["user.created", "user.updated", "user.deleted"] @@ -92,7 +91,7 @@ async def test_an_upsert_records_what_each_row_turned_out_to_be( recorded = { event.data["name"]: event.type - for event in await events.select().order_by(OutboxEvent.created_at).all() + for event in await events.select().order_by(OutboxEvent.id).all() } assert recorded == { "John": "user.created", diff --git a/tests/test_outbox.py b/tests/test_outbox.py index 7899d61..3d950ef 100644 --- a/tests/test_outbox.py +++ b/tests/test_outbox.py @@ -6,6 +6,7 @@ from contextvars import ContextVar from datetime import timedelta +from types import SimpleNamespace from uuid import uuid4 import anyio @@ -125,7 +126,7 @@ def on_conflict(self) -> OnConflictOptions: class OrganizationUserRepository(OutboxRepository[OrganizationUser]): outbox = OutboxConfig( - topic="events.organizations.{organization_id}.users.{id}.created", + topic="events.organizations.{organization_id}.users.{id}.{operation}", type_prefix="org_user", ) @@ -186,7 +187,7 @@ def events(): async def stored(events: EventRepository) -> list[OutboxEvent]: - return list(await events.select().order_by(OutboxEvent.created_at).all()) + return list(await events.select().order_by(OutboxEvent.id).all()) # --- configuration --- @@ -285,6 +286,21 @@ def test_format_topic_fills_an_id_placeholder(): assert format_topic("events.users.{id}", row) == f"events.users.{row.id}" +def test_format_topic_fills_an_operation_placeholder(): + row = OrganizationUser(id=uuid4(), organization_id=uuid4()) + + assert ( + format_topic("events.users.{id}.{operation}", row, Operation.DELETED) + == f"events.users.{row.id}.deleted" + ) + + +def test_an_operation_placeholder_wins_over_a_column_of_that_name(): + row = SimpleNamespace(operation="whatever") + + assert format_topic("{operation}", row, Operation.UPDATED) == "updated" + + def test_a_missing_attribute_raises_like_str_format_does(): row = OrganizationUser() @@ -319,15 +335,31 @@ async def test_each_row_of_a_bulk_write_gets_its_own_topic(org_users, events): async def test_a_templated_topic_is_filled_at_write_time(org_users, events): + first, second = uuid4(), uuid4() + user = await org_users.create(name="John", organization_id=first) + + await org_users.update_one({"organization_id": second}, id=user.id) + + topics = [event.topic for event in await stored(events)] + assert topics == [ + f"events.organizations.{first}.users.{user.id}.created", + f"events.organizations.{second}.users.{user.id}.updated", + ] + + +async def test_an_operation_placeholder_follows_the_write(org_users, events): organization = uuid4() user = await org_users.create(name="John", organization_id=organization) await org_users.update_one({"name": "Jane"}, id=user.id) + await org_users.remove(id=user.id) + prefix = f"events.organizations.{organization}.users.{user.id}" topics = [event.topic for event in await stored(events)] assert topics == [ - f"events.organizations.{organization}.users.{user.id}.created", - f"events.organizations.{organization}.users.{user.id}.created", + f"{prefix}.created", + f"{prefix}.updated", + f"{prefix}.deleted", ] From ac9825c6880a4ced5aad75d6b1a5713fe02e3161 Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Tue, 25 Aug 2026 14:28:46 +0200 Subject: [PATCH 09/11] f --- docs/outbox.md | 6 +++++- docs/reference/types.md | 15 +++++++++------ sqlargon/types/uuid.py | 25 +++++++++++++++++++++---- tests/e2e/README.md | 8 ++++---- tests/e2e/backends.py | 15 ++++++++------- tests/e2e/conftest.py | 5 +---- tests/e2e/models.py | 16 +++++++--------- tests/e2e/test_auditable.py | 2 -- tests/test_types.py | 17 +++++++++++++++++ 9 files changed, 72 insertions(+), 37 deletions(-) diff --git a/docs/outbox.md b/docs/outbox.md index cebe12b..4588cd1 100644 --- a/docs/outbox.md +++ b/docs/outbox.md @@ -391,7 +391,11 @@ it, and an Alembic autogenerate pass picks it up (see | `last_error` | Why the last attempt failed. | Every column carries a server default as well as a client one, so a row -inserted without naming them is still complete. +inserted without naming them is still complete. `id` defaults to `uuidv7()`, +which is a PostgreSQL 18 builtin — an older server +[falls back](reference/types.md#uuids) to a random v4 value, so a row the +library itself writes is still a UUIDv7 (the client default mints it), while +one inserted by raw SQL there is not, and so does not sort by write time. `OutboxEventRepository` exposes the table directly, for monitoring or for a dispatcher of your own: `pending_count()`, `claim_pending()`, diff --git a/docs/reference/types.md b/docs/reference/types.md index 8073d1a..d504c12 100644 --- a/docs/reference/types.md +++ b/docs/reference/types.md @@ -33,12 +33,15 @@ class User(Base): | Element | PostgreSQL | MySQL | SQLite | | --- | --- | --- | --- | | `GenerateUUID` | `GEN_RANDOM_UUID()` | `RANDOM_BYTES`-based v4 expression | `randomblob`-based v4 expression | -| `GenerateUUIDV7` | `uuidv7()` | `NOW(3)` + `RANDOM_BYTES` v7 expression | same v4 expression as above | - -`uuidv7()` requires a PostgreSQL server that ships the function (18+, or an extension), and -the MySQL v7 expression requires MySQL 5.6.4+. SQLite has no UUID v7 equivalent — it falls -back to the random v4 expression, so rely on the client-side `default=uuid7` there if -ordering matters. +| `GenerateUUIDV7` | `uuidv7()`, or `GEN_RANDOM_UUID()` below 18 | `NOW(3)` + `RANDOM_BYTES` v7 expression | same v4 expression as above | + +`uuidv7()` is a PostgreSQL 18 builtin. On an older server `GenerateUUIDV7` compiles to +`GEN_RANDOM_UUID()` instead, so the DDL still runs — the default just yields a random v4 +value rather than a time-ordered one. The version is read off the connected dialect, so a +dialect that has not seen a server yet (an offline `create_all()` dump, say) compiles the +builtin. The MySQL v7 expression requires MySQL 5.6.4+, and SQLite has no UUID v7 +equivalent — it falls back to the random v4 expression too. Where the fallback applies, +rely on the client-side `default=uuid7` if ordering matters. ## Timestamps diff --git a/sqlargon/types/uuid.py b/sqlargon/types/uuid.py index a477030..126b394 100644 --- a/sqlargon/types/uuid.py +++ b/sqlargon/types/uuid.py @@ -52,19 +52,36 @@ class GenerateUUIDV7(FunctionElement): name = "uuidv7_default" +#: The PostgreSQL release that made ``uuidv7()`` a server builtin. +POSTGRESQL_UUIDV7_VERSION = (18,) + + @compiles(GenerateUUID, "postgresql") @compiles(GenerateUUID) -def _generate_uuid_postgresql( - _element: GenerateUUID, _compiler: Any, **_kwargs: Any -) -> str: +def _generate_uuid_postgresql(_element: Any, _compiler: Any, **_kwargs: Any) -> str: return "GEN_RANDOM_UUID()" @compiles(GenerateUUIDV7, "postgresql") @compiles(GenerateUUIDV7) def _generate_uuidv7_postgresql( - _element: GenerateUUID, _compiler: Any, **_kwargs: Any + element: GenerateUUIDV7, compiler: Any, **kwargs: Any ) -> str: + """ + Generates a UUID v7 in PostgreSQL 18+, which has ``uuidv7()`` builtin. + + An older server falls back to :func:`_generate_uuid_postgresql`, so the + DDL it rejects still runs -- at the cost of a random, v4 value rather + than a time ordered one. The version is only known once the dialect has + seen a server, so the builtin is what an unconnected dialect compiles. + """ + version = compiler.dialect.server_version_info + if ( + compiler.dialect.name == "postgresql" + and version is not None + and tuple(version) < POSTGRESQL_UUIDV7_VERSION + ): + return _generate_uuid_postgresql(element, compiler, **kwargs) return "uuidv7()" diff --git a/tests/e2e/README.md b/tests/e2e/README.md index 734d0ef..ed81705 100644 --- a/tests/e2e/README.md +++ b/tests/e2e/README.md @@ -78,10 +78,10 @@ Known gaps: - The SQLite `has_any_key`/`has_all_keys` operators match JSON values, not object keys. - `GenerateUUIDV7` needs PostgreSQL 18 for `uuidv7()`, and falls back to a - random, v4 shaped value on SQLite. The `postgres17` backend runs the same - server without that column: the `server_side_uuidv7` capability is off, so - the `uuidv7()` table is not created and the test for it is skipped, pinning - that the rest of the suite still passes on a pre-18 server. + random, v4 shaped value below it and on SQLite. Every table is created on + the `postgres17` backend all the same, since the fallback is what its DDL + gets; only the test asserting a server generated value really is a v7 is + skipped, the `server_side_uuidv7` capability being off there. - **Vector search** needs pgvector on PostgreSQL and the `sqliteai-vector` loadable extension on SQLite, so `vector_search` is off for the MySQL family and for `postgres17` — `VectorDoc` and `VectorCollection` are `uuidv7()` diff --git a/tests/e2e/backends.py b/tests/e2e/backends.py index 2cf3c33..8ddf73d 100644 --- a/tests/e2e/backends.py +++ b/tests/e2e/backends.py @@ -18,11 +18,11 @@ from contextlib import AbstractContextManager from pathlib import Path -# ``uuidv7()`` is a PostgreSQL 18 builtin, so GenerateUUIDV7 needs at least it; -# postgres17 is run without the v7 server default to pin what still works there. -# The 18 image is the pgvector build -- that PostgreSQL plus the extension the -# vector suite needs -- so it stands in for the plain one rather than adding a -# backend; 17 has no use for it, the vector tables being uuidv7() defaulted. +# ``uuidv7()`` is a PostgreSQL 18 builtin, so GenerateUUIDV7 falls back to +# GEN_RANDOM_UUID() below it; postgres17 pins that the suite passes on the +# fallback. The 18 image is the pgvector build -- that PostgreSQL plus the +# extension the vector suite needs -- so it stands in for the plain one rather +# than adding a backend; 17 runs without pgvector, and so without vectors. POSTGRES_IMAGE = os.environ.get("SQLARGON_E2E_POSTGRES_IMAGE", "pgvector/pgvector:pg18") POSTGRES_17_IMAGE = os.environ.get( "SQLARGON_E2E_POSTGRES_17_IMAGE", "postgres:17-alpine" @@ -104,6 +104,7 @@ class Backend: delete_returning: bool native_locks: bool server_side_uuid: bool + #: a ``uuidv7()`` column default yields a real v7 value, not the fallback server_side_uuidv7: bool skip_locked: bool json_key_operators: bool @@ -161,8 +162,8 @@ def is_mysql_family(self) -> bool: delete_returning=True, native_locks=True, server_side_uuid=True, - # uuidv7() is a PostgreSQL 18 builtin the 17 server lacks, and the - # vector tables are defaulted from it, so they cannot be created here + # uuidv7() is a PostgreSQL 18 builtin the 17 server lacks, so the + # column default falls back to a random, v4 value here server_side_uuidv7=False, skip_locked=True, json_key_operators=True, diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 2927e30..e659ee7 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -24,7 +24,6 @@ from .models import ( SERVER_DEFAULT_TABLES, TABLES, - UUIDV7_SERVER_DEFAULT_TABLES, VECTOR_TABLES, XMIN_TABLES, AuditArticleRepository, @@ -87,8 +86,6 @@ def tables(backend: Backend) -> tuple[sa.Table, ...]: result = TABLES if backend.server_side_uuid: result = result + SERVER_DEFAULT_TABLES - if backend.server_side_uuidv7: - result = result + UUIDV7_SERVER_DEFAULT_TABLES if backend.dialect == "postgresql": result = result + XMIN_TABLES if backend.vector_search: @@ -214,7 +211,7 @@ def needs_server_side_uuid(backend: Backend) -> None: @pytest.fixture def needs_server_side_uuidv7(backend: Backend) -> None: if not backend.server_side_uuidv7: - pytest.skip(f"{backend.name} lacks uuidv7(), a PostgreSQL 18 server builtin") + pytest.skip(f"{backend.name} defaults a UUIDv7 column to a v4 fallback") @pytest.fixture diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 2cda746..9dbf48f 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -336,6 +336,10 @@ def _tables(*models: type[Base]) -> tuple[sa.Table, ...]: #: Tables every backend can hold; the only ones the e2e suite creates. #: +#: The ``uuidv7()`` defaulted ones are among them: a pre-18 PostgreSQL has no +#: such builtin, but ``GenerateUUIDV7`` falls back to ``GEN_RANDOM_UUID()`` +#: there, so the DDL still runs -- only the value is not a v7. +#: #: A child referencing an audited version comes before the table it points at, #: so the per test cleanup can empty them in this order without tripping the #: foreign key. @@ -349,14 +353,15 @@ def _tables(*models: type[Base]) -> tuple[sa.Table, ...]: AuditComment, AuditFollow, AuditArticle, + UUIDAuditArticle, + ServerDefaultsUUIDV7, I18nPost, I18nArticleTranslation, I18nArticle, ) #: Tables the vector suite needs, which only a backend that can search vectors -#: creates -- ``VectorDoc`` and ``VectorCollection`` carry a ``uuidv7()`` server -#: default on top of the VECTOR columns, so a pre-18 PostgreSQL rejects the DDL. +#: creates -- their VECTOR columns name a type the others have no extension for. VECTOR_TABLES: tuple[sa.Table, ...] = _tables(VectorNote, VectorDoc, VectorCollection) #: Tables whose DDL carries a server side UUID default, which not every @@ -365,10 +370,3 @@ def _tables(*models: type[Base]) -> tuple[sa.Table, ...]: #: Tables that only PostgreSQL can hold (the ``xmin`` system column). XMIN_TABLES: tuple[sa.Table, ...] = _tables(XminUser) - -#: Tables whose DDL carries the server side ``uuidv7()`` default, a -#: PostgreSQL 18 builtin the 17 backend rejects -- ``UUIDAuditArticle`` -#: versions its rows with it. -UUIDV7_SERVER_DEFAULT_TABLES: tuple[sa.Table, ...] = _tables( - ServerDefaultsUUIDV7, UUIDAuditArticle -) diff --git a/tests/e2e/test_auditable.py b/tests/e2e/test_auditable.py index d280917..c16ecb3 100644 --- a/tests/e2e/test_auditable.py +++ b/tests/e2e/test_auditable.py @@ -283,7 +283,6 @@ async def test_a_pinned_version_cannot_be_purged( # --- the UUIDv7 strategy --- -@pytest.mark.usefixtures("needs_server_side_uuidv7") async def test_uuid_strategy_appends_sortable_versions(uuid_audit_articles): entity_id = uuid4() @@ -300,7 +299,6 @@ async def test_uuid_strategy_appends_sortable_versions(uuid_audit_articles): assert (await uuid_audit_articles.get(id=entity_id)).name == "final" -@pytest.mark.usefixtures("needs_server_side_uuidv7") async def test_uuid_strategy_deletes_by_appending_a_tombstone(uuid_audit_articles): entity_id = uuid4() await uuid_audit_articles.create(id=entity_id, name="draft") diff --git a/tests/test_types.py b/tests/test_types.py index c99295b..648ca2c 100644 --- a/tests/test_types.py +++ b/tests/test_types.py @@ -117,6 +117,23 @@ def test_generate_uuid_postgresql(element, expected): ) +@pytest.mark.parametrize( + ("server_version_info", "expected"), + [ + ((17, 4), "GEN_RANDOM_UUID()"), + ((18, 0), "uuidv7()"), + (None, "uuidv7()"), + ], +) +def test_generate_uuidv7_falls_back_below_postgresql_18(server_version_info, expected): + dialect = postgresql.dialect() + dialect.server_version_info = server_version_info + + assert _compile(sa.select(GenerateUUIDV7()), dialect).startswith( + f"SELECT {expected}" + ) + + @pytest.mark.parametrize( ("element", "expected", "not_expected"), [ From ccc912c9abe98137487b06fa59b0addb8af99fd4 Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Tue, 25 Aug 2026 15:09:21 +0200 Subject: [PATCH 10/11] 1.0.1b1 --- pyproject.toml | 2 +- uv.lock | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index aa87bfa..786800e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "sqlargon" -version = "1.0.0" +version = "1.0.1b1" description = "SQLAlchemy repository pattern and utilities" readme = "README.md" requires-python = ">=3.10,<4.0" diff --git a/uv.lock b/uv.lock index 8a0f97d..ab4bef1 100644 --- a/uv.lock +++ b/uv.lock @@ -880,7 +880,7 @@ name = "exceptiongroup" version = "1.3.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions", marker = "python_full_version < '3.11'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/50/79/66800aadf48771f6b62f7eb014e352e5d06856655206165d775e675a02c9/exceptiongroup-1.3.1.tar.gz", hash = "sha256:8b412432c6055b0b7d14c310000ae93352ed6754f70fa8f7c34141f91c4e3219", size = 30371, upload-time = "2025-11-21T23:01:54.787Z" } wheels = [ @@ -2307,7 +2307,7 @@ wheels = [ [[package]] name = "sqlargon" -version = "1.0.0" +version = "1.0.1b1" source = { editable = "." } dependencies = [ { name = "alembic" }, From faac3d622bfa5edda0622ba7bb3a5918dd4bc4e8 Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Tue, 25 Aug 2026 15:19:21 +0200 Subject: [PATCH 11/11] f --- pyproject.toml | 2 +- uv.lock | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 786800e..ddd4532 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "sqlargon" -version = "1.0.1b1" +version = "1.1.0b1" description = "SQLAlchemy repository pattern and utilities" readme = "README.md" requires-python = ">=3.10,<4.0" diff --git a/uv.lock b/uv.lock index ab4bef1..26b771c 100644 --- a/uv.lock +++ b/uv.lock @@ -2307,7 +2307,7 @@ wheels = [ [[package]] name = "sqlargon" -version = "1.0.1b1" +version = "1.1.0b1" source = { editable = "." } dependencies = [ { name = "alembic" },