1use 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
38const 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#[derive(Clone)]
51pub struct PostgresReadModelStore {
52 pool: PgPool,
53}
54
55impl PostgresReadModelStore {
56 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 pub fn with_pool(pool: PgPool) -> Self {
71 PostgresReadModelStore { pool }
72 }
73
74 pub fn pool(&self) -> &PgPool {
76 &self.pool
77 }
78
79 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
90fn 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
113fn 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 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 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}