1use std::path::Path;
2
3use rusqlite::{Connection, OptionalExtension, TransactionBehavior, params};
4
5use super::{IncomingEdge, RecordQuery, Store, check_sequence};
6use crate::error::{Error, Result};
7use crate::event::{Event, EventKind, scrub_event};
8use crate::id::{Origin, RecordId};
9use crate::model::Record;
10
11pub struct SqliteStore {
15 origin: Origin,
16 conn: Connection,
17}
18
19const SCHEMA: &str = r#"
20CREATE TABLE IF NOT EXISTS pd_meta (key TEXT PRIMARY KEY, value TEXT NOT NULL);
21CREATE TABLE IF NOT EXISTS pd_records (
22 id TEXT PRIMARY KEY,
23 kind TEXT NOT NULL,
24 version INTEGER NOT NULL,
25 json TEXT NOT NULL
26);
27CREATE INDEX IF NOT EXISTS pd_records_kind ON pd_records(kind);
28CREATE TABLE IF NOT EXISTS pd_events (
29 record_id TEXT NOT NULL,
30 sequence INTEGER NOT NULL,
31 event_id TEXT NOT NULL UNIQUE,
32 idempotency_key TEXT UNIQUE,
33 json TEXT NOT NULL,
34 PRIMARY KEY (record_id, sequence)
35);
36CREATE TABLE IF NOT EXISTS pd_relations (
37 source TEXT NOT NULL,
38 relation_id TEXT NOT NULL,
39 target TEXT NOT NULL,
40 PRIMARY KEY (source, relation_id)
41);
42CREATE INDEX IF NOT EXISTS pd_relations_target ON pd_relations(target);
43"#;
44
45fn db(e: rusqlite::Error) -> Error {
46 Error::Storage(e.to_string())
47}
48
49impl SqliteStore {
50 pub fn open(path: impl AsRef<Path>, origin: Origin) -> Result<Self> {
51 Self::init(Connection::open(path).map_err(db)?, origin)
52 }
53
54 pub fn open_in_memory(origin: Origin) -> Result<Self> {
55 Self::init(Connection::open_in_memory().map_err(db)?, origin)
56 }
57
58 fn init(conn: Connection, origin: Origin) -> Result<Self> {
59 conn.execute_batch("PRAGMA journal_mode = WAL; PRAGMA foreign_keys = ON;").map_err(db)?;
60 conn.execute_batch(SCHEMA).map_err(db)?;
61 let existing: Option<String> = conn
63 .query_row("SELECT value FROM pd_meta WHERE key = 'origin'", [], |r| r.get(0))
64 .optional()
65 .map_err(db)?;
66 match existing {
67 Some(o) if o != origin.as_str() => {
68 return Err(Error::Storage(format!("database belongs to origin {o}, not {origin}")));
69 }
70 Some(_) => {}
71 None => {
72 conn.execute("INSERT INTO pd_meta (key, value) VALUES ('origin', ?1)", [origin.as_str()])
73 .map_err(db)?;
74 }
75 }
76 Ok(Self { origin, conn })
77 }
78
79 fn decode<T: serde::de::DeserializeOwned>(json: String) -> Result<T> {
80 serde_json::from_str(&json).map_err(|e| Error::Storage(format!("corrupt row: {e}")))
81 }
82}
83
84impl Store for SqliteStore {
85 fn origin(&self) -> &Origin {
86 &self.origin
87 }
88
89 fn get(&self, id: &RecordId) -> Result<Option<Record>> {
90 let json: Option<String> = self
91 .conn
92 .query_row("SELECT json FROM pd_records WHERE id = ?1", [id.to_string()], |r| r.get(0))
93 .optional()
94 .map_err(db)?;
95 json.map(Self::decode).transpose()
96 }
97
98 fn history(&self, id: &RecordId) -> Result<Vec<Event>> {
99 let mut stmt = self
100 .conn
101 .prepare_cached("SELECT json FROM pd_events WHERE record_id = ?1 ORDER BY sequence")
102 .map_err(db)?;
103 let rows = stmt.query_map([id.to_string()], |r| r.get::<_, String>(0)).map_err(db)?;
104 rows.map(|r| r.map_err(db).and_then(Self::decode)).collect()
105 }
106
107 fn event_by_idempotency_key(&self, key: &str) -> Result<Option<Event>> {
108 let json: Option<String> = self
109 .conn
110 .query_row("SELECT json FROM pd_events WHERE idempotency_key = ?1", [key], |r| r.get(0))
111 .optional()
112 .map_err(db)?;
113 json.map(Self::decode).transpose()
114 }
115
116 fn commit(&mut self, event: Event, snapshot: Record) -> Result<()> {
117 let id = event.record_id.to_string();
118 let tx = self.conn.transaction_with_behavior(TransactionBehavior::Immediate).map_err(db)?;
119 let current: u64 = tx
120 .query_row("SELECT version FROM pd_records WHERE id = ?1", [&id], |r| r.get(0))
121 .optional()
122 .map_err(db)?
123 .unwrap_or(0);
124 check_sequence(&event.record_id, current, &event, &snapshot)?;
125 if let Some(key) = &event.idempotency_key {
126 let used: Option<i64> = tx
127 .query_row("SELECT 1 FROM pd_events WHERE idempotency_key = ?1", [key], |r| r.get(0))
128 .optional()
129 .map_err(db)?;
130 if used.is_some() {
131 return Err(Error::IdempotencyConflict(key.clone()));
132 }
133 }
134 if let EventKind::Redacted { fields } = &event.kind {
135 let earlier: Vec<(i64, String)> = {
136 let mut stmt = tx.prepare("SELECT sequence, json FROM pd_events WHERE record_id = ?1").map_err(db)?;
137 let rows = stmt.query_map([&id], |r| Ok((r.get(0)?, r.get(1)?))).map_err(db)?;
138 rows.collect::<std::result::Result<_, _>>().map_err(db)?
139 };
140 for (seq, json) in earlier {
141 let mut e: Event = Self::decode(json)?;
142 scrub_event(&mut e, fields);
143 tx.execute(
144 "UPDATE pd_events SET json = ?1 WHERE record_id = ?2 AND sequence = ?3",
145 params![serde_json::to_string(&e)?, &id, seq],
146 )
147 .map_err(db)?;
148 }
149 }
150 tx.execute(
151 "INSERT INTO pd_events (record_id, sequence, event_id, idempotency_key, json) VALUES (?1, ?2, ?3, ?4, ?5)",
152 params![&id, event.sequence, event.id.to_string(), event.idempotency_key, serde_json::to_string(&event)?],
153 )
154 .map_err(db)?;
155 tx.execute(
156 "INSERT INTO pd_records (id, kind, version, json) VALUES (?1, ?2, ?3, ?4)
157 ON CONFLICT(id) DO UPDATE SET version = excluded.version, json = excluded.json",
158 params![&id, snapshot.kind.to_string(), snapshot.version, serde_json::to_string(&snapshot)?],
159 )
160 .map_err(db)?;
161 tx.execute("DELETE FROM pd_relations WHERE source = ?1", [&id]).map_err(db)?;
162 for rel in &snapshot.relations {
163 tx.execute(
164 "INSERT INTO pd_relations (source, relation_id, target) VALUES (?1, ?2, ?3)",
165 params![&id, rel.id.to_string(), rel.target.to_string()],
166 )
167 .map_err(db)?;
168 }
169 tx.commit().map_err(db)
170 }
171
172 fn list(&self, query: &RecordQuery) -> Result<Vec<Record>> {
173 let mut stmt = self.conn.prepare_cached("SELECT kind, json FROM pd_records ORDER BY id DESC").map_err(db)?;
174 let kinds: Vec<String> = query.kinds.iter().map(|k| k.to_string()).collect();
175 let rows = stmt.query_map([], |r| Ok((r.get::<_, String>(0)?, r.get::<_, String>(1)?))).map_err(db)?;
176 let mut records = Vec::new();
177 for row in rows {
178 let (kind, json) = row.map_err(db)?;
179 if kinds.is_empty() || kinds.contains(&kind) {
180 records.push(Self::decode(json)?);
181 }
182 }
183 Ok(query.page(records.into_iter()))
184 }
185
186 fn incoming(&self, target: &RecordId) -> Result<Vec<IncomingEdge>> {
187 let mut stmt =
188 self.conn.prepare_cached("SELECT DISTINCT source FROM pd_relations WHERE target = ?1").map_err(db)?;
189 let sources: Vec<String> = stmt
190 .query_map([target.to_string()], |r| r.get(0))
191 .map_err(db)?
192 .collect::<std::result::Result<_, _>>()
193 .map_err(db)?;
194 let mut out = Vec::new();
195 for s in sources {
196 let id: RecordId = s.parse()?;
197 if let Some(record) = self.get(&id)? {
198 for rel in record.relations.iter().filter(|r| &r.target == target) {
199 out.push(IncomingEdge {
200 source: record.id.clone(),
201 source_kind: record.kind.clone(),
202 relation: rel.clone(),
203 });
204 }
205 }
206 }
207 Ok(out)
208 }
209}