Skip to main content

arc_es_postgres/
read_model_store.rs

1//! # Postgres Read Model Store
2//!
3//! Implementation of [`ReadModelStore`] backed by an [`sqlx::PgPool`]. Mirrors
4//! the SQLite store's contract: every projection table has the standard
5//! projection shape `(id TEXT PK, version BIGINT, data JSONB)`, and writes are
6//! version-gated so replay and at-least-once delivery converge.
7//!
8//! ## Idempotency
9//!
10//! [`upsert`](ReadModelStore::upsert) translates to
11//!
12//! ```sql
13//! INSERT INTO {table} (id, version, data) VALUES ($1, $2, $3)
14//! ON CONFLICT (id) DO UPDATE
15//!    SET version = EXCLUDED.version, data = EXCLUDED.data
16//!  WHERE {table}.version < EXCLUDED.version
17//! ```
18//!
19//! Applying an older or equal version never regresses state.
20//!
21//! ## Queries
22//!
23//! [`find_by`](ReadModelStore::find_by) uses Postgres' `data ->> $field`
24//! text-extraction operator. `IS NOT DISTINCT FROM` gives null-safe equality.
25//!
26//! ## Identifiers
27//!
28//! sqlx binds values, not identifiers. Table names are spliced into SQL, so
29//! they are validated against `[A-Za-z0-9_]` first ([`check_ident`]). Field
30//! names are bound as the `->>` operand and so are parameterized, but are still
31//! validated for parity with the SQLite store.
32
33use arc_core::read_model_store::{ReadModelError, ReadModelResult, ReadModelStore, Row, Upsert};
34use async_trait::async_trait;
35use sqlx::postgres::{PgPool, PgPoolOptions};
36use sqlx::Row as _;
37
38/// DDL for the `users_view` projection table. Idempotent.
39const USERS_VIEW_SCHEMA: &str = r#"
40CREATE TABLE IF NOT EXISTS users_view (
41    id      TEXT   NOT NULL PRIMARY KEY,
42    version BIGINT NOT NULL,
43    data    JSONB  NOT NULL
44);
45CREATE UNIQUE INDEX IF NOT EXISTS idx_users_view_email
46    ON users_view ((data ->> 'email'));
47"#;
48
49/// Postgres implementation of [`ReadModelStore`].
50#[derive(Clone)]
51pub struct PostgresReadModelStore {
52    pool: PgPool,
53}
54
55impl PostgresReadModelStore {
56    /// Build a new store from a database URL, creating a small pool.
57    pub async fn new(database_url: &str) -> ReadModelResult<Self> {
58        let pool = PgPoolOptions::new()
59            .max_connections(10)
60            .connect(database_url)
61            .await
62            .map_err(|e| {
63                ReadModelError::other(format!("Failed to build read-model pool: {}", e))
64            })?;
65        Ok(PostgresReadModelStore { pool })
66    }
67
68    /// Build a store from an existing pool. Lets tests share one pool with the
69    /// event store against the same database.
70    pub fn with_pool(pool: PgPool) -> Self {
71        PostgresReadModelStore { pool }
72    }
73
74    /// Borrow the underlying pool.
75    pub fn pool(&self) -> &PgPool {
76        &self.pool
77    }
78
79    /// Create the `users_view` projection table and its email index if absent.
80    /// Idempotent.
81    pub async fn initialize_schema(&self) -> ReadModelResult<()> {
82        sqlx::raw_sql(USERS_VIEW_SCHEMA)
83            .execute(&self.pool)
84            .await
85            .map_err(|e| ReadModelError::schema_failed(e.to_string()))?;
86        Ok(())
87    }
88}
89
90/// Validate a table or column name against an allow-list of characters before
91/// splicing it into a SQL string. Guards against caller-supplied identifiers
92/// carrying anything other than `[A-Za-z0-9_]`.
93fn check_ident(label: &str, ident: &str) -> ReadModelResult<()> {
94    if ident.is_empty() {
95        return Err(ReadModelError::other(format!("{label} cannot be empty")));
96    }
97    if !ident.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
98        return Err(ReadModelError::other(format!(
99            "invalid {label} '{ident}': only [A-Za-z0-9_] permitted"
100        )));
101    }
102    Ok(())
103}
104
105fn extract_version(row: &Row) -> ReadModelResult<i64> {
106    row.get("version").and_then(|v| v.as_i64()).ok_or_else(|| {
107        ReadModelError::write_failed(
108            "Upsert.row missing required i64 field 'version' for version-gated upsert",
109        )
110    })
111}
112
113/// Map a JSON primitive to the text Postgres' `->>` operator returns for that
114/// value, so `find_by` can compare against an extracted field.
115fn find_by_text(value: &serde_json::Value) -> ReadModelResult<Option<String>> {
116    match value {
117        serde_json::Value::Null => Ok(None),
118        serde_json::Value::String(s) => Ok(Some(s.clone())),
119        serde_json::Value::Number(n) => Ok(Some(n.to_string())),
120        serde_json::Value::Bool(b) => Ok(Some(if *b { "true".into() } else { "false".into() })),
121        other => Err(ReadModelError::query_failed(format!(
122            "find_by only accepts primitive values, got: {other}"
123        ))),
124    }
125}
126
127#[async_trait]
128impl ReadModelStore for PostgresReadModelStore {
129    async fn upsert(&self, op: Upsert) -> ReadModelResult<()> {
130        check_ident("table name", &op.table)?;
131        let version = extract_version(&op.row)?;
132
133        let sql = format!(
134            "INSERT INTO {table} (id, version, data) VALUES ($1, $2, $3) \
135             ON CONFLICT (id) DO UPDATE SET version = EXCLUDED.version, data = EXCLUDED.data \
136             WHERE {table}.version < EXCLUDED.version",
137            table = op.table
138        );
139
140        sqlx::query(&sql)
141            .bind(&op.key)
142            .bind(version)
143            .bind(&op.row)
144            .execute(&self.pool)
145            .await
146            .map_err(|e| ReadModelError::write_failed(e.to_string()))?;
147        Ok(())
148    }
149
150    async fn delete(&self, table: &str, key: &str) -> ReadModelResult<()> {
151        check_ident("table name", table)?;
152        sqlx::query(&format!("DELETE FROM {table} WHERE id = $1"))
153            .bind(key)
154            .execute(&self.pool)
155            .await
156            .map_err(|e| ReadModelError::write_failed(e.to_string()))?;
157        Ok(())
158    }
159
160    async fn get(&self, table: &str, key: &str) -> ReadModelResult<Option<Row>> {
161        check_ident("table name", table)?;
162        let row = sqlx::query(&format!("SELECT data FROM {table} WHERE id = $1 LIMIT 1"))
163            .bind(key)
164            .fetch_optional(&self.pool)
165            .await
166            .map_err(|e| ReadModelError::query_failed(e.to_string()))?;
167
168        match row {
169            Some(r) => {
170                Ok(Some(r.try_get("data").map_err(|e| {
171                    ReadModelError::query_failed(e.to_string())
172                })?))
173            }
174            None => Ok(None),
175        }
176    }
177
178    async fn find_by(
179        &self,
180        table: &str,
181        field: &str,
182        value: &serde_json::Value,
183    ) -> ReadModelResult<Vec<Row>> {
184        check_ident("table name", table)?;
185        check_ident("field name", field)?;
186        let needle = find_by_text(value)?;
187
188        let rows = sqlx::query(&format!(
189            "SELECT data FROM {table} WHERE data ->> $1 IS NOT DISTINCT FROM $2"
190        ))
191        .bind(field)
192        .bind(needle)
193        .fetch_all(&self.pool)
194        .await
195        .map_err(|e| ReadModelError::query_failed(e.to_string()))?;
196
197        rows.iter()
198            .map(|r| {
199                r.try_get("data")
200                    .map_err(|e| ReadModelError::query_failed(e.to_string()))
201            })
202            .collect()
203    }
204
205    async fn list(&self, table: &str) -> ReadModelResult<Vec<Row>> {
206        check_ident("table name", table)?;
207        let rows = sqlx::query(&format!("SELECT data FROM {table}"))
208            .fetch_all(&self.pool)
209            .await
210            .map_err(|e| ReadModelError::query_failed(e.to_string()))?;
211
212        rows.iter()
213            .map(|r| {
214                r.try_get("data")
215                    .map_err(|e| ReadModelError::query_failed(e.to_string()))
216            })
217            .collect()
218    }
219
220    async fn truncate(&self, table: &str) -> ReadModelResult<()> {
221        check_ident("table name", table)?;
222        sqlx::query(&format!("DELETE FROM {table}"))
223            .execute(&self.pool)
224            .await
225            .map_err(|e| ReadModelError::schema_failed(e.to_string()))?;
226        Ok(())
227    }
228}
229
230#[cfg(test)]
231mod tests {
232    use super::*;
233    use serde_json::json;
234
235    #[test]
236    fn test_check_ident_accepts_plain_name() {
237        check_ident("table name", "users_view").unwrap();
238    }
239
240    #[test]
241    fn test_check_ident_rejects_injection() {
242        let err = check_ident("table name", "users; DROP TABLE users").unwrap_err();
243        assert!(
244            matches!(err, ReadModelError::Other { ref message } if message.contains("table name")),
245            "got {err:?}"
246        );
247    }
248
249    #[test]
250    fn test_check_ident_rejects_empty() {
251        assert!(check_ident("field name", "").is_err());
252    }
253
254    #[test]
255    fn test_extract_version_reads_field() {
256        let row = json!({ "id": "u1", "version": 7 });
257        assert_eq!(extract_version(&row).unwrap(), 7);
258    }
259
260    #[test]
261    fn test_extract_version_missing_is_error() {
262        let row = json!({ "id": "u1" });
263        let err = extract_version(&row).unwrap_err();
264        assert!(matches!(err, ReadModelError::WriteFailed { .. }));
265    }
266
267    #[test]
268    fn test_find_by_text_maps_primitives() {
269        assert_eq!(find_by_text(&json!("a@b.c")).unwrap(), Some("a@b.c".into()));
270        assert_eq!(find_by_text(&json!(42)).unwrap(), Some("42".into()));
271        assert_eq!(find_by_text(&json!(true)).unwrap(), Some("true".into()));
272        assert_eq!(find_by_text(&json!(null)).unwrap(), None);
273    }
274
275    #[test]
276    fn test_find_by_text_rejects_composite() {
277        assert!(find_by_text(&json!({ "x": 1 })).is_err());
278        assert!(find_by_text(&json!([1, 2])).is_err());
279    }
280
281    // ── Live-database tests ──────────────────────────────────────────────────
282    // Gated behind ARC_POSTGRES_TEST_DATABASE_URL; see lib.rs for usage.
283
284    async fn live_store() -> Option<PostgresReadModelStore> {
285        let url = std::env::var("ARC_POSTGRES_TEST_DATABASE_URL").ok()?;
286        let store = PostgresReadModelStore::new(&url).await.expect("connect");
287        store.initialize_schema().await.expect("schema");
288        sqlx::query("DELETE FROM users_view")
289            .execute(store.pool())
290            .await
291            .expect("clear");
292        Some(store)
293    }
294
295    fn user_row(id: &str, name: &str, email: &str, version: i64) -> Row {
296        json!({ "id": id, "name": name, "email": email, "version": version })
297    }
298
299    #[tokio::test]
300    #[serial_test::serial]
301    async fn test_live_upsert_and_get() {
302        let Some(store) = live_store().await else {
303            return;
304        };
305        store
306            .upsert(Upsert::new(
307                "users_view",
308                "u1",
309                user_row("u1", "Alice", "a@b.c", 1),
310            ))
311            .await
312            .unwrap();
313        let got = store.get("users_view", "u1").await.unwrap().unwrap();
314        assert_eq!(got["name"], "Alice");
315        assert_eq!(got["version"], 1);
316    }
317
318    #[tokio::test]
319    #[serial_test::serial]
320    async fn test_live_upsert_version_gate() {
321        let Some(store) = live_store().await else {
322            return;
323        };
324        store
325            .upsert(Upsert::new(
326                "users_view",
327                "u1",
328                user_row("u1", "Alice2", "a@b.c", 2),
329            ))
330            .await
331            .unwrap();
332        // Stale write must not regress.
333        store
334            .upsert(Upsert::new(
335                "users_view",
336                "u1",
337                user_row("u1", "Stale", "a@b.c", 1),
338            ))
339            .await
340            .unwrap();
341        let got = store.get("users_view", "u1").await.unwrap().unwrap();
342        assert_eq!(got["name"], "Alice2");
343        assert_eq!(got["version"], 2);
344    }
345
346    #[tokio::test]
347    #[serial_test::serial]
348    async fn test_live_find_by_email() {
349        let Some(store) = live_store().await else {
350            return;
351        };
352        store
353            .upsert(Upsert::new(
354                "users_view",
355                "u1",
356                user_row("u1", "Alice", "a@b.c", 1),
357            ))
358            .await
359            .unwrap();
360        store
361            .upsert(Upsert::new(
362                "users_view",
363                "u2",
364                user_row("u2", "Bob", "b@b.c", 1),
365            ))
366            .await
367            .unwrap();
368        let hits = store
369            .find_by("users_view", "email", &json!("b@b.c"))
370            .await
371            .unwrap();
372        assert_eq!(hits.len(), 1);
373        assert_eq!(hits[0]["id"], "u2");
374    }
375
376    #[tokio::test]
377    #[serial_test::serial]
378    async fn test_live_delete_and_truncate() {
379        let Some(store) = live_store().await else {
380            return;
381        };
382        store
383            .upsert(Upsert::new(
384                "users_view",
385                "u1",
386                user_row("u1", "Alice", "a@b.c", 1),
387            ))
388            .await
389            .unwrap();
390        store
391            .upsert(Upsert::new(
392                "users_view",
393                "u2",
394                user_row("u2", "Bob", "b@b.c", 1),
395            ))
396            .await
397            .unwrap();
398        store.delete("users_view", "u1").await.unwrap();
399        assert!(store.get("users_view", "u1").await.unwrap().is_none());
400        assert_eq!(store.list("users_view").await.unwrap().len(), 1);
401        store.truncate("users_view").await.unwrap();
402        assert!(store.list("users_view").await.unwrap().is_empty());
403    }
404}