diff --git a/Cargo.toml b/Cargo.toml index 4d96f77c..68e7fdb8 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -32,6 +32,8 @@ emitter = ["dep:event-emitter-rs"] bus = [] http = ["bus", "dep:axum", "dep:tokio"] grpc = ["bus", "dep:tonic", "dep:prost", "dep:tokio"] +postgres = ["dep:sqlx", "dep:tokio", "sqlx/postgres", "sqlx/runtime-tokio"] +sqlite = ["dep:sqlx", "dep:tokio", "sqlx/runtime-tokio", "sqlx/sqlite"] [dependencies] axum = { version = "0.7", optional = true } @@ -41,6 +43,7 @@ event-emitter-rs = { version = "0.1.4", optional = true } serde = { version = "1.0.210", features = ["derive"] } serde_json = "1.0.128" sourced_rust_macros = { workspace = true } +sqlx = { version = "0.8", default-features = false, optional = true } tonic = { version = "0.12", optional = true } prost = { version = "0.13", optional = true } tokio = { version = "1", features = ["rt-multi-thread", "net", "macros"], optional = true } diff --git a/compose.yaml b/compose.yaml new file mode 100644 index 00000000..0780cb5c --- /dev/null +++ b/compose.yaml @@ -0,0 +1,14 @@ +services: + postgres: + image: postgres:16 + environment: + POSTGRES_USER: sourced + POSTGRES_PASSWORD: sourced + POSTGRES_DB: sourced_rust + ports: + - "5432:5432" + healthcheck: + test: ["CMD-SHELL", "pg_isready -U sourced -d sourced_rust"] + interval: 2s + timeout: 5s + retries: 20 diff --git a/docs/async-repositories.md b/docs/async-repositories.md index b756c899..22d44806 100644 --- a/docs/async-repositories.md +++ b/docs/async-repositories.md @@ -39,3 +39,46 @@ stream-aware contract before SQL code lands. The Postgres repository should implement the async traits directly with `sqlx`. It should not hide database I/O behind the synchronous traits with `block_on`, `block_in_place`, or a blocking wrapper in normal async runtimes. + +## SQLite Adapter + +The optional `sqlite` feature exports `SqliteRepository`, an async-only +SQL-backed adapter for local persistence and conformance work: + +```rust +let repo = sourced_rust::SqliteRepository::connect_and_migrate("sqlite::memory:").await?; +``` + +`SqliteRepository::migrate` applies explicit SQLite migrations from +`migrations/sqlite`. Plain construction from an existing pool does not create +tables implicitly, so applications can control bootstrap order. + +The first SQLite pass persists aggregate events, transactional document read +models, processed-message marks, and snapshots in one SQL transaction when they +are staged through `AsyncCommitBatch`. It intentionally does not claim Postgres +production readiness: Postgres-specific column types, isolation behavior, error +mapping, deployment, and migration validation still belong to the Postgres +adapter and its own tests. + +## Postgres Adapter + +The optional `postgres` feature exports `PostgresRepository`, an async-only +SQLx adapter for the production SQL event-store path: + +```rust +let repo = + sourced_rust::PostgresRepository::connect_and_migrate(database_url).await?; +``` + +Local integration tests can use the root `compose.yaml` service: + +```bash +docker compose up -d postgres +DATABASE_URL=postgres://sourced:sourced@localhost:5432/sourced_rust \ + cargo test --features postgres --test postgres_repository +``` + +The first Postgres pass persists aggregate event streams and snapshots through +explicit migrations in `migrations/postgres`. It rejects non-empty read-model +write plans instead of creating generic read-model tables implicitly; durable +read-model persistence remains a separate adapter track. diff --git a/migrations/postgres/0001_initial.sql b/migrations/postgres/0001_initial.sql new file mode 100644 index 00000000..1da4e679 --- /dev/null +++ b/migrations/postgres/0001_initial.sql @@ -0,0 +1,37 @@ +CREATE TABLE IF NOT EXISTS aggregate_events ( + aggregate_type text NOT NULL, + aggregate_id text NOT NULL, + sequence bigint NOT NULL, + event_name text NOT NULL, + event_version integer NOT NULL DEFAULT 1, + payload bytea NOT NULL, + payload_codec text NOT NULL, + payload_codec_version integer NOT NULL, + metadata jsonb NOT NULL DEFAULT '{}'::jsonb, + recorded_at timestamptz NOT NULL DEFAULT now(), + PRIMARY KEY (aggregate_type, aggregate_id, sequence), + CHECK (aggregate_type <> ''), + CHECK (aggregate_id <> ''), + CHECK (sequence > 0), + CHECK (event_version > 0), + CHECK (payload_codec <> ''), + CHECK (payload_codec_version > 0) +); + +CREATE INDEX IF NOT EXISTS aggregate_events_event_version_idx + ON aggregate_events (aggregate_type, event_name, event_version); + +CREATE INDEX IF NOT EXISTS aggregate_events_recorded_at_idx + ON aggregate_events (recorded_at); + +CREATE TABLE IF NOT EXISTS aggregate_snapshots ( + aggregate_type text NOT NULL, + aggregate_id text NOT NULL, + version bigint NOT NULL, + data bytea NOT NULL, + updated_at timestamptz NOT NULL DEFAULT now(), + PRIMARY KEY (aggregate_type, aggregate_id), + CHECK (aggregate_type <> ''), + CHECK (aggregate_id <> ''), + CHECK (version > 0) +); diff --git a/migrations/sqlite/0001_initial.sql b/migrations/sqlite/0001_initial.sql new file mode 100644 index 00000000..cc31303a --- /dev/null +++ b/migrations/sqlite/0001_initial.sql @@ -0,0 +1,58 @@ +CREATE TABLE IF NOT EXISTS aggregate_events ( + aggregate_type TEXT NOT NULL, + aggregate_id TEXT NOT NULL, + sequence INTEGER NOT NULL, + event_name TEXT NOT NULL, + event_version INTEGER NOT NULL DEFAULT 1, + payload BLOB NOT NULL, + payload_codec TEXT NOT NULL, + payload_codec_version INTEGER NOT NULL, + metadata TEXT NOT NULL DEFAULT '{}', + recorded_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (aggregate_type, aggregate_id, sequence), + CHECK (aggregate_type <> ''), + CHECK (aggregate_id <> ''), + CHECK (sequence > 0), + CHECK (event_version > 0), + CHECK (payload_codec <> ''), + CHECK (payload_codec_version > 0) +); + +CREATE INDEX IF NOT EXISTS aggregate_events_event_version_idx + ON aggregate_events (aggregate_type, event_name, event_version); + +CREATE INDEX IF NOT EXISTS aggregate_events_recorded_at_idx + ON aggregate_events (recorded_at); + +CREATE TABLE IF NOT EXISTS transactional_read_models ( + collection TEXT NOT NULL, + id TEXT NOT NULL, + version INTEGER NOT NULL, + payload BLOB NOT NULL, + updated_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (collection, id), + CHECK (collection <> ''), + CHECK (id <> ''), + CHECK (version > 0) +); + +CREATE TABLE IF NOT EXISTS read_model_processed_messages ( + consumer_name TEXT NOT NULL, + message_id TEXT NOT NULL, + processed_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (consumer_name, message_id), + CHECK (consumer_name <> ''), + CHECK (message_id <> '') +); + +CREATE TABLE IF NOT EXISTS aggregate_snapshots ( + aggregate_type TEXT NOT NULL, + aggregate_id TEXT NOT NULL, + version INTEGER NOT NULL, + data BLOB NOT NULL, + updated_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (aggregate_type, aggregate_id), + CHECK (aggregate_type <> ''), + CHECK (aggregate_id <> ''), + CHECK (version > 0) +); diff --git a/src/lib.rs b/src/lib.rs index 0b87ffdb..1d0770f1 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -17,9 +17,15 @@ pub mod lock; pub mod microsvc; mod outbox; mod outbox_worker; +#[cfg(feature = "postgres")] +pub mod postgres_repo; pub mod queued_repo; pub mod read_model; pub mod snapshot; +#[cfg(feature = "sqlite")] +pub mod sqlite_repo; +#[cfg(any(feature = "postgres", feature = "sqlite"))] +mod sqlx_repo; // Re-export entity types at crate root for convenience pub use entity::{ @@ -46,6 +52,10 @@ pub use aggregate::{ }; pub use hashmap_repo::HashMapRepository; +#[cfg(feature = "postgres")] +pub use postgres_repo::PostgresRepository; +#[cfg(feature = "sqlite")] +pub use sqlite_repo::SqliteRepository; // Re-export lock traits and types at crate root for convenience pub use lock::{InMemoryLock, InMemoryLockManager, Lock, LockError, LockManager}; diff --git a/src/postgres_repo/mod.rs b/src/postgres_repo/mod.rs new file mode 100644 index 00000000..806f8afa --- /dev/null +++ b/src/postgres_repo/mod.rs @@ -0,0 +1,529 @@ +//! Postgres-backed async aggregate repository. +//! +//! This adapter is the production-oriented SQL event-store path. It is +//! feature-gated behind `postgres`, async-only, and intentionally does not +//! create read-model tables in the first pass. + +#![expect( + clippy::manual_async_fn, + reason = "async trait impls return impl Future + Send to preserve public Send bounds" +)] + +use std::future::Future; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use sqlx::postgres::{PgPoolOptions, PgRow}; +use sqlx::{PgPool, Postgres, Row, Transaction}; + +use crate::entity::{Entity, EventRecord}; +use crate::repository::{ + AsyncCommitBatch, AsyncGetStream, AsyncSnapshotStore, AsyncSnapshotWrite, + AsyncTransactionalCommit, PreparedEventAppend, RepositoryError, StreamIdentity, +}; +use crate::snapshot::SnapshotRecord; +use crate::sqlx_repo::{ + self, deserialize_event_metadata, is_postgres_unique_violation, reject_duplicate_streams, + repository_i32_from_u64 as sqlx_repository_i32_from_u64, + repository_i64_from_u64 as sqlx_repository_i64_from_u64, + repository_u16_from_i32 as sqlx_repository_u16_from_i32, + repository_u64_from_i32 as sqlx_repository_u64_from_i32, + repository_u64_from_i64 as sqlx_repository_u64_from_i64, serialize_event_metadata, + validate_entity_id_matches_identity, validate_prepared_appends, validate_snapshot_identity, + validate_supported_event_codec, +}; + +const POSTGRES_SCHEMA: &str = include_str!("../../migrations/postgres/0001_initial.sql"); +const POSTGRES_BACKEND: &str = "postgres"; +const BIGINT_STORAGE: &str = "bigint storage"; +const INTEGER_STORAGE: &str = "integer storage"; + +/// Postgres-backed async repository. +#[derive(Clone)] +pub struct PostgresRepository { + pool: PgPool, +} + +impl PostgresRepository { + /// Create a repository from an existing migrated pool. + pub fn new(pool: PgPool) -> Self { + Self { pool } + } + + /// Open a Postgres pool without applying migrations. + pub async fn connect(database_url: &str) -> Result { + let pool = PgPoolOptions::new() + .max_connections(5) + .connect(database_url) + .await + .map_err(|err| repository_storage_error("connect", err))?; + Ok(Self::new(pool)) + } + + /// Open a Postgres pool and apply the explicit Postgres migrations. + pub async fn connect_and_migrate(database_url: &str) -> Result { + let repo = Self::connect(database_url).await?; + repo.migrate().await?; + Ok(repo) + } + + /// Apply Postgres migrations to this repository's pool. + pub async fn migrate(&self) -> Result<(), RepositoryError> { + Self::migrate_pool(&self.pool).await + } + + /// Apply Postgres migrations to an existing pool. + pub async fn migrate_pool(pool: &PgPool) -> Result<(), RepositoryError> { + for statement in POSTGRES_SCHEMA.split(';') { + let statement = statement.trim(); + if statement.is_empty() { + continue; + } + sqlx::query(statement) + .execute(pool) + .await + .map_err(|err| repository_storage_error("migrate", err))?; + } + Ok(()) + } + + /// Access the underlying SQLx pool for application-specific setup or tests. + pub fn pool(&self) -> &PgPool { + &self.pool + } +} + +impl AsyncGetStream for PostgresRepository { + fn get_stream<'a>( + &'a self, + identity: &'a StreamIdentity, + ) -> impl Future, RepositoryError>> + Send + 'a { + async move { + let rows = sqlx::query( + r#" + SELECT event_name, + event_version, + payload, + payload_codec, + payload_codec_version, + metadata::text AS metadata, + sequence, + EXTRACT(EPOCH FROM recorded_at)::double precision AS recorded_at_epoch + FROM aggregate_events + WHERE aggregate_type = $1 AND aggregate_id = $2 + ORDER BY sequence ASC + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .fetch_all(&self.pool) + .await + .map_err(|err| repository_storage_error("load stream", err))?; + + if rows.is_empty() { + return Ok(None); + } + + let mut events = Vec::with_capacity(rows.len()); + for row in rows { + events.push(event_from_row(row)?); + } + + let mut entity = Entity::new(); + entity.set_id(identity.aggregate_id()); + entity.load_from_history(events); + Ok(Some(entity)) + } + } + + fn get_streams<'a>( + &'a self, + identities: &'a [StreamIdentity], + ) -> impl Future, RepositoryError>> + Send + 'a { + async move { + let mut entities = Vec::with_capacity(identities.len()); + for identity in identities { + if let Some(entity) = self.get_stream(identity).await? { + entities.push(entity); + } + } + Ok(entities) + } + } +} + +impl AsyncTransactionalCommit for PostgresRepository { + fn commit_batch_async<'a>( + &'a self, + batch: AsyncCommitBatch<'a>, + ) -> impl Future> + Send + 'a { + async move { + reject_duplicate_streams(&batch.streams)?; + validate_entity_id_matches_identity(&batch.streams)?; + reject_read_model_plans(&batch)?; + + let prepared = batch + .streams + .iter() + .map(PreparedEventAppend::from_stream_write) + .collect::>(); + validate_prepared_appends(&prepared)?; + + let mut tx = self + .pool + .begin() + .await + .map_err(|err| repository_storage_error("begin commit transaction", err))?; + + for append in &prepared { + let actual = stream_version_in_tx(&mut tx, &append.identity).await?; + if actual != append.expected_version { + return Err(RepositoryError::ConcurrentWrite { + id: append.identity.to_string(), + expected: append.expected_version, + actual, + }); + } + } + + for append in &prepared { + for event in &append.events { + insert_event_in_tx( + &self.pool, + &mut tx, + &append.identity, + append.expected_version, + event, + ) + .await?; + } + } + + for write in batch.snapshots { + match write { + AsyncSnapshotWrite::Save { identity, record } => { + save_snapshot_in_tx(&mut tx, &identity, record).await?; + } + } + } + + tx.commit() + .await + .map_err(|err| repository_storage_error("commit transaction", err))?; + + for stream in batch.streams { + stream.entity.mark_committed(); + } + + Ok(()) + } + } +} + +impl AsyncSnapshotStore for PostgresRepository { + fn get_snapshot_async<'a>( + &'a self, + identity: &'a StreamIdentity, + ) -> impl Future, RepositoryError>> + Send + 'a { + async move { + let row = sqlx::query( + r#" + SELECT aggregate_id, version, data + FROM aggregate_snapshots + WHERE aggregate_type = $1 AND aggregate_id = $2 + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .fetch_optional(&self.pool) + .await + .map_err(|err| repository_storage_error("load snapshot", err))?; + + let Some(row) = row else { + return Ok(None); + }; + + Ok(Some(snapshot_from_row(row)?)) + } + } + + fn save_snapshot_async<'a>( + &'a self, + identity: &'a StreamIdentity, + record: SnapshotRecord, + ) -> impl Future> + Send + 'a { + async move { + let mut tx = self + .pool + .begin() + .await + .map_err(|err| repository_storage_error("begin snapshot transaction", err))?; + save_snapshot_in_tx(&mut tx, identity, record).await?; + tx.commit() + .await + .map_err(|err| repository_storage_error("commit snapshot transaction", err))?; + Ok(()) + } + } + + fn delete_snapshot_async<'a>( + &'a self, + identity: &'a StreamIdentity, + ) -> impl Future> + Send + 'a { + async move { + let result = sqlx::query( + r#" + DELETE FROM aggregate_snapshots + WHERE aggregate_type = $1 AND aggregate_id = $2 + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .execute(&self.pool) + .await + .map_err(|err| repository_storage_error("delete snapshot", err))?; + + Ok(result.rows_affected() > 0) + } + } +} + +fn reject_read_model_plans(batch: &AsyncCommitBatch<'_>) -> Result<(), RepositoryError> { + if batch.read_model_plans.iter().any(|plan| !plan.is_empty()) { + return Err(RepositoryError::Model( + "PostgresRepository first pass does not persist read-model write plans".into(), + )); + } + Ok(()) +} + +async fn stream_version_in_tx( + tx: &mut Transaction<'_, Postgres>, + identity: &StreamIdentity, +) -> Result { + let row = sqlx::query( + r#" + SELECT MAX(sequence) AS version + FROM aggregate_events + WHERE aggregate_type = $1 AND aggregate_id = $2 + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .fetch_one(&mut **tx) + .await + .map_err(|err| repository_storage_error("load stream version", err))?; + + let version: Option = row + .try_get("version") + .map_err(|err| repository_storage_error("decode stream version row", err))?; + version + .map(|value| sqlx_repository_u64_from_i64(POSTGRES_BACKEND, value, "sequence")) + .unwrap_or(Ok(0)) +} + +async fn stream_version_pool( + pool: &PgPool, + identity: &StreamIdentity, +) -> Result { + let row = sqlx::query( + r#" + SELECT MAX(sequence) AS version + FROM aggregate_events + WHERE aggregate_type = $1 AND aggregate_id = $2 + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .fetch_one(pool) + .await + .map_err(|err| repository_storage_error("load stream version", err))?; + + let version: Option = row + .try_get("version") + .map_err(|err| repository_storage_error("decode stream version row", err))?; + version + .map(|value| sqlx_repository_u64_from_i64(POSTGRES_BACKEND, value, "sequence")) + .unwrap_or(Ok(0)) +} + +async fn insert_event_in_tx( + pool: &PgPool, + tx: &mut Transaction<'_, Postgres>, + identity: &StreamIdentity, + expected_version: u64, + event: &EventRecord, +) -> Result<(), RepositoryError> { + let metadata = serialize_event_metadata(&event.metadata)?; + + let result = sqlx::query( + r#" + INSERT INTO aggregate_events ( + aggregate_type, + aggregate_id, + sequence, + event_name, + event_version, + payload, + payload_codec, + payload_codec_version, + metadata, + recorded_at + ) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9::jsonb, to_timestamp($10)) + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .bind(sqlx_repository_i64_from_u64( + POSTGRES_BACKEND, + event.sequence, + "sequence", + BIGINT_STORAGE, + )?) + .bind(&event.event_name) + .bind(sqlx_repository_i32_from_u64( + POSTGRES_BACKEND, + event.event_version, + "event_version", + INTEGER_STORAGE, + )?) + .bind(&event.payload) + .bind(&event.payload_codec) + .bind(i32::from(event.payload_codec_version)) + .bind(metadata) + .bind(system_time_to_epoch_secs(event.timestamp)?) + .execute(&mut **tx) + .await; + + match result { + Ok(_) => Ok(()), + Err(err) if is_postgres_unique_violation(&err) => { + let actual = stream_version_pool(pool, identity) + .await + .unwrap_or_default(); + Err(RepositoryError::ConcurrentWrite { + id: identity.to_string(), + expected: expected_version, + actual, + }) + } + Err(err) => Err(repository_storage_error("insert event", err)), + } +} + +fn event_from_row(row: PgRow) -> Result { + let payload_codec: String = row + .try_get("payload_codec") + .map_err(|err| repository_storage_error("decode payload codec row", err))?; + let payload_codec_version = sqlx_repository_u16_from_i32( + POSTGRES_BACKEND, + row.try_get("payload_codec_version") + .map_err(|err| repository_storage_error("decode payload codec version row", err))?, + "payload_codec_version", + )?; + let metadata_json: String = row + .try_get("metadata") + .map_err(|err| repository_storage_error("decode metadata row", err))?; + let metadata = deserialize_event_metadata(&metadata_json)?; + let event = EventRecord { + event_name: row + .try_get("event_name") + .map_err(|err| repository_storage_error("decode event name row", err))?, + payload_codec, + payload_codec_version, + payload: row + .try_get("payload") + .map_err(|err| repository_storage_error("decode payload row", err))?, + event_version: sqlx_repository_u64_from_i32( + POSTGRES_BACKEND, + row.try_get("event_version") + .map_err(|err| repository_storage_error("decode event version row", err))?, + "event_version", + )?, + sequence: sqlx_repository_u64_from_i64( + POSTGRES_BACKEND, + row.try_get("sequence") + .map_err(|err| repository_storage_error("decode sequence row", err))?, + "sequence", + )?, + timestamp: system_time_from_epoch_secs( + row.try_get("recorded_at_epoch") + .map_err(|err| repository_storage_error("decode recorded_at row", err))?, + )?, + metadata, + }; + validate_supported_event_codec(&event)?; + Ok(event) +} + +async fn save_snapshot_in_tx( + tx: &mut Transaction<'_, Postgres>, + identity: &StreamIdentity, + record: SnapshotRecord, +) -> Result<(), RepositoryError> { + validate_snapshot_identity(identity, &record)?; + + sqlx::query( + r#" + INSERT INTO aggregate_snapshots (aggregate_type, aggregate_id, version, data) + VALUES ($1, $2, $3, $4) + ON CONFLICT(aggregate_type, aggregate_id) DO UPDATE SET + version = excluded.version, + data = excluded.data, + updated_at = now() + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .bind(sqlx_repository_i64_from_u64( + POSTGRES_BACKEND, + record.version, + "snapshot version", + BIGINT_STORAGE, + )?) + .bind(record.data) + .execute(&mut **tx) + .await + .map_err(|err| repository_storage_error("save snapshot", err))?; + + Ok(()) +} + +fn snapshot_from_row(row: PgRow) -> Result { + Ok(SnapshotRecord { + aggregate_id: row + .try_get("aggregate_id") + .map_err(|err| repository_storage_error("decode snapshot aggregate id row", err))?, + version: sqlx_repository_u64_from_i64( + POSTGRES_BACKEND, + row.try_get("version") + .map_err(|err| repository_storage_error("decode snapshot version row", err))?, + "snapshot version", + )?, + data: row + .try_get("data") + .map_err(|err| repository_storage_error("decode snapshot data row", err))?, + }) +} + +fn system_time_to_epoch_secs(timestamp: SystemTime) -> Result { + let duration = timestamp.duration_since(UNIX_EPOCH).map_err(|err| { + RepositoryError::Model(format!( + "event timestamp before UNIX epoch cannot be stored in postgres: {err}" + )) + })?; + Ok(duration.as_secs_f64()) +} + +fn system_time_from_epoch_secs(value: f64) -> Result { + if !value.is_finite() || value < 0.0 { + return Err(RepositoryError::Model(format!( + "postgres recorded_at epoch value {value} is invalid" + ))); + } + Ok(UNIX_EPOCH + Duration::from_secs_f64(value)) +} + +fn repository_storage_error(operation: &str, err: sqlx::Error) -> RepositoryError { + sqlx_repo::repository_storage_error(POSTGRES_BACKEND, operation, err) +} diff --git a/src/sqlite_repo/mod.rs b/src/sqlite_repo/mod.rs new file mode 100644 index 00000000..e2f9c704 --- /dev/null +++ b/src/sqlite_repo/mod.rs @@ -0,0 +1,966 @@ +//! SQLite-backed async repository and transactional document stores. +//! +//! This adapter is a local SQL persistence backend for the async repository +//! boundary. It is feature-gated behind `sqlite` and is intentionally async-only. + +#![expect( + clippy::manual_async_fn, + reason = "async trait impls return impl Future + Send to preserve public Send bounds" +)] + +use std::collections::HashSet; +use std::future::Future; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use sqlx::sqlite::SqlitePoolOptions; +use sqlx::{Row, Sqlite, SqlitePool, Transaction}; + +use crate::entity::{Entity, EventRecord}; +use crate::read_model::{ + ProcessedMessageMark, ReadModel, ReadModelAdapterCapabilities, ReadModelCommitOutcome, + ReadModelError, ReadModelMutation, ReadModelWritePlan, Versioned, +}; +use crate::repository::{ + AsyncCommitBatch, AsyncGetStream, AsyncReadModelSessionStore, AsyncReadModelStore, + AsyncSnapshotStore, AsyncSnapshotWrite, AsyncTransactionalCommit, PreparedEventAppend, + RepositoryError, StreamIdentity, +}; +use crate::snapshot::SnapshotRecord; +use crate::sqlx_repo::{ + self, deserialize_event_metadata, is_sqlite_unique_constraint, + read_model_i64_from_u64 as sqlx_read_model_i64_from_u64, + read_model_u64_from_i64 as sqlx_read_model_u64_from_i64, reject_duplicate_streams, + repository_i64_from_u64 as sqlx_repository_i64_from_u64, + repository_u16_from_i64 as sqlx_repository_u16_from_i64, + repository_u64_from_i64 as sqlx_repository_u64_from_i64, serialize_event_metadata, + validate_entity_id_matches_identity, validate_prepared_appends, validate_snapshot_identity, + validate_supported_event_codec, +}; + +const SQLITE_SCHEMA: &str = include_str!("../../migrations/sqlite/0001_initial.sql"); +const SQLITE_BACKEND: &str = "sqlite"; +const SIGNED_INTEGER_STORAGE: &str = "signed integer storage"; + +/// SQLite-backed async repository. +#[derive(Clone)] +pub struct SqliteRepository { + pool: SqlitePool, +} + +impl SqliteRepository { + /// Create a repository from an existing migrated pool. + pub fn new(pool: SqlitePool) -> Self { + Self { pool } + } + + /// Open a SQLite pool without applying migrations. + pub async fn connect(database_url: &str) -> Result { + let pool = SqlitePoolOptions::new() + .max_connections(default_pool_size(database_url)) + .connect(database_url) + .await + .map_err(|err| repository_storage_error("connect", err))?; + Ok(Self::new(pool)) + } + + /// Open a SQLite pool and apply the explicit SQLite migrations. + pub async fn connect_and_migrate(database_url: &str) -> Result { + let repo = Self::connect(database_url).await?; + repo.migrate().await?; + Ok(repo) + } + + /// Apply SQLite migrations to this repository's pool. + pub async fn migrate(&self) -> Result<(), RepositoryError> { + Self::migrate_pool(&self.pool).await + } + + /// Apply SQLite migrations to an existing pool. + pub async fn migrate_pool(pool: &SqlitePool) -> Result<(), RepositoryError> { + for statement in SQLITE_SCHEMA.split(';') { + let statement = statement.trim(); + if statement.is_empty() { + continue; + } + sqlx::query(statement) + .execute(pool) + .await + .map_err(|err| repository_storage_error("migrate", err))?; + } + Ok(()) + } + + /// Access the underlying SQLx pool for application-specific setup or tests. + pub fn pool(&self) -> &SqlitePool { + &self.pool + } +} + +impl AsyncGetStream for SqliteRepository { + fn get_stream<'a>( + &'a self, + identity: &'a StreamIdentity, + ) -> impl Future, RepositoryError>> + Send + 'a { + async move { + let rows = sqlx::query( + r#" + SELECT event_name, event_version, payload, payload_codec, + payload_codec_version, metadata, sequence, recorded_at + FROM aggregate_events + WHERE aggregate_type = ? AND aggregate_id = ? + ORDER BY sequence ASC + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .fetch_all(&self.pool) + .await + .map_err(|err| repository_storage_error("load stream", err))?; + + if rows.is_empty() { + return Ok(None); + } + + let mut events = Vec::with_capacity(rows.len()); + for row in rows { + events.push(event_from_row(row)?); + } + + let mut entity = Entity::new(); + entity.set_id(identity.aggregate_id()); + entity.load_from_history(events); + Ok(Some(entity)) + } + } + + fn get_streams<'a>( + &'a self, + identities: &'a [StreamIdentity], + ) -> impl Future, RepositoryError>> + Send + 'a { + async move { + let mut entities = Vec::with_capacity(identities.len()); + for identity in identities { + if let Some(entity) = self.get_stream(identity).await? { + entities.push(entity); + } + } + Ok(entities) + } + } +} + +impl AsyncTransactionalCommit for SqliteRepository { + fn commit_batch_async<'a>( + &'a self, + batch: AsyncCommitBatch<'a>, + ) -> impl Future> + Send + 'a { + async move { + reject_duplicate_streams(&batch.streams)?; + validate_entity_id_matches_identity(&batch.streams)?; + + let prepared = batch + .streams + .iter() + .map(PreparedEventAppend::from_stream_write) + .collect::>(); + validate_prepared_appends(&prepared)?; + + for plan in &batch.read_model_plans { + validate_document_write_plan(plan)?; + } + + let mut tx = self + .pool + .begin() + .await + .map_err(|err| repository_storage_error("begin commit transaction", err))?; + + for append in &prepared { + let actual = stream_version_in_tx(&mut tx, &append.identity).await?; + if actual != append.expected_version { + return Err(RepositoryError::ConcurrentWrite { + id: append.identity.to_string(), + expected: append.expected_version, + actual, + }); + } + } + + for append in &prepared { + for event in &append.events { + insert_event_in_tx(&mut tx, &append.identity, append.expected_version, event) + .await?; + } + } + + for plan in batch.read_model_plans { + let outcome = apply_document_write_plan_in_tx(&mut tx, plan).await?; + if let Some(mark) = outcome.duplicate_message() { + return Err(RepositoryError::Model(format!( + "processed message already handled by consumer `{}`: `{}`", + mark.consumer_name, mark.message_id + ))); + } + } + + for write in batch.snapshots { + match write { + AsyncSnapshotWrite::Save { identity, record } => { + save_snapshot_in_tx(&mut tx, &identity, record).await?; + } + } + } + + tx.commit() + .await + .map_err(|err| repository_storage_error("commit transaction", err))?; + + for stream in batch.streams { + stream.entity.mark_committed(); + } + + Ok(()) + } + } +} + +impl AsyncReadModelStore for SqliteRepository { + fn get_model_async<'a, M>( + &'a self, + id: &'a str, + ) -> impl Future>, ReadModelError>> + Send + 'a + where + M: ReadModel + 'a, + { + async move { self.load_document_model::(id).await } + } + + fn get_by_primary_key_async<'a, M>( + &'a self, + id: &'a str, + ) -> impl Future>, ReadModelError>> + Send + 'a + where + M: ReadModel + 'a, + { + async move { self.load_document_model::(id).await } + } + + fn upsert_async<'a, M>( + &'a self, + model: &'a M, + ) -> impl Future, ReadModelError>> + Send + 'a + where + M: ReadModel + 'a, + { + async move { + let bytes = + serde_json::to_vec(model).map_err(|err| ReadModelError::Serde(err.to_string()))?; + let mut tx = begin_read_model_tx(&self.pool).await?; + let version = upsert_document_in_tx(&mut tx, M::COLLECTION, model.id(), bytes).await?; + commit_read_model_tx(tx).await?; + Ok(Versioned { + data: model.clone(), + version, + }) + } + } + + fn insert_async<'a, M>( + &'a self, + model: &'a M, + ) -> impl Future, ReadModelError>> + Send + 'a + where + M: ReadModel + 'a, + { + async move { + let bytes = + serde_json::to_vec(model).map_err(|err| ReadModelError::Serde(err.to_string()))?; + let mut tx = begin_read_model_tx(&self.pool).await?; + let existing = document_version_in_tx(&mut tx, M::COLLECTION, model.id()).await?; + if let Some(actual) = existing { + return Err(ReadModelError::ConcurrencyConflict { + collection: M::COLLECTION.to_string(), + id: model.id().to_string(), + expected: 0, + actual, + }); + } + insert_document_in_tx(&mut tx, M::COLLECTION, model.id(), bytes, 1).await?; + commit_read_model_tx(tx).await?; + Ok(Versioned { + data: model.clone(), + version: 1, + }) + } + } + + fn update_async<'a, M>( + &'a self, + model: &'a M, + expected_version: u64, + ) -> impl Future, ReadModelError>> + Send + 'a + where + M: ReadModel + 'a, + { + async move { + let bytes = + serde_json::to_vec(model).map_err(|err| ReadModelError::Serde(err.to_string()))?; + let mut tx = begin_read_model_tx(&self.pool).await?; + let actual = document_version_in_tx(&mut tx, M::COLLECTION, model.id()) + .await? + .ok_or_else(|| ReadModelError::NotFound { + collection: M::COLLECTION.to_string(), + id: model.id().to_string(), + })?; + if actual != expected_version { + return Err(ReadModelError::ConcurrencyConflict { + collection: M::COLLECTION.to_string(), + id: model.id().to_string(), + expected: expected_version, + actual, + }); + } + let new_version = next_document_version(M::COLLECTION, model.id(), Some(actual))?; + update_document_in_tx(&mut tx, M::COLLECTION, model.id(), bytes, new_version).await?; + commit_read_model_tx(tx).await?; + Ok(Versioned { + data: model.clone(), + version: new_version, + }) + } + } + + fn delete_async<'a, M>( + &'a self, + id: &'a str, + ) -> impl Future> + Send + 'a + where + M: ReadModel + 'a, + { + async move { + let result = sqlx::query( + r#" + DELETE FROM transactional_read_models + WHERE collection = ? AND id = ? + "#, + ) + .bind(M::COLLECTION) + .bind(id) + .execute(&self.pool) + .await + .map_err(|err| read_model_storage_error("delete document", err))?; + + Ok(result.rows_affected() > 0) + } + } +} + +impl SqliteRepository { + async fn load_document_model( + &self, + id: &str, + ) -> Result>, ReadModelError> { + let row = sqlx::query( + r#" + SELECT payload, version + FROM transactional_read_models + WHERE collection = ? AND id = ? + "#, + ) + .bind(M::COLLECTION) + .bind(id) + .fetch_optional(&self.pool) + .await + .map_err(|err| read_model_storage_error("load document", err))?; + + let Some(row) = row else { + return Ok(None); + }; + + let payload: Vec = row + .try_get("payload") + .map_err(|err| read_model_storage_error("decode document payload row", err))?; + let version = sqlx_read_model_u64_from_i64( + SQLITE_BACKEND, + row.try_get("version") + .map_err(|err| read_model_storage_error("decode document version row", err))?, + "version", + )?; + let data = serde_json::from_slice(&payload) + .map_err(|err| ReadModelError::Serde(err.to_string()))?; + + Ok(Some(Versioned { data, version })) + } +} + +impl AsyncReadModelSessionStore for SqliteRepository { + fn read_model_capabilities_async(&self) -> ReadModelAdapterCapabilities { + document_capabilities() + } + + fn commit_write_plan_async( + &self, + plan: ReadModelWritePlan, + ) -> impl Future> + Send + '_ { + async move { + validate_document_write_plan(&plan)?; + let mut tx = begin_read_model_tx(&self.pool).await?; + let outcome = apply_document_write_plan_in_tx(&mut tx, plan).await?; + if outcome.was_applied() { + commit_read_model_tx(tx).await?; + } + Ok(outcome) + } + } + + fn is_processed_async<'a>( + &'a self, + consumer_name: &'a str, + message_id: &'a str, + ) -> impl Future> + Send + 'a { + async move { processed_message_exists_pool(&self.pool, consumer_name, message_id).await } + } +} + +impl AsyncSnapshotStore for SqliteRepository { + fn get_snapshot_async<'a>( + &'a self, + identity: &'a StreamIdentity, + ) -> impl Future, RepositoryError>> + Send + 'a { + async move { + let row = sqlx::query( + r#" + SELECT aggregate_id, version, data + FROM aggregate_snapshots + WHERE aggregate_type = ? AND aggregate_id = ? + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .fetch_optional(&self.pool) + .await + .map_err(|err| repository_storage_error("load snapshot", err))?; + + let Some(row) = row else { + return Ok(None); + }; + + Ok(Some(snapshot_from_row(row)?)) + } + } + + fn save_snapshot_async<'a>( + &'a self, + identity: &'a StreamIdentity, + record: SnapshotRecord, + ) -> impl Future> + Send + 'a { + async move { + let mut tx = self + .pool + .begin() + .await + .map_err(|err| repository_storage_error("begin snapshot transaction", err))?; + save_snapshot_in_tx(&mut tx, identity, record).await?; + tx.commit() + .await + .map_err(|err| repository_storage_error("commit snapshot transaction", err))?; + Ok(()) + } + } + + fn delete_snapshot_async<'a>( + &'a self, + identity: &'a StreamIdentity, + ) -> impl Future> + Send + 'a { + async move { + let result = sqlx::query( + r#" + DELETE FROM aggregate_snapshots + WHERE aggregate_type = ? AND aggregate_id = ? + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .execute(&self.pool) + .await + .map_err(|err| repository_storage_error("delete snapshot", err))?; + + Ok(result.rows_affected() > 0) + } + } +} + +fn default_pool_size(database_url: &str) -> u32 { + if database_url.contains(":memory:") { + 1 + } else { + 5 + } +} + +async fn stream_version_in_tx( + tx: &mut Transaction<'_, Sqlite>, + identity: &StreamIdentity, +) -> Result { + let row = sqlx::query( + r#" + SELECT MAX(sequence) AS version + FROM aggregate_events + WHERE aggregate_type = ? AND aggregate_id = ? + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .fetch_one(&mut **tx) + .await + .map_err(|err| repository_storage_error("load stream version", err))?; + + let version: Option = row + .try_get("version") + .map_err(|err| repository_storage_error("decode stream version row", err))?; + version + .map(|value| sqlx_repository_u64_from_i64(SQLITE_BACKEND, value, "sequence")) + .unwrap_or(Ok(0)) +} + +async fn insert_event_in_tx( + tx: &mut Transaction<'_, Sqlite>, + identity: &StreamIdentity, + expected_version: u64, + event: &EventRecord, +) -> Result<(), RepositoryError> { + let metadata = serialize_event_metadata(&event.metadata)?; + + let result = sqlx::query( + r#" + INSERT INTO aggregate_events ( + aggregate_type, + aggregate_id, + sequence, + event_name, + event_version, + payload, + payload_codec, + payload_codec_version, + metadata, + recorded_at + ) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .bind(sqlx_repository_i64_from_u64( + SQLITE_BACKEND, + event.sequence, + "sequence", + SIGNED_INTEGER_STORAGE, + )?) + .bind(&event.event_name) + .bind(sqlx_repository_i64_from_u64( + SQLITE_BACKEND, + event.event_version, + "event_version", + SIGNED_INTEGER_STORAGE, + )?) + .bind(&event.payload) + .bind(&event.payload_codec) + .bind(i64::from(event.payload_codec_version)) + .bind(metadata) + .bind(system_time_to_storage(event.timestamp)?) + .execute(&mut **tx) + .await; + + match result { + Ok(_) => Ok(()), + Err(err) if is_sqlite_unique_constraint(&err) => { + let actual = stream_version_in_tx(tx, identity).await?; + Err(RepositoryError::ConcurrentWrite { + id: identity.to_string(), + expected: expected_version, + actual, + }) + } + Err(err) => Err(repository_storage_error("insert event", err)), + } +} + +fn event_from_row(row: sqlx::sqlite::SqliteRow) -> Result { + let payload_codec: String = row + .try_get("payload_codec") + .map_err(|err| repository_storage_error("decode payload codec row", err))?; + let payload_codec_version = sqlx_repository_u16_from_i64( + SQLITE_BACKEND, + row.try_get("payload_codec_version") + .map_err(|err| repository_storage_error("decode payload codec version row", err))?, + "payload_codec_version", + )?; + let metadata_json: String = row + .try_get("metadata") + .map_err(|err| repository_storage_error("decode metadata row", err))?; + let metadata = deserialize_event_metadata(&metadata_json)?; + let event = EventRecord { + event_name: row + .try_get("event_name") + .map_err(|err| repository_storage_error("decode event name row", err))?, + payload_codec, + payload_codec_version, + payload: row + .try_get("payload") + .map_err(|err| repository_storage_error("decode payload row", err))?, + event_version: sqlx_repository_u64_from_i64( + SQLITE_BACKEND, + row.try_get("event_version") + .map_err(|err| repository_storage_error("decode event version row", err))?, + "event_version", + )?, + sequence: sqlx_repository_u64_from_i64( + SQLITE_BACKEND, + row.try_get("sequence") + .map_err(|err| repository_storage_error("decode sequence row", err))?, + "sequence", + )?, + timestamp: system_time_from_storage( + row.try_get::("recorded_at") + .map_err(|err| repository_storage_error("decode recorded_at row", err))? + .as_str(), + ), + metadata, + }; + validate_supported_event_codec(&event)?; + Ok(event) +} + +fn document_capabilities() -> ReadModelAdapterCapabilities { + ReadModelAdapterCapabilities { + relational_rows: false, + document_rows: true, + sparse_patches: false, + deletes: false, + processed_messages: true, + } +} + +fn validate_document_write_plan(plan: &ReadModelWritePlan) -> Result<(), ReadModelError> { + for mutation in &plan.mutations { + if !matches!(mutation, ReadModelMutation::Document(_)) { + return Err(ReadModelError::Metadata( + "SqliteRepository currently supports only document read-model mutations".into(), + )); + } + } + plan.validate_for(&document_capabilities()) +} + +async fn begin_read_model_tx(pool: &SqlitePool) -> Result, ReadModelError> { + pool.begin() + .await + .map_err(|err| read_model_storage_error("begin transaction", err)) +} + +async fn commit_read_model_tx(tx: Transaction<'_, Sqlite>) -> Result<(), ReadModelError> { + tx.commit() + .await + .map_err(|err| read_model_storage_error("commit transaction", err)) +} + +async fn apply_document_write_plan_in_tx( + tx: &mut Transaction<'_, Sqlite>, + plan: ReadModelWritePlan, +) -> Result { + validate_document_write_plan(&plan)?; + + let mut marks_in_plan = HashSet::with_capacity(plan.processed_messages.len()); + for mark in &plan.processed_messages { + let key = processed_message_key(mark); + if !marks_in_plan.insert(key) || processed_message_exists_in_tx(tx, mark).await? { + return Ok(ReadModelCommitOutcome::skipped_duplicate(mark.clone())); + } + } + + for mutation in plan.mutations { + match mutation { + ReadModelMutation::Document(mutation) => { + upsert_document_in_tx(tx, &mutation.collection, &mutation.id, mutation.bytes) + .await?; + } + _ => { + return Err(ReadModelError::Metadata( + "SqliteRepository currently supports only document read-model mutations".into(), + )); + } + } + } + + for mark in plan.processed_messages { + let result = insert_processed_message_in_tx(tx, &mark).await; + if let Err(err) = result { + if is_sqlite_unique_constraint(&err) { + return Ok(ReadModelCommitOutcome::skipped_duplicate(mark)); + } + return Err(read_model_storage_error("insert processed message", err)); + } + } + + Ok(ReadModelCommitOutcome::applied()) +} + +async fn upsert_document_in_tx( + tx: &mut Transaction<'_, Sqlite>, + collection: &str, + id: &str, + bytes: Vec, +) -> Result { + let current = document_version_in_tx(tx, collection, id).await?; + let new_version = next_document_version(collection, id, current)?; + match current { + Some(_) => update_document_in_tx(tx, collection, id, bytes, new_version).await?, + None => insert_document_in_tx(tx, collection, id, bytes, new_version).await?, + } + Ok(new_version) +} + +async fn document_version_in_tx( + tx: &mut Transaction<'_, Sqlite>, + collection: &str, + id: &str, +) -> Result, ReadModelError> { + let row = sqlx::query( + r#" + SELECT version + FROM transactional_read_models + WHERE collection = ? AND id = ? + "#, + ) + .bind(collection) + .bind(id) + .fetch_optional(&mut **tx) + .await + .map_err(|err| read_model_storage_error("load document version", err))?; + + row.map(|row| { + sqlx_read_model_u64_from_i64( + SQLITE_BACKEND, + row.try_get("version") + .map_err(|err| read_model_storage_error("decode document version row", err))?, + "version", + ) + }) + .transpose() +} + +async fn insert_document_in_tx( + tx: &mut Transaction<'_, Sqlite>, + collection: &str, + id: &str, + bytes: Vec, + version: u64, +) -> Result<(), ReadModelError> { + sqlx::query( + r#" + INSERT INTO transactional_read_models (collection, id, version, payload) + VALUES (?, ?, ?, ?) + "#, + ) + .bind(collection) + .bind(id) + .bind(sqlx_read_model_i64_from_u64( + SQLITE_BACKEND, + version, + "version", + SIGNED_INTEGER_STORAGE, + )?) + .bind(bytes) + .execute(&mut **tx) + .await + .map_err(|err| read_model_storage_error("insert document", err))?; + + Ok(()) +} + +async fn update_document_in_tx( + tx: &mut Transaction<'_, Sqlite>, + collection: &str, + id: &str, + bytes: Vec, + version: u64, +) -> Result<(), ReadModelError> { + sqlx::query( + r#" + UPDATE transactional_read_models + SET version = ?, payload = ?, updated_at = CURRENT_TIMESTAMP + WHERE collection = ? AND id = ? + "#, + ) + .bind(sqlx_read_model_i64_from_u64( + SQLITE_BACKEND, + version, + "version", + SIGNED_INTEGER_STORAGE, + )?) + .bind(bytes) + .bind(collection) + .bind(id) + .execute(&mut **tx) + .await + .map_err(|err| read_model_storage_error("update document", err))?; + + Ok(()) +} + +async fn processed_message_exists_pool( + pool: &SqlitePool, + consumer_name: &str, + message_id: &str, +) -> Result { + let row = sqlx::query( + r#" + SELECT 1 + FROM read_model_processed_messages + WHERE consumer_name = ? AND message_id = ? + "#, + ) + .bind(consumer_name) + .bind(message_id) + .fetch_optional(pool) + .await + .map_err(|err| read_model_storage_error("load processed message", err))?; + Ok(row.is_some()) +} + +async fn processed_message_exists_in_tx( + tx: &mut Transaction<'_, Sqlite>, + mark: &ProcessedMessageMark, +) -> Result { + let row = sqlx::query( + r#" + SELECT 1 + FROM read_model_processed_messages + WHERE consumer_name = ? AND message_id = ? + "#, + ) + .bind(&mark.consumer_name) + .bind(&mark.message_id) + .fetch_optional(&mut **tx) + .await + .map_err(|err| read_model_storage_error("load processed message", err))?; + Ok(row.is_some()) +} + +async fn insert_processed_message_in_tx( + tx: &mut Transaction<'_, Sqlite>, + mark: &ProcessedMessageMark, +) -> Result<(), sqlx::Error> { + sqlx::query( + r#" + INSERT INTO read_model_processed_messages (consumer_name, message_id) + VALUES (?, ?) + "#, + ) + .bind(&mark.consumer_name) + .bind(&mark.message_id) + .execute(&mut **tx) + .await?; + Ok(()) +} + +fn processed_message_key(mark: &ProcessedMessageMark) -> (String, String) { + (mark.consumer_name.clone(), mark.message_id.clone()) +} + +async fn save_snapshot_in_tx( + tx: &mut Transaction<'_, Sqlite>, + identity: &StreamIdentity, + record: SnapshotRecord, +) -> Result<(), RepositoryError> { + validate_snapshot_identity(identity, &record)?; + + sqlx::query( + r#" + INSERT INTO aggregate_snapshots (aggregate_type, aggregate_id, version, data) + VALUES (?, ?, ?, ?) + ON CONFLICT(aggregate_type, aggregate_id) DO UPDATE SET + version = excluded.version, + data = excluded.data, + updated_at = CURRENT_TIMESTAMP + "#, + ) + .bind(identity.aggregate_type()) + .bind(identity.aggregate_id()) + .bind(sqlx_repository_i64_from_u64( + SQLITE_BACKEND, + record.version, + "snapshot version", + SIGNED_INTEGER_STORAGE, + )?) + .bind(record.data) + .execute(&mut **tx) + .await + .map_err(|err| repository_storage_error("save snapshot", err))?; + + Ok(()) +} + +fn snapshot_from_row(row: sqlx::sqlite::SqliteRow) -> Result { + Ok(SnapshotRecord { + aggregate_id: row + .try_get("aggregate_id") + .map_err(|err| repository_storage_error("decode snapshot aggregate id row", err))?, + version: sqlx_repository_u64_from_i64( + SQLITE_BACKEND, + row.try_get("version") + .map_err(|err| repository_storage_error("decode snapshot version row", err))?, + "snapshot version", + )?, + data: row + .try_get("data") + .map_err(|err| repository_storage_error("decode snapshot data row", err))?, + }) +} + +fn next_document_version( + collection: &str, + id: &str, + current_version: Option, +) -> Result { + match current_version { + Some(version) => version.checked_add(1).ok_or_else(|| { + ReadModelError::Storage(format!("read model version overflow for {collection}:{id}")) + }), + None => Ok(1), + } +} + +fn system_time_to_storage(timestamp: SystemTime) -> Result { + let duration = timestamp.duration_since(UNIX_EPOCH).map_err(|err| { + RepositoryError::Model(format!( + "event timestamp before UNIX epoch cannot be stored in sqlite: {err}" + )) + })?; + Ok(format!( + "{}.{:09}", + duration.as_secs(), + duration.subsec_nanos() + )) +} + +fn system_time_from_storage(value: &str) -> SystemTime { + let Some((secs, nanos)) = value.split_once('.') else { + return UNIX_EPOCH; + }; + let Ok(secs) = secs.parse::() else { + return UNIX_EPOCH; + }; + let Ok(nanos) = nanos.parse::() else { + return UNIX_EPOCH; + }; + UNIX_EPOCH + Duration::new(secs, nanos) +} + +fn repository_storage_error(operation: &str, err: sqlx::Error) -> RepositoryError { + sqlx_repo::repository_storage_error(SQLITE_BACKEND, operation, err) +} + +fn read_model_storage_error(operation: &str, err: sqlx::Error) -> ReadModelError { + sqlx_repo::read_model_storage_error(SQLITE_BACKEND, operation, err) +} diff --git a/src/sqlx_repo/mod.rs b/src/sqlx_repo/mod.rs new file mode 100644 index 00000000..258670b5 --- /dev/null +++ b/src/sqlx_repo/mod.rs @@ -0,0 +1,221 @@ +use std::collections::{HashMap, HashSet}; + +use crate::entity::{ + EventRecord, EventRecordError, BITCODE_PAYLOAD_CODEC, BITCODE_PAYLOAD_CODEC_VERSION, +}; +#[cfg(feature = "sqlite")] +use crate::read_model::ReadModelError; +use crate::repository::{AsyncStreamWrite, PreparedEventAppend, RepositoryError, StreamIdentity}; +use crate::snapshot::SnapshotRecord; + +pub(crate) fn reject_duplicate_streams( + streams: &[AsyncStreamWrite<'_>], +) -> Result<(), RepositoryError> { + let mut seen = HashSet::with_capacity(streams.len()); + for stream in streams { + let key = stream.identity.storage_key(); + if !seen.insert(key) { + return Err(RepositoryError::DuplicateStreamInBatch { + id: stream.identity.to_string(), + }); + } + } + Ok(()) +} + +pub(crate) fn validate_entity_id_matches_identity( + streams: &[AsyncStreamWrite<'_>], +) -> Result<(), RepositoryError> { + for stream in streams { + if stream.entity.id() != stream.identity.aggregate_id() { + return Err(RepositoryError::Model(format!( + "stream identity `{}` does not match entity id `{}`", + stream.identity, + stream.entity.id() + ))); + } + } + Ok(()) +} + +pub(crate) fn validate_prepared_appends( + appends: &[PreparedEventAppend], +) -> Result<(), RepositoryError> { + for append in appends { + for (offset, event) in append.events.iter().enumerate() { + validate_supported_event_codec(event)?; + let expected_sequence = append.expected_version + offset as u64 + 1; + if event.sequence != expected_sequence { + return Err(RepositoryError::Model(format!( + "event `{}` for stream `{}` has sequence {}, expected {}", + event.event_name, append.identity, event.sequence, expected_sequence + ))); + } + } + } + Ok(()) +} + +pub(crate) fn validate_supported_event_codec(event: &EventRecord) -> Result<(), RepositoryError> { + if event.payload_codec != BITCODE_PAYLOAD_CODEC + || event.payload_codec_version != BITCODE_PAYLOAD_CODEC_VERSION + { + return Err(EventRecordError::unsupported_codec( + &event.payload_codec, + event.payload_codec_version, + ) + .into()); + } + Ok(()) +} + +pub(crate) fn serialize_event_metadata( + metadata: &HashMap, +) -> Result { + serde_json::to_string(metadata) + .map_err(|err| RepositoryError::Model(format!("serialize event metadata: {err}"))) +} + +pub(crate) fn deserialize_event_metadata( + metadata_json: &str, +) -> Result, RepositoryError> { + serde_json::from_str(metadata_json) + .map_err(|err| RepositoryError::Model(format!("deserialize event metadata: {err}"))) +} + +pub(crate) fn validate_snapshot_identity( + identity: &StreamIdentity, + record: &SnapshotRecord, +) -> Result<(), RepositoryError> { + if record.aggregate_id != identity.aggregate_id() { + return Err(RepositoryError::Model(format!( + "snapshot aggregate id `{}` does not match stream identity `{}`", + record.aggregate_id, identity + ))); + } + Ok(()) +} + +pub(crate) fn repository_i64_from_u64( + backend: &str, + value: u64, + field: &str, + storage: &str, +) -> Result { + i64::try_from(value).map_err(|_| { + RepositoryError::Model(format!("{backend} {field} value {value} exceeds {storage}")) + }) +} + +#[cfg(feature = "postgres")] +pub(crate) fn repository_i32_from_u64( + backend: &str, + value: u64, + field: &str, + storage: &str, +) -> Result { + i32::try_from(value).map_err(|_| { + RepositoryError::Model(format!("{backend} {field} value {value} exceeds {storage}")) + }) +} + +pub(crate) fn repository_u64_from_i64( + backend: &str, + value: i64, + field: &str, +) -> Result { + u64::try_from(value) + .map_err(|_| RepositoryError::Model(format!("{backend} {field} value {value} is negative"))) +} + +#[cfg(feature = "postgres")] +pub(crate) fn repository_u64_from_i32( + backend: &str, + value: i32, + field: &str, +) -> Result { + u64::try_from(value) + .map_err(|_| RepositoryError::Model(format!("{backend} {field} value {value} is negative"))) +} + +#[cfg(feature = "sqlite")] +pub(crate) fn repository_u16_from_i64( + backend: &str, + value: i64, + field: &str, +) -> Result { + u16::try_from(value) + .map_err(|_| RepositoryError::Model(format!("{backend} {field} value {value} is invalid"))) +} + +#[cfg(feature = "postgres")] +pub(crate) fn repository_u16_from_i32( + backend: &str, + value: i32, + field: &str, +) -> Result { + u16::try_from(value) + .map_err(|_| RepositoryError::Model(format!("{backend} {field} value {value} is invalid"))) +} + +#[cfg(feature = "sqlite")] +pub(crate) fn read_model_i64_from_u64( + backend: &str, + value: u64, + field: &str, + storage: &str, +) -> Result { + i64::try_from(value).map_err(|_| { + ReadModelError::Storage(format!("{backend} {field} value {value} exceeds {storage}")) + }) +} + +#[cfg(feature = "sqlite")] +pub(crate) fn read_model_u64_from_i64( + backend: &str, + value: i64, + field: &str, +) -> Result { + u64::try_from(value).map_err(|_| { + ReadModelError::Storage(format!("{backend} {field} value {value} is negative")) + }) +} + +#[cfg(feature = "sqlite")] +pub(crate) fn is_sqlite_unique_constraint(err: &sqlx::Error) -> bool { + match err { + sqlx::Error::Database(db_err) => { + let message = db_err.message(); + let code = db_err.code().map(|code| code.into_owned()); + message.contains("UNIQUE constraint failed") + || message.contains("PRIMARY KEY") + || matches!(code.as_deref(), Some("1555" | "2067")) + } + _ => false, + } +} + +#[cfg(feature = "postgres")] +pub(crate) fn is_postgres_unique_violation(err: &sqlx::Error) -> bool { + match err { + sqlx::Error::Database(db_err) => db_err.code().as_deref() == Some("23505"), + _ => false, + } +} + +pub(crate) fn repository_storage_error( + backend: &str, + operation: &str, + err: sqlx::Error, +) -> RepositoryError { + RepositoryError::Model(format!("{backend} {operation} failed: {err}")) +} + +#[cfg(feature = "sqlite")] +pub(crate) fn read_model_storage_error( + backend: &str, + operation: &str, + err: sqlx::Error, +) -> ReadModelError { + ReadModelError::Storage(format!("{backend} {operation} failed: {err}")) +} diff --git a/tests/postgres_repository/main.rs b/tests/postgres_repository/main.rs new file mode 100644 index 00000000..ef9bbfad --- /dev/null +++ b/tests/postgres_repository/main.rs @@ -0,0 +1,359 @@ +#![cfg(feature = "postgres")] + +use std::env; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::{SystemTime, UNIX_EPOCH}; + +use serde::{Deserialize, Serialize}; +use sourced_rust::{ + impl_aggregate, Aggregate, AsyncAggregateBuilder, AsyncCommitBatch, AsyncGetStream, + AsyncSnapshotStore, AsyncStreamWrite, AsyncTransactionalCommit, Entity, EventRecord, + PostgresRepository, ReadModel, ReadModelSession, RepositoryError, SnapshotRecord, + StreamIdentity, +}; + +static NEXT_ID: AtomicU64 = AtomicU64::new(1); + +#[derive(Default)] +struct Counter { + entity: Entity, + value: i32, +} + +impl Counter { + fn increment(&mut self, id: &str, by: i32) { + self.entity.set_id(id); + self.entity.digest("Incremented", &by).unwrap(); + self.value += by; + } + + fn replay(&mut self, event: &EventRecord) -> Result<(), String> { + if event.event_name == "Incremented" { + let by = event.decode::().map_err(|err| err.to_string())?; + self.value += by; + } + Ok(()) + } +} + +impl_aggregate!(Counter, entity, replay, aggregate_type = "postgres.counter"); + +#[derive(Default)] +struct CounterProjection { + entity: Entity, +} + +impl CounterProjection { + fn touch(&mut self, id: &str) { + self.entity.set_id(id); + self.entity.digest_empty("Touched").unwrap(); + } + + fn replay(&mut self, _event: &EventRecord) -> Result<(), String> { + Ok(()) + } +} + +impl_aggregate!( + CounterProjection, + entity, + replay, + aggregate_type = "postgres.counter_projection" +); + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +struct CounterView { + id: String, + value: i32, +} + +impl ReadModel for CounterView { + const COLLECTION: &'static str = "postgres_counter_views"; + + fn id(&self) -> &str { + &self.id + } +} + +async fn repository() -> Option { + let Ok(database_url) = env::var("DATABASE_URL") else { + eprintln!("skipping Postgres integration test: DATABASE_URL is not set"); + return None; + }; + + Some( + PostgresRepository::connect_and_migrate(&database_url) + .await + .unwrap(), + ) +} + +fn unique_id(prefix: &str) -> String { + let nanos = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos(); + let id = NEXT_ID.fetch_add(1, Ordering::Relaxed); + format!("{prefix}-{nanos}-{id}") +} + +#[tokio::test] +async fn migration_is_idempotent_and_uses_postgres_column_types() { + let Some(repo) = repository().await else { + return; + }; + repo.migrate().await.unwrap(); + + let rows = sqlx::query( + r#" + SELECT column_name, udt_name + FROM information_schema.columns + WHERE table_name = 'aggregate_events' + AND column_name IN ('payload', 'metadata', 'recorded_at') + "#, + ) + .fetch_all(repo.pool()) + .await + .unwrap(); + + let mut columns = rows + .into_iter() + .map(|row| { + ( + sqlx::Row::try_get::(&row, "column_name").unwrap(), + sqlx::Row::try_get::(&row, "udt_name").unwrap(), + ) + }) + .collect::>(); + columns.sort(); + + assert_eq!( + columns, + vec![ + ("metadata".into(), "jsonb".into()), + ("payload".into(), "bytea".into()), + ("recorded_at".into(), "timestamptz".into()), + ] + ); + + let read_model_table: Option = + sqlx::query_scalar("SELECT to_regclass('public.transactional_read_models')::text") + .fetch_one(repo.pool()) + .await + .unwrap(); + assert!(read_model_table.is_none()); +} + +#[tokio::test] +async fn aggregate_stream_round_trips_with_metadata() { + let Some(repo) = repository().await else { + return; + }; + let counter_repo = repo.clone().async_aggregate::(); + let id = unique_id("counter"); + + let mut counter = Counter::default(); + counter.entity.set_correlation_id("corr-postgres"); + counter.increment(&id, 2); + counter.increment(&id, 3); + + counter_repo.commit(&mut counter).await.unwrap(); + + let loaded = counter_repo.get(&id).await.unwrap().unwrap(); + assert_eq!(loaded.value, 5); + assert_eq!(loaded.entity().events().len(), 2); + assert_eq!(loaded.entity().events()[0].sequence, 1); + assert_eq!(loaded.entity().events()[1].sequence, 2); + assert_eq!( + loaded.entity().events()[0].correlation_id(), + Some("corr-postgres") + ); +} + +#[tokio::test] +async fn optimistic_conflict_rolls_back_other_stream_and_snapshot() { + let Some(repo) = repository().await else { + return; + }; + let counter_repo = repo.clone().async_aggregate::(); + let counter_id = unique_id("conflict"); + let other_id = unique_id("rollback"); + + let mut original = Counter::default(); + original.increment(&counter_id, 1); + counter_repo.commit(&mut original).await.unwrap(); + + let mut stale = counter_repo.get(&counter_id).await.unwrap().unwrap(); + let mut winner = counter_repo.get(&counter_id).await.unwrap().unwrap(); + stale.increment(&counter_id, 10); + winner.increment(&counter_id, 20); + counter_repo.commit(&mut winner).await.unwrap(); + + let mut other = CounterProjection::default(); + other.touch(&other_id); + + let stale_identity = StreamIdentity::new(Counter::aggregate_type(), &counter_id).unwrap(); + let other_identity = + StreamIdentity::new(CounterProjection::aggregate_type(), &other_id).unwrap(); + let err = repo + .commit_batch_async(AsyncCommitBatch { + streams: vec![ + AsyncStreamWrite::new(stale_identity, stale.entity_mut()), + AsyncStreamWrite::new(other_identity.clone(), other.entity_mut()), + ], + read_model_plans: Vec::new(), + snapshots: vec![sourced_rust::AsyncSnapshotWrite::Save { + identity: other_identity.clone(), + record: SnapshotRecord { + aggregate_id: other_id.clone(), + version: 1, + data: vec![1], + }, + }], + }) + .await + .unwrap_err(); + + assert!(matches!(err, RepositoryError::ConcurrentWrite { .. })); + assert!(repo.get_stream(&other_identity).await.unwrap().is_none()); + assert!(repo + .get_snapshot_async(&other_identity) + .await + .unwrap() + .is_none()); + assert_eq!(stale.entity().committed_version(), 1); + assert_eq!(stale.entity().new_events().len(), 1); +} + +#[tokio::test] +async fn duplicate_stream_identity_is_rejected_before_sql_writes() { + let Some(repo) = repository().await else { + return; + }; + let id = unique_id("duplicate"); + let identity = StreamIdentity::new(Counter::aggregate_type(), &id).unwrap(); + let mut first = Entity::with_id(&id); + first.digest_empty("First").unwrap(); + let mut second = Entity::with_id(&id); + second.digest_empty("Second").unwrap(); + + let err = repo + .commit_batch_async(AsyncCommitBatch::new(vec![ + AsyncStreamWrite::new(identity.clone(), &mut first), + AsyncStreamWrite::new(identity.clone(), &mut second), + ])) + .await + .unwrap_err(); + + assert_eq!( + err, + RepositoryError::DuplicateStreamInBatch { + id: format!("{}:{id}", Counter::aggregate_type()) + } + ); + assert!(repo.get_stream(&identity).await.unwrap().is_none()); +} + +#[tokio::test] +async fn read_model_plans_are_rejected_in_first_pass() { + let Some(repo) = repository().await else { + return; + }; + let id = unique_id("read-model"); + let mut entity = Entity::with_id(&id); + entity.digest_empty("Touched").unwrap(); + let identity = StreamIdentity::new(Counter::aggregate_type(), &id).unwrap(); + let mut session = ReadModelSession::new(); + session.document(&CounterView { id, value: 1 }).unwrap(); + + let err = repo + .commit_batch_async(AsyncCommitBatch { + streams: vec![AsyncStreamWrite::new(identity.clone(), &mut entity)], + read_model_plans: vec![session.into_write_plan().unwrap()], + snapshots: Vec::new(), + }) + .await + .unwrap_err(); + + assert!( + matches!(err, RepositoryError::Model(message) if message.contains("does not persist read-model write plans")) + ); + assert!(repo.get_stream(&identity).await.unwrap().is_none()); +} + +#[tokio::test] +async fn snapshots_persist_by_full_stream_identity() { + let Some(repo) = repository().await else { + return; + }; + let id = unique_id("snapshot"); + let counter = StreamIdentity::new("postgres.counter", &id).unwrap(); + let projection = StreamIdentity::new("postgres.counter_projection", &id).unwrap(); + + repo.save_snapshot_async( + &counter, + SnapshotRecord { + aggregate_id: id.clone(), + version: 1, + data: vec![1], + }, + ) + .await + .unwrap(); + repo.save_snapshot_async( + &projection, + SnapshotRecord { + aggregate_id: id, + version: 2, + data: vec![2], + }, + ) + .await + .unwrap(); + + let loaded_counter = repo.get_snapshot_async(&counter).await.unwrap().unwrap(); + let loaded_projection = repo.get_snapshot_async(&projection).await.unwrap().unwrap(); + + assert_eq!(loaded_counter.version, 1); + assert_eq!(loaded_counter.data, vec![1]); + assert_eq!(loaded_projection.version, 2); + assert_eq!(loaded_projection.data, vec![2]); +} + +#[tokio::test] +async fn unsupported_codec_rows_fail_on_read() { + let Some(repo) = repository().await else { + return; + }; + let id = unique_id("bad-codec"); + sqlx::query( + r#" + INSERT INTO aggregate_events ( + aggregate_type, + aggregate_id, + sequence, + event_name, + event_version, + payload, + payload_codec, + payload_codec_version, + metadata, + recorded_at + ) + VALUES ($1, $2, 1, 'BadEvent', 1, $3, 'json', 1, '{}'::jsonb, now()) + "#, + ) + .bind("postgres.counter") + .bind(&id) + .bind(vec![0_u8]) + .execute(repo.pool()) + .await + .unwrap(); + + let identity = StreamIdentity::new("postgres.counter", &id).unwrap(); + let err = repo.get_stream(&identity).await.unwrap_err(); + + assert!( + matches!(err, RepositoryError::Model(message) if message.contains("unsupported payload codec")) + ); +} diff --git a/tests/sqlite_repository/main.rs b/tests/sqlite_repository/main.rs new file mode 100644 index 00000000..6de97145 --- /dev/null +++ b/tests/sqlite_repository/main.rs @@ -0,0 +1,288 @@ +#![cfg(feature = "sqlite")] + +use serde::{Deserialize, Serialize}; +use sourced_rust::{ + impl_aggregate, Aggregate, AsyncAggregateBuilder, AsyncCommitBatch, AsyncGetStream, + AsyncReadModelSessionStore, AsyncReadModelStore, AsyncSnapshotStore, AsyncStreamWrite, + AsyncTransactionalCommit, Entity, EventRecord, ReadModel, ReadModelSession, RepositoryError, + SnapshotRecord, SqliteRepository, StreamIdentity, +}; + +#[derive(Default)] +struct Counter { + entity: Entity, + value: i32, +} + +impl Counter { + fn increment(&mut self, id: &str, by: i32) { + self.entity.set_id(id); + self.entity.digest("Incremented", &by).unwrap(); + self.value += by; + } + + fn replay(&mut self, event: &EventRecord) -> Result<(), String> { + if event.event_name == "Incremented" { + let by = event.decode::().map_err(|err| err.to_string())?; + self.value += by; + } + Ok(()) + } +} + +impl_aggregate!(Counter, entity, replay, aggregate_type = "sqlite.counter"); + +#[derive(Default)] +struct CounterProjection { + entity: Entity, +} + +impl CounterProjection { + fn touch(&mut self, id: &str) { + self.entity.set_id(id); + self.entity.digest_empty("Touched").unwrap(); + } + + fn replay(&mut self, _event: &EventRecord) -> Result<(), String> { + Ok(()) + } +} + +impl_aggregate!( + CounterProjection, + entity, + replay, + aggregate_type = "sqlite.counter_projection" +); + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +struct CounterView { + id: String, + value: i32, +} + +impl ReadModel for CounterView { + const COLLECTION: &'static str = "sqlite_counter_views"; + + fn id(&self) -> &str { + &self.id + } +} + +async fn repository() -> SqliteRepository { + SqliteRepository::connect_and_migrate("sqlite::memory:") + .await + .unwrap() +} + +#[tokio::test] +async fn migration_is_idempotent_and_aggregate_stream_round_trips() { + let repo = repository().await; + repo.migrate().await.unwrap(); + let counter_repo = repo.clone().async_aggregate::(); + + let mut counter = Counter::default(); + counter.entity.set_correlation_id("corr-1"); + counter.increment("counter-1", 2); + counter.increment("counter-1", 3); + + counter_repo.commit(&mut counter).await.unwrap(); + + let loaded = counter_repo.get("counter-1").await.unwrap().unwrap(); + assert_eq!(loaded.value, 5); + assert_eq!(loaded.entity().events().len(), 2); + assert_eq!(loaded.entity().events()[0].sequence, 1); + assert_eq!(loaded.entity().events()[1].sequence, 2); + assert_eq!(loaded.entity().events()[0].correlation_id(), Some("corr-1")); +} + +#[tokio::test] +async fn aggregate_stream_identity_separates_same_id_across_types() { + let repo = repository().await; + let counter_repo = repo.clone().async_aggregate::(); + let projection_repo = repo.clone().async_aggregate::(); + + let mut counter = Counter::default(); + counter.increment("shared-id", 7); + let mut projection = CounterProjection::default(); + projection.touch("shared-id"); + + counter_repo.commit(&mut counter).await.unwrap(); + projection_repo.commit(&mut projection).await.unwrap(); + + let loaded_counter = counter_repo.get("shared-id").await.unwrap().unwrap(); + let loaded_projection = projection_repo.get("shared-id").await.unwrap().unwrap(); + + assert_eq!(loaded_counter.value, 7); + assert_eq!(loaded_counter.entity().events().len(), 1); + assert_eq!(loaded_projection.entity().events().len(), 1); +} + +#[tokio::test] +async fn optimistic_conflict_rolls_back_other_stream_and_read_model_plan() { + let repo = repository().await; + let counter_repo = repo.clone().async_aggregate::(); + + let mut original = Counter::default(); + original.increment("conflict-1", 1); + counter_repo.commit(&mut original).await.unwrap(); + + let mut stale = counter_repo.get("conflict-1").await.unwrap().unwrap(); + let mut winner = counter_repo.get("conflict-1").await.unwrap().unwrap(); + stale.increment("conflict-1", 10); + winner.increment("conflict-1", 20); + counter_repo.commit(&mut winner).await.unwrap(); + + let mut other = CounterProjection::default(); + other.touch("should-not-commit"); + + let view = CounterView { + id: "should-not-commit".into(), + value: 99, + }; + let mut read_models = ReadModelSession::new(); + read_models.document(&view).unwrap(); + + let stale_identity = StreamIdentity::new(Counter::aggregate_type(), "conflict-1").unwrap(); + let other_identity = + StreamIdentity::new(CounterProjection::aggregate_type(), "should-not-commit").unwrap(); + let err = repo + .commit_batch_async(AsyncCommitBatch { + streams: vec![ + AsyncStreamWrite::new(stale_identity.clone(), stale.entity_mut()), + AsyncStreamWrite::new(other_identity.clone(), other.entity_mut()), + ], + read_model_plans: vec![read_models.into_write_plan().unwrap()], + snapshots: Vec::new(), + }) + .await + .unwrap_err(); + + assert!(matches!(err, RepositoryError::ConcurrentWrite { .. })); + assert!(repo.get_stream(&other_identity).await.unwrap().is_none()); + assert!(repo + .get_model_async::("should-not-commit") + .await + .unwrap() + .is_none()); + assert_eq!(stale.entity().committed_version(), 1); + assert_eq!(stale.entity().new_events().len(), 1); +} + +#[tokio::test] +async fn read_model_session_persists_documents_and_processed_marks() { + let repo = repository().await; + let view = CounterView { + id: "view-1".into(), + value: 42, + }; + let mut session = ReadModelSession::new(); + session + .document(&view) + .unwrap() + .mark_processed("projection", "event-1"); + + let outcome = session.commit_async(&repo).await.unwrap(); + let loaded = repo + .get_model_async::("view-1") + .await + .unwrap() + .unwrap(); + let processed = repo + .is_processed_async("projection", "event-1") + .await + .unwrap(); + + assert!(outcome.was_applied()); + assert_eq!(loaded.version, 1); + assert_eq!(loaded.data, view); + assert!(processed); + + let mut duplicate = ReadModelSession::new(); + duplicate + .document(&CounterView { + id: "view-1".into(), + value: 100, + }) + .unwrap() + .mark_processed("projection", "event-1"); + let duplicate_outcome = duplicate.commit_async(&repo).await.unwrap(); + let still_loaded = repo + .get_model_async::("view-1") + .await + .unwrap() + .unwrap(); + + assert!(duplicate_outcome.was_skipped()); + assert_eq!(still_loaded.data.value, 42); +} + +#[tokio::test] +async fn snapshots_persist_by_full_stream_identity() { + let repo = repository().await; + let counter = StreamIdentity::new("sqlite.counter", "same-id").unwrap(); + let projection = StreamIdentity::new("sqlite.counter_projection", "same-id").unwrap(); + + repo.save_snapshot_async( + &counter, + SnapshotRecord { + aggregate_id: "same-id".into(), + version: 1, + data: vec![1], + }, + ) + .await + .unwrap(); + repo.save_snapshot_async( + &projection, + SnapshotRecord { + aggregate_id: "same-id".into(), + version: 2, + data: vec![2], + }, + ) + .await + .unwrap(); + + let loaded_counter = repo.get_snapshot_async(&counter).await.unwrap().unwrap(); + let loaded_projection = repo.get_snapshot_async(&projection).await.unwrap().unwrap(); + + assert_eq!(loaded_counter.version, 1); + assert_eq!(loaded_counter.data, vec![1]); + assert_eq!(loaded_projection.version, 2); + assert_eq!(loaded_projection.data, vec![2]); +} + +#[tokio::test] +async fn unsupported_codec_rows_fail_on_read() { + let repo = repository().await; + sqlx::query( + r#" + INSERT INTO aggregate_events ( + aggregate_type, + aggregate_id, + sequence, + event_name, + event_version, + payload, + payload_codec, + payload_codec_version, + metadata, + recorded_at + ) + VALUES (?, ?, 1, 'BadEvent', 1, x'00', 'json', 1, '{}', '0.000000000') + "#, + ) + .bind("sqlite.counter") + .bind("bad-codec") + .execute(repo.pool()) + .await + .unwrap(); + + let identity = StreamIdentity::new("sqlite.counter", "bad-codec").unwrap(); + let err = repo.get_stream(&identity).await.unwrap_err(); + + assert!( + matches!(err, RepositoryError::Model(message) if message.contains("unsupported payload codec")) + ); +}