From ddb4b80d067d3a99eb5c432c7171a8c2bfde472c Mon Sep 17 00:00:00 2001 From: Radzim Kowalow Date: Sat, 22 Aug 2026 22:44:56 +0200 Subject: [PATCH 1/2] feat: i18n --- 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 | 212 +++++++++++++++++++++++++++++ tests/test_i18n.py | 190 ++++++++++++++++++++++++++ 7 files changed, 870 insertions(+) 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 tests/test_i18n.py 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..d9163c2 --- /dev/null +++ b/sqlargon/i18n/translation.py @@ -0,0 +1,212 @@ +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` -- + ``contains``, ``has_any_key``, ``has_all_keys``, ``json_value`` and + indexing -- still address the whole locale map. + """ + + 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/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 From 17886f30742869328cf3cec91c2aab5cdd3315c6 Mon Sep 17 00:00:00 2001 From: RaRhAeu <37556570+RaRhAeu@users.noreply.github.com> Date: Mon, 24 Aug 2026 22:53:42 +0200 Subject: [PATCH 2/2] feat: json operations (#36) --- README.md | 4 +- docs/reference/types.md | 61 +++- sqlargon/i18n/translation.py | 8 +- sqlargon/types/json.py | 617 ++++++++++++++++++++++++++++++++++- tests/e2e/test_types.py | 156 +++++++++ tests/test_types.py | 487 ++++++++++++++++++++++++++- 6 files changed, 1310 insertions(+), 23 deletions(-) diff --git a/README.md b/README.md index 09c43e4..c294002 100644 --- a/README.md +++ b/README.md @@ -366,7 +366,9 @@ ordering. `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`. diff --git a/docs/reference/types.md b/docs/reference/types.md index 78127ec..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 diff --git a/sqlargon/i18n/translation.py b/sqlargon/i18n/translation.py index d9163c2..88d2a8a 100644 --- a/sqlargon/i18n/translation.py +++ b/sqlargon/i18n/translation.py @@ -158,8 +158,12 @@ class TranslatedString(TypeDecorator[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` -- - ``contains``, ``has_any_key``, ``has_all_keys``, ``json_value`` and - indexing -- still address the whole locale map. + 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 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/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/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}