diff --git a/README.md b/README.md index 02d4d37..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 @@ -331,10 +366,50 @@ 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`. +## 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/index.md b/docs/index.md index 2243ebb1..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 @@ -113,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 index edc7a7e..8cbcf11 100644 --- a/docs/outbox.md +++ b/docs/outbox.md @@ -79,6 +79,43 @@ 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` diff --git a/docs/reference/api.md b/docs/reference/api.md index 491e3cb..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 @@ -122,6 +126,8 @@ Generated from the source. See [Usage](../usage.md) for a narrative introduction ::: sqlargon.outbox.Operation +::: sqlargon.outbox.format_topic + ::: sqlargon.integrations.eventiq.to_cloud_event ::: sqlargon.integrations.eventiq.eventiq_publisher @@ -132,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/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 b391775..f5d93fd 100644 --- a/mkdocs.yaml +++ b/mkdocs.yaml @@ -25,6 +25,8 @@ nav: - "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 1baeb32..41816d6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,6 +26,10 @@ 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 = [ @@ -42,6 +46,8 @@ 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 = [ @@ -153,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 72cb326..5bd6da3 100644 --- a/sqlargon/__init__.py +++ b/sqlargon/__init__.py @@ -1,25 +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 ( + 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, @@ -45,7 +56,14 @@ __all__ = [ "AbstractUnitOfWork", + "AnyAuditableBase", "AnyDatabase", + "AnyVersionedBase", + "AppendOnlyError", + "AuditableBase", + "AuditableMixin", + "AuditableModel", + "AuditableRepository", "Base", "BaseDatabase", "ConcurrentModificationError", @@ -53,6 +71,7 @@ "DatabaseCluster", "DefaultRouter", "DeletedRowExistsError", + "IntegerAuditableMixin", "Model", "ModelRouter", "ORMModel", @@ -69,6 +88,8 @@ "SoftDeleteBase", "SoftDeleteModel", "SoftDeleteRepository", + "UUIDAuditableBase", + "UUIDAuditableMixin", "UUIDVersionedMixin", "VersionedBase", "VersionedMixin", @@ -79,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/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/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 index 17a98c9..2738f92 100644 --- a/sqlargon/outbox/__init__.py +++ b/sqlargon/outbox/__init__.py @@ -1,4 +1,10 @@ -from .config import ALL_OPERATIONS, AttributeSource, Operation, OutboxConfig +from .config import ( + ALL_OPERATIONS, + AttributeSource, + Operation, + OutboxConfig, + format_topic, +) from .models import OutboxEvent from .relay import OutboxRelay, Publisher from .repository import OutboxEventRepository, OutboxRepository @@ -13,4 +19,5 @@ "OutboxRelay", "OutboxRepository", "Publisher", + "format_topic", ] diff --git a/sqlargon/outbox/config.py b/sqlargon/outbox/config.py index 6ef5c8d..790185c 100644 --- a/sqlargon/outbox/config.py +++ b/sqlargon/outbox/config.py @@ -5,7 +5,13 @@ from types import MappingProxyType from typing import Any -__all__ = ["ALL_OPERATIONS", "AttributeSource", "Operation", "OutboxConfig"] +__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 @@ -38,32 +44,14 @@ class Operation(str, Enum): ALL_OPERATIONS = frozenset(Operation) -@dataclass(frozen=True, slots=True) -class OutboxConfig: - """How the writes of one repository are turned into events. - - ``topic`` and ``type_prefix`` both default to the model's table name, - giving events of type ``user.created`` on the ``user`` topic. ``exclude`` - keeps columns out of the payload -- a password hash has no business - leaving the database -- while ``include`` states the payload columns - outright and wins over ``exclude``. - - ``attributes`` names the values that ride next to ``type`` and ``source`` - rather than inside the payload, for a schema whose CloudEvent carries a - tenant or a trace of its own:: +def format_topic(topic: str, row: Any) -> str: + if "{" not in topic: + return topic + return topic.format(**vars(row)) - OutboxConfig( - exclude={"tenant_id"}, - attributes={ - "tenant_id": "tenant_id", - "traceparent": lambda _: trace_id.get(), - }, - ) - - They are read when the write happens, since the relay publishes long - after the context the write ran in is gone. - """ +@dataclass(frozen=True, slots=True) +class OutboxConfig: topic: str | None = None type_prefix: str | None = None source: str | None = None diff --git a/sqlargon/outbox/repository.py b/sqlargon/outbox/repository.py index d10ac22..86a9731 100644 --- a/sqlargon/outbox/repository.py +++ b/sqlargon/outbox/repository.py @@ -5,13 +5,13 @@ import sqlalchemy as sa from pydantic_core import to_jsonable_python -from sqlalchemy.engine.result import IteratorResult, SimpleResultMetaData 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 +from .config import Operation, OutboxConfig, format_topic from .models import OutboxEvent if TYPE_CHECKING: @@ -27,21 +27,6 @@ from sqlargon.typing import MultipleValues, OnConflictOptions, SingleValue, Values -def _as_result(rows: Sequence[object]) -> IteratorResult[Any]: - """Re-wrap already fetched rows as a single-column result. - - The capture hooks have to consume the result of the write they wrap in - order to build the events, so the rows are handed 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 OutboxEventRepository(SQLAlchemyRepository[OutboxEvent]): """Repository for :class:`OutboxEvent` rows, used by the relay.""" @@ -175,7 +160,12 @@ def events(self) -> OutboxEventRepository: @property def topic(self) -> str: - """The topic every event of this repository is published on.""" + """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 @@ -233,14 +223,16 @@ def _build_events( 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, - while the context the write ran in is still the current one. + 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": self.topic, + "topic": format_topic(topic, row), "type": event_type, "source": self.outbox.source, "data": to_jsonable_python( 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 index 17e4c79..fcc6a81 100644 --- a/sqlargon/repository/__init__.py +++ b/sqlargon/repository/__init__.py @@ -1,8 +1,11 @@ +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", 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/base.py b/sqlargon/repository/base.py index 68f4be6..75cccef 100644 --- a/sqlargon/repository/base.py +++ b/sqlargon/repository/base.py @@ -1,7 +1,7 @@ from __future__ import annotations 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, @@ -16,6 +16,7 @@ false, or_, ) +from sqlalchemy.engine.result import IteratorResult, SimpleResultMetaData from sqlalchemy.orm import QueryableAttribute, selectinload from sqlargon.mixins import CreatedUpdatedMixin @@ -58,6 +59,21 @@ __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]): """Repository over a model, bound to a database or cluster. diff --git a/sqlargon/repository/soft_delete.py b/sqlargon/repository/soft_delete.py index 2d6643a..7d4170c 100644 --- a/sqlargon/repository/soft_delete.py +++ b/sqlargon/repository/soft_delete.py @@ -63,13 +63,26 @@ def __init__(self) -> None: 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): + if abstract: + return + mixin, name = cls._required_mixin() + if not issubclass(cls.model, mixin): msg = ( - f"{cls.model.__name__} must inherit from SoftDeleteMixin " + 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"} diff --git a/sqlargon/repository/versioned.py b/sqlargon/repository/versioned.py index ce85e95..e216a2a 100644 --- a/sqlargon/repository/versioned.py +++ b/sqlargon/repository/versioned.py @@ -67,11 +67,11 @@ 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): + mixin, name = cls._required_mixin() + if not issubclass(cls.model, mixin): msg = ( - f"{cls.model.__name__} must inherit from VersionedMixin " - f"(UUIDVersionedMixin or XminVersionedMixin) to be used " - f"with {cls.__name__}" + f"{cls.model.__name__} must inherit from {name} " + f"to be used with {cls.__name__}" ) raise TypeError(msg) if cls._version_col() is None: @@ -82,6 +82,14 @@ def __init_subclass__(cls, *, abstract: bool = False, **kwargs: Any) -> None: ) 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``.""" 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/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 3722241..e00d874 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -13,19 +13,28 @@ 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, ) @@ -33,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.""" @@ -81,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), @@ -191,3 +207,65 @@ def xmin_users() -> 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 ae83a5c..6d3e42e 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -11,22 +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): @@ -135,10 +148,75 @@ 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" @@ -159,13 +237,62 @@ class OutboxUserRepository(OutboxRepository[OutboxUser]): ) +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. +#: +#: 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 + 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 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_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_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_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_outbox.py b/tests/test_outbox.py index b05c7a5..7899d61 100644 --- a/tests/test_outbox.py +++ b/tests/test_outbox.py @@ -21,6 +21,7 @@ OutboxEventRepository, OutboxRelay, OutboxRepository, + format_topic, ) from sqlargon.types import GUID from sqlargon.typing import OnConflictOptions @@ -56,6 +57,15 @@ class OutboxTag(UUIDModelMixin, CreatedUpdatedMixin, Base): 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" @@ -113,6 +123,13 @@ 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 @@ -126,6 +143,7 @@ async def tables(db: Database): OutboxUser.__table__, InsertOnlyNote.__table__, OutboxTag.__table__, + OrganizationUser.__table__, OutboxEvent.__table__, ) async with db.engine.begin() as conn: @@ -157,6 +175,11 @@ def tags(): return TagRepository() +@pytest.fixture +def org_users(): + return OrganizationUserRepository() + + @pytest.fixture def events(): return EventRepository() @@ -240,6 +263,74 @@ def test_attributes_cannot_shadow_the_core_ones(): 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 --- 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 c7e2cbd..6f2db8b 100644 --- a/uv.lock +++ b/uv.lock @@ -1485,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" @@ -2128,6 +2137,12 @@ standard = [ { name = "opentelemetry-instrumentation-sqlalchemy" }, { name = "sqlakeyset" }, ] +vectors = [ + { name = "pgvector" }, +] +vectors-sqlite = [ + { name = "sqliteai-vector" }, +] [package.dev-dependencies] dev = [ @@ -2152,6 +2167,7 @@ dev = [ { name = "pytest-timeout" }, { name = "pytest-xdist" }, { name = "ruff" }, + { name = "sqliteai-vector" }, { name = "testcontainers" }, { name = "watchdog" }, ] @@ -2165,6 +2181,7 @@ docs = [ ] e2e = [ { name = "cryptography" }, + { name = "sqliteai-vector" }, { name = "testcontainers" }, ] lint = [ @@ -2202,14 +2219,16 @@ requires-dist = [ { 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", "outbox", "eventiq", "opentelemetry", "standard"] +provides-extras = ["fastapi", "postgres", "sqlite", "mysql", "pagination", "cron", "outbox", "eventiq", "opentelemetry", "standard", "vectors", "vectors-sqlite"] [package.metadata.requires-dev] dev = [ @@ -2234,6 +2253,7 @@ dev = [ { 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" }, ] @@ -2247,6 +2267,7 @@ docs = [ ] e2e = [ { name = "cryptography" }, + { name = "sqliteai-vector", specifier = ">=1.0.0,<2" }, { name = "testcontainers", specifier = ">=4.15.0" }, ] lint = [ @@ -2266,6 +2287,18 @@ test = [ { 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]] name = "starlette" version = "1.3.1"