1use std::marker::PhantomData;
4use std::sync::Mutex;
5
6use chrono::Utc;
7use rusqlite::{params, Connection, TransactionBehavior};
8use serde::{de::DeserializeOwned, Serialize};
9use uuid::Uuid;
10
11use substrate_core::event_store_port::{EventEnvelope, EventStorePort};
12
13use crate::error::StoreError;
14use crate::schema;
15
16pub struct SqliteEventStore<E> {
18 conn: Mutex<Connection>,
19 _event: PhantomData<E>,
20}
21
22impl<E> SqliteEventStore<E>
23where
24 E: Serialize + DeserializeOwned + Clone + Send + Sync,
25{
26 pub fn open(path: &str) -> Result<Self, StoreError> {
28 let conn = Connection::open(path)?;
29 schema::init(&conn)?;
30 Ok(Self {
31 conn: Mutex::new(conn),
32 _event: PhantomData,
33 })
34 }
35
36 pub fn open_in_memory() -> Result<Self, StoreError> {
38 let conn = Connection::open_in_memory()?;
39 schema::init(&conn)?;
40 Ok(Self {
41 conn: Mutex::new(conn),
42 _event: PhantomData,
43 })
44 }
45
46 pub fn load_all_global(&self) -> Result<Vec<EventEnvelope<E>>, StoreError> {
48 let conn = self.conn.lock().unwrap();
49 let mut stmt = conn.prepare(
50 "SELECT aggregate_id, aggregate_seq, global_seq, payload, occurred_at \
51 FROM event_log ORDER BY global_seq ASC",
52 )?;
53 let rows = stmt.query_map([], |row| {
54 Ok((
55 row.get::<_, String>(0)?,
56 row.get::<_, i64>(1)?,
57 row.get::<_, i64>(2)?,
58 row.get::<_, String>(3)?,
59 row.get::<_, i64>(4)?,
60 ))
61 })?;
62 rows.map(|row| {
63 let (aggregate_id, aggregate_seq, global_seq, payload, occurred_at) = row?;
64 Ok(EventEnvelope {
65 aggregate_id: Uuid::parse_str(&aggregate_id).unwrap_or_else(|_| Uuid::nil()),
66 aggregate_seq: aggregate_seq as u64,
67 global_seq: global_seq as u64,
68 event: serde_json::from_str(&payload)?,
69 occurred_at,
70 })
71 })
72 .collect()
73 }
74}
75
76impl<E> EventStorePort for SqliteEventStore<E>
77where
78 E: Serialize + DeserializeOwned + Clone + Send + Sync,
79{
80 type Error = StoreError;
81 type Event = E;
82
83 fn append(
84 &self,
85 aggregate_id: Uuid,
86 expected_seq: u64,
87 event: &Self::Event,
88 ) -> Result<EventEnvelope<Self::Event>, Self::Error> {
89 let mut conn = self.conn.lock().unwrap();
90 let tx = conn.transaction_with_behavior(TransactionBehavior::Immediate)?;
91
92 let current: i64 = tx
93 .query_row(
94 "SELECT COUNT(*) FROM event_log WHERE aggregate_id = ?1",
95 params![aggregate_id.to_string()],
96 |row| row.get(0),
97 )
98 .unwrap_or(0);
99
100 if current as u64 != expected_seq {
101 tx.rollback()?;
102 return Err(StoreError::DuplicateEventSeq {
103 aggregate_id: aggregate_id.to_string(),
104 expected: expected_seq,
105 });
106 }
107
108 let global_seq: i64 = tx
109 .query_row(
110 "SELECT COALESCE(MAX(global_seq), -1) + 1 FROM event_log",
111 [],
112 |row| row.get(0),
113 )
114 .unwrap_or(0);
115
116 let payload = serde_json::to_string(event)?;
117 let occurred_at = Utc::now().timestamp();
118
119 tx.execute(
120 "INSERT INTO event_log (aggregate_id, aggregate_seq, global_seq, payload, occurred_at) \
121 VALUES (?1, ?2, ?3, ?4, ?5)",
122 params![
123 aggregate_id.to_string(),
124 expected_seq as i64,
125 global_seq,
126 payload,
127 occurred_at,
128 ],
129 )?;
130
131 tx.commit()?;
132
133 Ok(EventEnvelope {
134 aggregate_id,
135 aggregate_seq: expected_seq,
136 global_seq: global_seq as u64,
137 event: event.clone(),
138 occurred_at,
139 })
140 }
141
142 fn load(&self, aggregate_id: Uuid) -> Result<Vec<EventEnvelope<Self::Event>>, Self::Error> {
143 let conn = self.conn.lock().unwrap();
144 let mut stmt = conn.prepare(
145 "SELECT aggregate_id, aggregate_seq, global_seq, payload, occurred_at \
146 FROM event_log WHERE aggregate_id = ?1 ORDER BY aggregate_seq ASC",
147 )?;
148 let rows = stmt.query_map(params![aggregate_id.to_string()], |row| {
149 Ok((
150 row.get::<_, String>(0)?,
151 row.get::<_, i64>(1)?,
152 row.get::<_, i64>(2)?,
153 row.get::<_, String>(3)?,
154 row.get::<_, i64>(4)?,
155 ))
156 })?;
157 rows.map(|row| {
158 let (aggregate_id, aggregate_seq, global_seq, payload, occurred_at) = row?;
159 Ok(EventEnvelope {
160 aggregate_id: Uuid::parse_str(&aggregate_id).unwrap_or_else(|_| Uuid::nil()),
161 aggregate_seq: aggregate_seq as u64,
162 global_seq: global_seq as u64,
163 event: serde_json::from_str(&payload)?,
164 occurred_at,
165 })
166 })
167 .collect()
168 }
169}