1use std::num::NonZeroUsize;
10use std::path::{Path, PathBuf};
11
12use deadpool::Runtime;
13use deadpool::managed::{Manager, Metrics, Object, Pool, RecycleError, RecycleResult};
14use deadpool_sync::SyncWrapper;
15use miden_node_tracing::Instrument;
16use rusqlite::{Connection, OpenFlags, TransactionBehavior};
17
18use crate::sqlite::tx::{ReadTx, WriteTx};
19use crate::{DatabaseError, default_connection_pool_size};
20
21const STATEMENT_CACHE_CAPACITY: usize = 512;
25
26#[derive(Debug, thiserror::Error)]
34pub(crate) enum SqliteManagerError {
35 #[error("failed to open the sqlite database")]
37 Open(#[source] rusqlite::Error),
38 #[error("failed to configure the sqlite connection")]
40 Configure(#[source] rusqlite::Error),
41 #[error("the pooled sqlite connection is poisoned")]
43 Poisoned,
44}
45
46struct SqliteManager {
47 path: PathBuf,
48 read_only: bool,
51}
52
53impl Manager for SqliteManager {
54 type Type = SyncWrapper<Connection>;
55 type Error = SqliteManagerError;
56
57 async fn create(&self) -> Result<Self::Type, Self::Error> {
58 let path = self.path.clone();
59 let read_only = self.read_only;
60 SyncWrapper::new(Runtime::Tokio1, move || {
61 let conn = Connection::open_with_flags(&path, OpenFlags::SQLITE_OPEN_READ_WRITE)
62 .map_err(SqliteManagerError::Open)?;
63 configure_connection(&conn, read_only).map_err(SqliteManagerError::Configure)?;
64 Ok(conn)
65 })
66 .await
67 }
68
69 async fn recycle(
70 &self,
71 conn: &mut Self::Type,
72 _metrics: &Metrics,
73 ) -> RecycleResult<Self::Error> {
74 if conn.is_mutex_poisoned() {
75 return Err(RecycleError::Backend(SqliteManagerError::Poisoned));
76 }
77 conn.interact(|conn| {
80 if !conn.is_autocommit() {
81 let _ = conn.execute_batch("ROLLBACK");
82 }
83 })
84 .await
85 .map_err(|_| RecycleError::Backend(SqliteManagerError::Poisoned))?;
86 Ok(())
87 }
88}
89
90fn configure_connection(conn: &Connection, read_only: bool) -> rusqlite::Result<()> {
96 if read_only {
99 conn.execute_batch(
102 "PRAGMA busy_timeout = 5000;
103 PRAGMA foreign_keys = ON;
104 PRAGMA query_only = ON;",
105 )?;
106 } else {
107 conn.execute_batch(
109 "PRAGMA busy_timeout = 5000;
110 PRAGMA journal_mode = WAL;
111 PRAGMA foreign_keys = ON;",
112 )?;
113 }
114 conn.set_prepared_statement_cache_capacity(STATEMENT_CACHE_CAPACITY);
115 rusqlite::vtab::array::load_module(conn)?;
118 Ok(())
119}
120
121pub fn open(database_filepath: &Path) -> Result<(DbWriter, DbReader), DatabaseError> {
128 open_with_pool_size(database_filepath, default_connection_pool_size())
129}
130
131pub fn open_with_pool_size(
139 database_filepath: &Path,
140 connection_pool_size: NonZeroUsize,
141) -> Result<(DbWriter, DbReader), DatabaseError> {
142 let writer = Pool::builder(SqliteManager {
143 path: database_filepath.to_path_buf(),
144 read_only: false,
145 })
146 .max_size(1)
147 .build()?;
148 let readers = Pool::builder(SqliteManager {
149 path: database_filepath.to_path_buf(),
150 read_only: true,
151 })
152 .max_size(connection_pool_size.get())
153 .build()?;
154 Ok((DbWriter { writer }, DbReader { readers }))
155}
156
157#[derive(Clone)]
162pub struct DbReader {
163 readers: Pool<SqliteManager>,
164}
165
166impl DbReader {
167 async fn checkout_reader(&self) -> Result<Object<SqliteManager>, DatabaseError> {
169 self.readers
170 .get()
171 .in_current_span()
172 .await
173 .map_err(|err| DatabaseError::ConnectionPoolObtainError(Box::new(err)))
174 }
175
176 pub async fn read<R, E, F>(&self, msg: impl ToString + Send, query: F) -> Result<R, E>
179 where
180 F: FnOnce(&ReadTx<'_>) -> Result<R, E> + Send + 'static,
181 R: Send + 'static,
182 E: From<DatabaseError> + Send + 'static,
183 {
184 let conn = self.checkout_reader().await.map_err(E::from)?;
185 let msg = msg.to_string();
186 let span = miden_node_tracing::Span::current();
187 conn.interact(move |conn| {
188 let _guard = span.enter();
189 let tx = conn
190 .transaction_with_behavior(TransactionBehavior::Deferred)
191 .map_err(|err| E::from(DatabaseError::from(err)))?;
192 query(&ReadTx::new(&tx))
193 })
195 .await
196 .map_err(|err| E::from(DatabaseError::interact(&msg, &err)))?
197 }
198
199 pub async fn begin_read(&self) -> Result<ReadTransaction, DatabaseError> {
202 let conn = self.checkout_reader().await?;
203 run_tx_stmt(&conn, "BEGIN DEFERRED").await?;
204 Ok(ReadTransaction { conn })
205 }
206}
207
208pub struct DbWriter {
215 writer: Pool<SqliteManager>,
216}
217
218impl DbWriter {
219 async fn checkout_writer(&self) -> Result<Object<SqliteManager>, DatabaseError> {
221 self.writer
222 .get()
223 .in_current_span()
224 .await
225 .map_err(|err| DatabaseError::ConnectionPoolObtainError(Box::new(err)))
226 }
227
228 pub async fn write<R, E, F>(&self, msg: impl ToString + Send, query: F) -> Result<R, E>
231 where
232 F: FnOnce(&WriteTx<'_>) -> Result<R, E> + Send + 'static,
233 R: Send + 'static,
234 E: From<DatabaseError> + Send + 'static,
235 {
236 let conn = self.checkout_writer().await.map_err(E::from)?;
237 let msg = msg.to_string();
238 let span = miden_node_tracing::Span::current();
239 conn.interact(move |conn| {
240 let _guard = span.enter();
241 let tx = conn
242 .transaction_with_behavior(TransactionBehavior::Immediate)
243 .map_err(|err| E::from(DatabaseError::from(err)))?;
244 let result = query(&WriteTx::new(&tx))?;
245 tx.commit().map_err(|err| E::from(DatabaseError::from(err)))?;
246 Ok(result)
247 })
248 .await
249 .map_err(|err| E::from(DatabaseError::interact(&msg, &err)))?
250 }
251
252 pub async fn begin_write(&self) -> Result<WriteTransaction, DatabaseError> {
256 let conn = self.checkout_writer().await?;
257 run_tx_stmt(&conn, "BEGIN IMMEDIATE").await?;
258 Ok(WriteTransaction { conn })
259 }
260}
261
262async fn run_tx_stmt(
267 conn: &Object<SqliteManager>,
268 stmt: &'static str,
269) -> Result<(), DatabaseError> {
270 conn.interact(move |conn| conn.execute_batch(stmt))
271 .await
272 .map_err(|err| DatabaseError::interact(stmt, &err))?
273 .map_err(DatabaseError::from)
274}
275
276pub struct ReadTransaction {
283 conn: Object<SqliteManager>,
284}
285
286impl ReadTransaction {
287 pub async fn run<R, E, F>(&self, msg: impl ToString + Send, query: F) -> Result<R, E>
289 where
290 F: FnOnce(&ReadTx<'_>) -> Result<R, E> + Send + 'static,
291 R: Send + 'static,
292 E: From<DatabaseError> + Send + 'static,
293 {
294 let msg = msg.to_string();
295 let span = miden_node_tracing::Span::current();
296 self.conn
297 .interact(move |conn| {
298 let _guard = span.enter();
299 query(&ReadTx::new(conn))
300 })
301 .await
302 .map_err(|err| E::from(DatabaseError::interact(&msg, &err)))?
303 }
304
305 pub async fn close(self) -> Result<(), DatabaseError> {
307 run_tx_stmt(&self.conn, "ROLLBACK").await
308 }
309}
310
311pub struct WriteTransaction {
322 conn: Object<SqliteManager>,
323}
324
325impl WriteTransaction {
326 pub async fn run<R, E, F>(&self, msg: impl ToString + Send, query: F) -> Result<R, E>
328 where
329 F: FnOnce(&WriteTx<'_>) -> Result<R, E> + Send + 'static,
330 R: Send + 'static,
331 E: From<DatabaseError> + Send + 'static,
332 {
333 let msg = msg.to_string();
334 let span = miden_node_tracing::Span::current();
335 self.conn
336 .interact(move |conn| {
337 let _guard = span.enter();
338 query(&WriteTx::new(conn))
339 })
340 .await
341 .map_err(|err| E::from(DatabaseError::interact(&msg, &err)))?
342 }
343
344 pub async fn commit(self) -> Result<(), DatabaseError> {
346 run_tx_stmt(&self.conn, "COMMIT").await
347 }
348
349 pub async fn rollback(self) -> Result<(), DatabaseError> {
351 run_tx_stmt(&self.conn, "ROLLBACK").await
352 }
353}
354
355#[cfg(test)]
356mod tests {
357 use std::num::NonZeroUsize;
358 use std::path::{Path, PathBuf};
359
360 use rusqlite::Connection;
361
362 use super::{DbReader, DbWriter, open_with_pool_size};
363 use crate::DatabaseError;
364
365 struct TempDb {
368 path: PathBuf,
369 }
370
371 impl TempDb {
372 fn new(name: &str) -> Self {
373 let path = std::env::temp_dir()
374 .join(format!("miden-node-db-pool-{name}-{}.sqlite3", std::process::id()));
375 let db = Self { path };
376 db.remove_files();
377 let conn = Connection::open(&db.path).expect("create db file");
378 conn.execute_batch("CREATE TABLE items (id INTEGER PRIMARY KEY);")
379 .expect("create table");
380 db
381 }
382
383 fn path(&self) -> &Path {
384 &self.path
385 }
386
387 fn remove_files(&self) {
388 let _ = fs_err::remove_file(&self.path);
389 let _ = fs_err::remove_file(self.path.with_extension("sqlite3-wal"));
390 let _ = fs_err::remove_file(self.path.with_extension("sqlite3-shm"));
391 }
392 }
393
394 impl Drop for TempDb {
395 fn drop(&mut self) {
396 self.remove_files();
397 }
398 }
399
400 fn open_db(temp: &TempDb) -> (DbWriter, DbReader) {
401 open_with_pool_size(temp.path(), NonZeroUsize::new(4).unwrap()).unwrap()
402 }
403
404 async fn count_items(reader: &DbReader) -> i64 {
405 reader
406 .read::<_, DatabaseError, _>("count", |r| {
407 Ok(r.query("SELECT COUNT(*) FROM items", &[], |row| row.get::<i64>(0))?
408 .into_iter()
409 .next()
410 .unwrap_or(0))
411 })
412 .await
413 .unwrap()
414 }
415
416 async fn insert_committed(writer: &DbWriter, id: i64) {
417 let tx = writer.begin_write().await.unwrap();
418 tx.run::<_, DatabaseError, _>("insert", move |w| {
419 w.execute("INSERT INTO items (id) VALUES (?1)", &[&id])?;
420 Ok(())
421 })
422 .await
423 .unwrap();
424 tx.commit().await.unwrap();
425 }
426
427 #[tokio::test]
428 async fn held_write_transaction_commits_across_awaits() {
429 let temp = TempDb::new("commit");
430 let (writer, reader) = open_db(&temp);
431
432 let tx = writer.begin_write().await.unwrap();
433 tx.run::<_, DatabaseError, _>("insert-1", |w| {
434 w.execute("INSERT INTO items (id) VALUES (?1)", &[&1i64])?;
435 Ok(())
436 })
437 .await
438 .unwrap();
439
440 tokio::task::yield_now().await;
442
443 tx.run::<_, DatabaseError, _>("insert-2", |w| {
444 w.execute("INSERT INTO items (id) VALUES (?1)", &[&2i64])?;
445 Ok(())
446 })
447 .await
448 .unwrap();
449
450 tx.commit().await.unwrap();
451
452 assert_eq!(count_items(&reader).await, 2);
453 }
454
455 #[tokio::test]
456 async fn dropped_write_transaction_rolls_back() {
457 let temp = TempDb::new("rollback");
458 let (writer, reader) = open_db(&temp);
459
460 {
461 let tx = writer.begin_write().await.unwrap();
462 tx.run::<_, DatabaseError, _>("insert", |w| {
463 w.execute("INSERT INTO items (id) VALUES (?1)", &[&1i64])?;
464 Ok(())
465 })
466 .await
467 .unwrap();
468 }
470
471 insert_committed(&writer, 2).await;
475 assert_eq!(count_items(&reader).await, 1);
476 }
477
478 #[tokio::test]
479 async fn reads_proceed_while_write_transaction_is_held() {
480 let temp = TempDb::new("concurrent");
481 let (writer, reader) = open_db(&temp);
482 insert_committed(&writer, 1).await;
483
484 let tx = writer.begin_write().await.unwrap();
486 tx.run::<_, DatabaseError, _>("insert-uncommitted", |w| {
487 w.execute("INSERT INTO items (id) VALUES (?1)", &[&2i64])?;
488 Ok(())
489 })
490 .await
491 .unwrap();
492
493 assert_eq!(count_items(&reader).await, 1);
496
497 tx.commit().await.unwrap();
498 assert_eq!(count_items(&reader).await, 2);
499 }
500
501 #[tokio::test]
502 async fn reader_connections_are_query_only() {
503 let temp = TempDb::new("query_only");
504 let (_writer, reader) = open_db(&temp);
505
506 let query_only = reader
507 .read::<_, DatabaseError, _>("pragma", |r| {
508 Ok(r.query("PRAGMA query_only", &[], |row| row.get::<i64>(0))?
509 .into_iter()
510 .next()
511 .unwrap_or(0))
512 })
513 .await
514 .unwrap();
515 assert_eq!(query_only, 1, "reader connections must be query_only");
516
517 let result = reader
519 .read::<(), DatabaseError, _>("rejected-write", |r| {
520 r.query("INSERT INTO items (id) VALUES (99)", &[], |_| Ok(()))?;
521 Ok(())
522 })
523 .await;
524 assert!(result.is_err(), "writes on a reader connection must fail");
525 }
526}