1use std::{
4 collections::{BTreeMap, BTreeSet, VecDeque},
5 pin::Pin,
6};
7
8use ankurah_core::{
9 error::RetrievalError,
10 storage::{StorageDump, StorageDumpItem},
11};
12use ankurah_proto::{
13 AttestationSet, Attested, Clock, CollectionId, EntityId, EntityState, Event, EventId, OperationSet, State, StateBuffers,
14};
15use async_trait::async_trait;
16use futures_util::{stream, Stream};
17use rusqlite::Connection;
18
19use crate::{SqliteConnectionManager, SqliteError, SqliteStorageEngine};
20
21const PAGE_SIZE: i64 = 512;
22
23type Pool = bb8::Pool<SqliteConnectionManager>;
24type BoxDumpStream = Pin<Box<dyn Stream<Item = Result<StorageDumpItem, RetrievalError>> + Send + 'static>>;
25
26#[async_trait]
27impl StorageDump for SqliteStorageEngine {
28 type DumpStream = BoxDumpStream;
29
30 async fn dump(&self) -> Result<Self::DumpStream, RetrievalError> {
31 let conn = self.pool().get().await.map_err(|error| SqliteError::Pool(error.to_string()))?;
32 let (event_collections, state_collections) = conn.with_connection(dump_collections).await?;
33 drop(conn);
34
35 let cursor = SqliteDumpCursor {
36 pool: self.pool().clone(),
37 phase: DumpPhase::Events,
38 event_collections,
39 state_collections,
40 collection_index: 0,
41 after: None,
42 pending: VecDeque::new(),
43 };
44 Ok(Box::pin(stream::try_unfold(cursor, |mut cursor| async move { Ok(cursor.next().await?.map(|item| (item, cursor))) })))
45 }
46}
47
48#[derive(Clone, Copy)]
49enum DumpPhase {
50 Events,
51 States,
52 Done,
53}
54
55struct SqliteDumpCursor {
56 pool: Pool,
57 phase: DumpPhase,
58 event_collections: Vec<CollectionId>,
59 state_collections: Vec<CollectionId>,
60 collection_index: usize,
61 after: Option<String>,
62 pending: VecDeque<StorageDumpItem>,
63}
64
65impl SqliteDumpCursor {
66 async fn next(&mut self) -> Result<Option<StorageDumpItem>, RetrievalError> {
67 loop {
68 if let Some(item) = self.pending.pop_front() {
69 return Ok(Some(item));
70 }
71 match self.phase {
72 DumpPhase::Events => {
73 let Some(collection) = self.event_collections.get(self.collection_index).cloned() else {
74 self.phase = DumpPhase::States;
75 self.collection_index = 0;
76 self.after = None;
77 continue;
78 };
79 let conn = self.pool.get().await.map_err(|error| SqliteError::Pool(error.to_string()))?;
80 let after = self.after.clone();
81 let (last, page) = conn.with_connection(move |conn| event_page(conn, &collection, after.as_deref())).await?;
82 if page.is_empty() {
83 self.collection_index += 1;
84 self.after = None;
85 continue;
86 }
87 self.after = last;
88 self.pending.extend(page);
89 }
90 DumpPhase::States => {
91 let Some(collection) = self.state_collections.get(self.collection_index).cloned() else {
92 self.phase = DumpPhase::Done;
93 continue;
94 };
95 let conn = self.pool.get().await.map_err(|error| SqliteError::Pool(error.to_string()))?;
96 let after = self.after.clone();
97 let (last, page) = conn.with_connection(move |conn| state_page(conn, &collection, after.as_deref())).await?;
98 if page.is_empty() {
99 self.collection_index += 1;
100 self.after = None;
101 continue;
102 }
103 self.after = last;
104 self.pending.extend(page);
105 }
106 DumpPhase::Done => return Ok(None),
107 }
108 }
109 }
110}
111
112fn event_page(
113 conn: &Connection,
114 collection: &CollectionId,
115 after: Option<&str>,
116) -> Result<(Option<String>, Vec<StorageDumpItem>), SqliteError> {
117 let table = quote_identifier(&format!("{collection}_event"));
118 let (query, arguments): (String, Vec<rusqlite::types::Value>) = if let Some(after) = after {
119 (
120 format!("SELECT id, entity_id, operations, parent, attestations FROM {table} WHERE id > ? ORDER BY id LIMIT ?"),
121 vec![after.to_owned().into(), PAGE_SIZE.into()],
122 )
123 } else {
124 (format!("SELECT id, entity_id, operations, parent, attestations FROM {table} ORDER BY id LIMIT ?"), vec![PAGE_SIZE.into()])
125 };
126 let mut statement = conn.prepare(&query)?;
127 let mut rows = statement.query(rusqlite::params_from_iter(arguments))?;
128 let mut last = None;
129 let mut page = Vec::new();
130 while let Some(row) = rows.next()? {
131 let stored_id: String = row.get(0)?;
132 let declared_id = EventId::from_base64(&stored_id).map_err(|error| SqliteError::Dump(error.to_string()))?;
133 let entity_id = EntityId::from_base64(row.get::<_, String>(1)?).map_err(|error| SqliteError::Dump(error.to_string()))?;
134 let operations = bincode::deserialize::<OperationSet>(&row.get::<_, Vec<u8>>(2)?)?;
135 let parent = serde_json::from_str::<Clock>(&row.get::<_, String>(3)?)?;
136 let attestations = bincode::deserialize::<AttestationSet>(&row.get::<_, Vec<u8>>(4)?)?;
137 let event = Attested { payload: Event { collection: collection.clone(), entity_id, operations, parent }, attestations };
138 if event.payload.id() != declared_id {
139 return Err(SqliteError::Dump(format!("stored event id does not match payload for {collection}/{declared_id}")));
140 }
141 last = Some(stored_id);
142 page.push(StorageDumpItem::Event(event));
143 }
144 Ok((last, page))
145}
146
147fn state_page(
148 conn: &Connection,
149 collection: &CollectionId,
150 after: Option<&str>,
151) -> Result<(Option<String>, Vec<StorageDumpItem>), SqliteError> {
152 let table = quote_identifier(collection.as_str());
153 let (query, arguments): (String, Vec<rusqlite::types::Value>) = if let Some(after) = after {
154 (
155 format!("SELECT id, state_buffer, head, attestations FROM {table} WHERE id > ? ORDER BY id LIMIT ?"),
156 vec![after.to_owned().into(), PAGE_SIZE.into()],
157 )
158 } else {
159 (format!("SELECT id, state_buffer, head, attestations FROM {table} ORDER BY id LIMIT ?"), vec![PAGE_SIZE.into()])
160 };
161 let mut statement = conn.prepare(&query)?;
162 let mut rows = statement.query(rusqlite::params_from_iter(arguments))?;
163 let mut last = None;
164 let mut page = Vec::new();
165 while let Some(row) = rows.next()? {
166 let stored_id: String = row.get(0)?;
167 let entity_id = EntityId::from_base64(&stored_id).map_err(|error| SqliteError::Dump(error.to_string()))?;
168 let state_buffers = bincode::deserialize::<BTreeMap<String, Vec<u8>>>(&row.get::<_, Vec<u8>>(1)?)?;
169 let head = serde_json::from_str::<Clock>(&row.get::<_, String>(2)?)?;
170 let attestations = bincode::deserialize::<AttestationSet>(&row.get::<_, Vec<u8>>(3)?)?;
171 last = Some(stored_id);
172 page.push(StorageDumpItem::State(Attested {
173 payload: EntityState {
174 entity_id,
175 collection: collection.clone(),
176 state: State { state_buffers: StateBuffers(state_buffers), head },
177 },
178 attestations,
179 }));
180 }
181 Ok((last, page))
182}
183
184fn dump_collections(conn: &Connection) -> Result<(Vec<CollectionId>, Vec<CollectionId>), SqliteError> {
185 let columns = table_columns(conn)?;
186 let state_columns = ["id", "state_buffer", "head", "attestations"];
187 let event_columns = ["id", "entity_id", "operations", "parent", "attestations"];
188 let mut events = Vec::new();
189 let mut states = Vec::new();
190 for (name, present) in columns {
191 if state_columns.iter().all(|column| present.contains(*column)) {
192 states.push(CollectionId::from(name.clone()));
193 }
194 if event_columns.iter().all(|column| present.contains(*column)) {
195 if let Some(collection) = name.strip_suffix("_event") {
196 events.push(CollectionId::from(collection));
197 }
198 }
199 }
200 Ok((events, states))
201}
202
203fn table_columns(conn: &Connection) -> Result<BTreeMap<String, BTreeSet<String>>, SqliteError> {
204 let mut tables = conn.prepare("SELECT name FROM sqlite_master WHERE type = 'table' AND name NOT LIKE 'sqlite_%' ORDER BY name")?;
205 let names = tables.query_map([], |row| row.get::<_, String>(0))?.collect::<Result<Vec<_>, _>>()?;
206 let mut columns = BTreeMap::new();
207 for name in names {
208 let query = format!(r#"PRAGMA table_info("{}")"#, name.replace('"', "\"\""));
209 let mut statement = conn.prepare(&query)?;
210 let present = statement.query_map([], |row| row.get::<_, String>(1))?.collect::<Result<BTreeSet<_>, _>>()?;
211 columns.insert(name, present);
212 }
213 Ok(columns)
214}
215
216fn quote_identifier(identifier: &str) -> String { format!("\"{}\"", identifier.replace('"', "\"\"")) }