Skip to main content

eventuary_sqlite/
reader.rs

1use std::collections::{HashMap, VecDeque};
2use std::sync::Arc;
3use std::time::Duration;
4
5use chrono::{DateTime, Utc};
6use rusqlite::types::Value;
7use tokio::sync::Mutex;
8use tokio::sync::Notify;
9use tokio::sync::mpsc;
10
11use eventuary_core::io::filter::EventFilter;
12use eventuary_core::io::stream::SpawnedStream;
13use eventuary_core::io::{Acker, Cursor, Message, Reader};
14use eventuary_core::{
15    Error, Result, SerializedEvent, SerializedPayload, StartFrom, StartableSubscription,
16    TopicPattern,
17};
18
19use crate::database::SqliteConn;
20use crate::relation::SqliteRelationName;
21
22#[derive(
23    Debug, Clone, Copy, Eq, PartialEq, Ord, PartialOrd, Hash, serde::Serialize, serde::Deserialize,
24)]
25#[serde(transparent)]
26pub struct SqliteCursor {
27    pub sequence: i64,
28}
29
30impl SqliteCursor {
31    pub fn new(sequence: i64) -> Self {
32        Self { sequence }
33    }
34
35    pub fn sequence(&self) -> i64 {
36        self.sequence
37    }
38}
39
40impl Cursor for SqliteCursor {}
41
42#[derive(Debug, Clone)]
43pub struct SqliteSubscription {
44    pub start: StartFrom<SqliteCursor>,
45    pub filter: EventFilter,
46    pub batch_size: Option<usize>,
47    pub limit: Option<usize>,
48}
49
50impl Default for SqliteSubscription {
51    fn default() -> Self {
52        Self {
53            start: StartFrom::Latest,
54            filter: EventFilter::default(),
55            batch_size: None,
56            limit: None,
57        }
58    }
59}
60
61impl StartableSubscription<SqliteCursor> for SqliteSubscription {
62    fn with_start(mut self, start: StartFrom<SqliteCursor>) -> Self {
63        self.start = start;
64        self
65    }
66}
67
68#[derive(Debug, Clone)]
69pub struct SqliteReaderConfig {
70    pub events_relation: SqliteRelationName,
71    pub poll_interval: Duration,
72    pub default_batch_size: usize,
73}
74
75impl Default for SqliteReaderConfig {
76    fn default() -> Self {
77        Self {
78            events_relation: SqliteRelationName::new("events").expect("default events relation"),
79            poll_interval: Duration::from_millis(100),
80            default_batch_size: 100,
81        }
82    }
83}
84
85#[derive(Clone)]
86pub struct SqliteCursorAcker {
87    state: Arc<Mutex<CursorState>>,
88    notify: Arc<Notify>,
89    sequence: i64,
90}
91
92struct CursorState {
93    last_acked: i64,
94    pending_nack: bool,
95}
96
97impl Acker for SqliteCursorAcker {
98    async fn ack(&self) -> Result<()> {
99        let mut state = self.state.lock().await;
100        if self.sequence > state.last_acked {
101            state.last_acked = self.sequence;
102        }
103        state.pending_nack = false;
104        self.notify.notify_waiters();
105        Ok(())
106    }
107
108    async fn nack(&self) -> Result<()> {
109        let mut state = self.state.lock().await;
110        state.pending_nack = true;
111        self.notify.notify_waiters();
112        Ok(())
113    }
114}
115
116pub struct SqliteReader {
117    conn: SqliteConn,
118    config: SqliteReaderConfig,
119}
120
121impl SqliteReader {
122    pub fn new(conn: SqliteConn, config: SqliteReaderConfig) -> Self {
123        Self { conn, config }
124    }
125}
126
127impl Reader for SqliteReader {
128    type Subscription = SqliteSubscription;
129    type Acker = SqliteCursorAcker;
130    type Cursor = SqliteCursor;
131    type Stream = SpawnedStream<SqliteCursorAcker, SqliteCursor>;
132
133    async fn read(&self, subscription: Self::Subscription) -> Result<Self::Stream> {
134        let conn = Arc::clone(&self.conn);
135        let events_relation = self.config.events_relation.render();
136        let poll_interval = self.config.poll_interval;
137        let batch_size = subscription
138            .batch_size
139            .unwrap_or(self.config.default_batch_size)
140            .clamp(1, 1000);
141        let filter = subscription.filter.clone();
142        let limit = subscription.limit;
143        let (tx, rx) = mpsc::channel(64);
144
145        let (mut after_seq, lower_bound_ts) =
146            match resolve_initial_position(&conn, &events_relation, &subscription).await {
147                Ok(pos) => pos,
148                Err(e) => {
149                    let _ = tx.send(Err(e)).await;
150                    return Ok(SpawnedStream::from_receiver(rx));
151                }
152            };
153
154        let state = Arc::new(Mutex::new(CursorState {
155            last_acked: after_seq,
156            pending_nack: false,
157        }));
158        let notify = Arc::new(Notify::new());
159
160        let handle = tokio::spawn(async move {
161            let mut delivered = 0usize;
162            let mut buffer: VecDeque<(SerializedEvent, i64)> = VecDeque::new();
163            loop {
164                if buffer.is_empty() {
165                    let fetched = match fetch_batch(
166                        &conn,
167                        &events_relation,
168                        after_seq,
169                        batch_size,
170                        lower_bound_ts,
171                        &filter,
172                    )
173                    .await
174                    {
175                        Ok(b) => b,
176                        Err(e) => {
177                            let _ = tx.send(Err(e)).await;
178                            return;
179                        }
180                    };
181                    if fetched.is_empty() {
182                        tokio::time::sleep(poll_interval).await;
183                        continue;
184                    }
185                    buffer.extend(fetched);
186                }
187
188                while let Some((serialized, sequence)) = buffer.front() {
189                    let sequence = *sequence;
190                    let event = match serialized.to_event() {
191                        Ok(e) => e,
192                        Err(e) => {
193                            let _ = tx
194                                .send(Err(Error::Serialization(format!(
195                                    "decode event at sequence {sequence}: {e}"
196                                ))))
197                                .await;
198                            return;
199                        }
200                    };
201                    if !filter.matches(&event) {
202                        buffer.pop_front();
203                        after_seq = sequence;
204                        continue;
205                    }
206                    if let Some(l) = limit
207                        && delivered >= l
208                    {
209                        return;
210                    }
211                    let acker = SqliteCursorAcker {
212                        state: Arc::clone(&state),
213                        notify: Arc::clone(&notify),
214                        sequence,
215                    };
216                    let cursor = SqliteCursor { sequence };
217                    if tx
218                        .send(Ok(Message::new(event, acker, cursor)))
219                        .await
220                        .is_err()
221                    {
222                        return;
223                    }
224                    delivered += 1;
225
226                    loop {
227                        {
228                            let guard = state.lock().await;
229                            if guard.last_acked >= sequence {
230                                after_seq = sequence;
231                                buffer.pop_front();
232                                break;
233                            }
234                            if guard.pending_nack {
235                                break;
236                            }
237                            if tx.is_closed() {
238                                return;
239                            }
240                        }
241                        notify.notified().await;
242                    }
243                }
244            }
245        });
246
247        Ok(SpawnedStream::new(rx, handle))
248    }
249}
250
251async fn resolve_initial_position(
252    conn: &SqliteConn,
253    events_relation: &str,
254    subscription: &SqliteSubscription,
255) -> Result<(i64, Option<DateTime<Utc>>)> {
256    match subscription.start.clone() {
257        StartFrom::After(cursor) => Ok((cursor.sequence, None)),
258        StartFrom::Earliest => Ok((0, None)),
259        StartFrom::Latest => {
260            let conn = Arc::clone(conn);
261            let org = subscription
262                .filter
263                .organization
264                .as_ref()
265                .map(|o| o.as_str().to_owned());
266            let relation = events_relation.to_owned();
267            tokio::task::spawn_blocking(move || {
268                let guard = conn.lock().map_err(|e| Error::Store(e.to_string()))?;
269                let seq: i64 = match org {
270                    Some(o) => guard
271                        .query_row(
272                            &format!(
273                                "SELECT COALESCE(MAX(sequence), 0) FROM {relation} WHERE organization = ?1"
274                            ),
275                            rusqlite::params![o],
276                            |r| r.get(0),
277                        )
278                        .map_err(|e| Error::Store(e.to_string()))?,
279                    None => guard
280                        .query_row(
281                            &format!("SELECT COALESCE(MAX(sequence), 0) FROM {relation}"),
282                            [],
283                            |r| r.get(0),
284                        )
285                        .map_err(|e| Error::Store(e.to_string()))?,
286                };
287                Ok::<(i64, Option<DateTime<Utc>>), Error>((seq, None))
288            })
289            .await
290            .map_err(|e| Error::Store(format!("blocking task panicked: {e}")))?
291        }
292        StartFrom::Timestamp(ts) => {
293            let conn = Arc::clone(conn);
294            let org = subscription
295                .filter
296                .organization
297                .as_ref()
298                .map(|o| o.as_str().to_owned());
299            let ts_str = ts.to_rfc3339();
300            let relation = events_relation.to_owned();
301            tokio::task::spawn_blocking(move || {
302                let guard = conn.lock().map_err(|e| Error::Store(e.to_string()))?;
303                let seq: i64 = match org {
304                    Some(o) => guard
305                        .query_row(
306                            &format!(
307                                "SELECT COALESCE(MIN(sequence), 1) - 1 FROM {relation} \
308                                 WHERE organization = ?1 AND timestamp >= ?2"
309                            ),
310                            rusqlite::params![o, ts_str],
311                            |r| r.get(0),
312                        )
313                        .map_err(|e| Error::Store(e.to_string()))?,
314                    None => guard
315                        .query_row(
316                            &format!(
317                                "SELECT COALESCE(MIN(sequence), 1) - 1 FROM {relation} \
318                                 WHERE timestamp >= ?1"
319                            ),
320                            rusqlite::params![ts_str],
321                            |r| r.get(0),
322                        )
323                        .map_err(|e| Error::Store(e.to_string()))?,
324                };
325                Ok::<(i64, Option<DateTime<Utc>>), Error>((seq.max(0), Some(ts)))
326            })
327            .await
328            .map_err(|e| Error::Store(format!("blocking task panicked: {e}")))?
329        }
330    }
331}
332
333async fn fetch_batch(
334    conn: &SqliteConn,
335    events_relation: &str,
336    after_seq: i64,
337    take: usize,
338    lower_bound_ts: Option<DateTime<Utc>>,
339    filter: &EventFilter,
340) -> Result<Vec<(SerializedEvent, i64)>> {
341    let conn = Arc::clone(conn);
342    let relation = events_relation.to_owned();
343    let org = filter.organization.as_ref().map(|o| o.as_str().to_owned());
344    let exact_topic: Option<String> = filter.topic.as_ref().map(|p| match p {
345        TopicPattern::Exact(t) => t.as_str().to_owned(),
346    });
347    let ns_prefix = filter.namespace.as_ref().and_then(|p| match p {
348        eventuary_core::NamespacePattern::Prefix(ns) if !ns.is_root() => {
349            Some(ns.as_str().to_owned())
350        }
351        _ => None,
352    });
353    let ts_str = lower_bound_ts.map(|t| t.to_rfc3339());
354
355    tokio::task::spawn_blocking(move || {
356        let guard = conn.lock().map_err(|e| Error::Store(e.to_string()))?;
357
358        let mut sql = format!(
359            "SELECT sequence, id, organization, namespace, topic, event_key, payload, content_type, metadata, \
360             timestamp, version, parent_id, correlation_id, causation_id \
361             FROM {relation} WHERE sequence > ?1"
362        );
363        let mut params: Vec<Value> = vec![Value::Integer(after_seq)];
364        let mut idx = 2usize;
365
366        if let Some(o) = &org {
367            sql.push_str(&format!(" AND organization = ?{idx}"));
368            params.push(Value::Text(o.clone()));
369            idx += 1;
370        }
371        if let Some(t) = &exact_topic {
372            sql.push_str(&format!(" AND topic = ?{idx}"));
373            params.push(Value::Text(t.clone()));
374            idx += 1;
375        }
376        if let Some(prefix) = &ns_prefix {
377            sql.push_str(&format!(
378                " AND (namespace = ?{idx} OR namespace LIKE ?{} || '/%')",
379                idx
380            ));
381            params.push(Value::Text(prefix.clone()));
382            idx += 1;
383        }
384        if let Some(ts) = &ts_str {
385            sql.push_str(&format!(" AND timestamp >= ?{idx}"));
386            params.push(Value::Text(ts.clone()));
387            idx += 1;
388        }
389        sql.push_str(&format!(" ORDER BY sequence ASC LIMIT ?{idx}"));
390        params.push(Value::Integer(take as i64));
391
392        let mut stmt = guard
393            .prepare(&sql)
394            .map_err(|e| Error::Store(e.to_string()))?;
395        let rows = stmt
396            .query_map(rusqlite::params_from_iter(params.iter()), |row| {
397                let sequence: i64 = row.get(0)?;
398                let id: String = row.get(1)?;
399                let organization: String = row.get(2)?;
400                let namespace: String = row.get(3)?;
401                let topic: String = row.get(4)?;
402                let key: Option<String> = row.get(5)?;
403                let payload_str: String = row.get(6)?;
404                let content_type: String = row.get(7)?;
405                let metadata_str: String = row.get(8)?;
406                let timestamp_str: String = row.get(9)?;
407                let version: i64 = row.get(10)?;
408                let parent_id: Option<String> = row.get(11)?;
409                let correlation_id: Option<String> = row.get(12)?;
410                let causation_id: Option<String> = row.get(13)?;
411                Ok((
412                    sequence,
413                    id,
414                    organization,
415                    namespace,
416                    topic,
417                    key,
418                    payload_str,
419                    content_type,
420                    metadata_str,
421                    timestamp_str,
422                    version,
423                    parent_id,
424                    correlation_id,
425                    causation_id,
426                ))
427            })
428            .map_err(|e| Error::Store(e.to_string()))?;
429
430        let mut out = Vec::new();
431        for row in rows {
432            let (
433                sequence,
434                id,
435                organization,
436                namespace,
437                topic,
438                key,
439                payload_str,
440                content_type,
441                metadata_str,
442                timestamp_str,
443                version,
444                parent_id,
445                correlation_id,
446                causation_id,
447            ) = row.map_err(|e| Error::Store(e.to_string()))?;
448
449            let payload: SerializedPayload = serde_json::from_str(&payload_str)
450                .map_err(|e| Error::Serialization(format!("decode payload: {e}")))?;
451            let _ = content_type;
452            let id = uuid::Uuid::parse_str(&id)
453                .map_err(|e| Error::Serialization(format!("decode id: {e}")))?;
454            let parent_id = parent_id
455                .as_deref()
456                .map(uuid::Uuid::parse_str)
457                .transpose()
458                .map_err(|e| Error::Serialization(format!("decode parent_id: {e}")))?;
459            let metadata: HashMap<String, String> = serde_json::from_str(&metadata_str)
460                .map_err(|e| Error::Serialization(format!("decode metadata: {e}")))?;
461            let timestamp = DateTime::parse_from_rfc3339(&timestamp_str)
462                .map(|d| d.with_timezone(&Utc))
463                .map_err(|e| Error::Serialization(format!("decode timestamp: {e}")))?;
464            out.push((
465                SerializedEvent {
466                    id,
467                    organization,
468                    namespace,
469                    topic,
470                    payload,
471                    metadata,
472                    timestamp,
473                    version: version as u64,
474                    key,
475                    parent_id,
476                    correlation_id,
477                    causation_id,
478                },
479                sequence,
480            ));
481        }
482        Ok(out)
483    })
484    .await
485    .map_err(|e| Error::Store(format!("blocking task panicked: {e}")))?
486}
487
488#[cfg(test)]
489mod tests {
490    use super::*;
491    use eventuary_core::io::{Cursor, CursorId};
492
493    #[test]
494    fn sqlite_cursor_id_is_global() {
495        assert_eq!(SqliteCursor::new(42).id(), CursorId::global());
496    }
497}