use std::num::NonZeroUsize;
use std::path::{Path, PathBuf};
use deadpool::Runtime;
use deadpool::managed::{Manager, Metrics, Object, Pool, RecycleError, RecycleResult};
use deadpool_sync::SyncWrapper;
use rusqlite::{Connection, OpenFlags, TransactionBehavior};
use tracing::Instrument;
use crate::sqlite::tx::{ReadTx, WriteTx};
use crate::{DatabaseError, default_connection_pool_size};
const STATEMENT_CACHE_CAPACITY: usize = 512;
#[derive(Debug, thiserror::Error)]
pub(crate) enum SqliteManagerError {
#[error("failed to open the sqlite database")]
Open(#[source] rusqlite::Error),
#[error("failed to configure the sqlite connection")]
Configure(#[source] rusqlite::Error),
#[error("the pooled sqlite connection is poisoned")]
Poisoned,
}
struct SqliteManager {
path: PathBuf,
read_only: bool,
}
impl Manager for SqliteManager {
type Type = SyncWrapper<Connection>;
type Error = SqliteManagerError;
async fn create(&self) -> Result<Self::Type, Self::Error> {
let path = self.path.clone();
let read_only = self.read_only;
SyncWrapper::new(Runtime::Tokio1, move || {
let conn = Connection::open_with_flags(&path, OpenFlags::SQLITE_OPEN_READ_WRITE)
.map_err(SqliteManagerError::Open)?;
configure_connection(&conn, read_only).map_err(SqliteManagerError::Configure)?;
Ok(conn)
})
.await
}
async fn recycle(
&self,
conn: &mut Self::Type,
_metrics: &Metrics,
) -> RecycleResult<Self::Error> {
if conn.is_mutex_poisoned() {
return Err(RecycleError::Backend(SqliteManagerError::Poisoned));
}
conn.interact(|conn| {
if !conn.is_autocommit() {
let _ = conn.execute_batch("ROLLBACK");
}
})
.await
.map_err(|_| RecycleError::Backend(SqliteManagerError::Poisoned))?;
Ok(())
}
}
fn configure_connection(conn: &Connection, read_only: bool) -> rusqlite::Result<()> {
if read_only {
conn.execute_batch(
"PRAGMA busy_timeout = 5000;
PRAGMA foreign_keys = ON;
PRAGMA query_only = ON;",
)?;
} else {
conn.execute_batch(
"PRAGMA busy_timeout = 5000;
PRAGMA journal_mode = WAL;
PRAGMA foreign_keys = ON;",
)?;
}
conn.set_prepared_statement_cache_capacity(STATEMENT_CACHE_CAPACITY);
rusqlite::vtab::array::load_module(conn)?;
Ok(())
}
#[derive(Clone)]
pub struct Database {
writer: Pool<SqliteManager>,
readers: Pool<SqliteManager>,
}
impl Database {
pub fn new(database_filepath: &Path) -> Result<Self, DatabaseError> {
Self::new_with_pool_size(database_filepath, default_connection_pool_size())
}
pub fn new_with_pool_size(
database_filepath: &Path,
connection_pool_size: NonZeroUsize,
) -> Result<Self, DatabaseError> {
let writer = Pool::builder(SqliteManager {
path: database_filepath.to_path_buf(),
read_only: false,
})
.max_size(1)
.build()?;
let readers = Pool::builder(SqliteManager {
path: database_filepath.to_path_buf(),
read_only: true,
})
.max_size(connection_pool_size.get())
.build()?;
Ok(Self { writer, readers })
}
async fn checkout_writer(&self) -> Result<Object<SqliteManager>, DatabaseError> {
self.writer
.get()
.in_current_span()
.await
.map_err(|err| DatabaseError::ConnectionPoolObtainError(Box::new(err)))
}
async fn checkout_reader(&self) -> Result<Object<SqliteManager>, DatabaseError> {
self.readers
.get()
.in_current_span()
.await
.map_err(|err| DatabaseError::ConnectionPoolObtainError(Box::new(err)))
}
pub async fn read<R, E, F>(&self, msg: impl ToString + Send, query: F) -> Result<R, E>
where
F: FnOnce(&ReadTx<'_>) -> Result<R, E> + Send + 'static,
R: Send + 'static,
E: From<DatabaseError> + Send + 'static,
{
let conn = self.checkout_reader().await.map_err(E::from)?;
let msg = msg.to_string();
let span = tracing::Span::current();
conn.interact(move |conn| {
let _guard = span.enter();
let tx = conn
.transaction_with_behavior(TransactionBehavior::Deferred)
.map_err(|err| E::from(DatabaseError::from(err)))?;
query(&ReadTx::new(&tx))
})
.await
.map_err(|err| E::from(DatabaseError::interact(&msg, &err)))?
}
pub async fn write<R, E, F>(&self, msg: impl ToString + Send, query: F) -> Result<R, E>
where
F: FnOnce(&WriteTx<'_>) -> Result<R, E> + Send + 'static,
R: Send + 'static,
E: From<DatabaseError> + Send + 'static,
{
let conn = self.checkout_writer().await.map_err(E::from)?;
let msg = msg.to_string();
let span = tracing::Span::current();
conn.interact(move |conn| {
let _guard = span.enter();
let tx = conn
.transaction_with_behavior(TransactionBehavior::Immediate)
.map_err(|err| E::from(DatabaseError::from(err)))?;
let result = query(&WriteTx::new(&tx))?;
tx.commit().map_err(|err| E::from(DatabaseError::from(err)))?;
Ok(result)
})
.await
.map_err(|err| E::from(DatabaseError::interact(&msg, &err)))?
}
pub async fn begin_read(&self) -> Result<ReadTransaction, DatabaseError> {
let conn = self.checkout_reader().await?;
run_tx_stmt(&conn, "BEGIN DEFERRED").await?;
Ok(ReadTransaction { conn })
}
pub async fn begin_write(&self) -> Result<WriteTransaction, DatabaseError> {
let conn = self.checkout_writer().await?;
run_tx_stmt(&conn, "BEGIN IMMEDIATE").await?;
Ok(WriteTransaction { conn })
}
}
async fn run_tx_stmt(
conn: &Object<SqliteManager>,
stmt: &'static str,
) -> Result<(), DatabaseError> {
conn.interact(move |conn| conn.execute_batch(stmt))
.await
.map_err(|err| DatabaseError::interact(stmt, &err))?
.map_err(DatabaseError::from)
}
pub struct ReadTransaction {
conn: Object<SqliteManager>,
}
impl ReadTransaction {
pub async fn run<R, E, F>(&self, msg: impl ToString + Send, query: F) -> Result<R, E>
where
F: FnOnce(&ReadTx<'_>) -> Result<R, E> + Send + 'static,
R: Send + 'static,
E: From<DatabaseError> + Send + 'static,
{
let msg = msg.to_string();
let span = tracing::Span::current();
self.conn
.interact(move |conn| {
let _guard = span.enter();
query(&ReadTx::new(conn))
})
.await
.map_err(|err| E::from(DatabaseError::interact(&msg, &err)))?
}
pub async fn close(self) -> Result<(), DatabaseError> {
run_tx_stmt(&self.conn, "ROLLBACK").await
}
}
pub struct WriteTransaction {
conn: Object<SqliteManager>,
}
impl WriteTransaction {
pub async fn run<R, E, F>(&self, msg: impl ToString + Send, query: F) -> Result<R, E>
where
F: FnOnce(&WriteTx<'_>) -> Result<R, E> + Send + 'static,
R: Send + 'static,
E: From<DatabaseError> + Send + 'static,
{
let msg = msg.to_string();
let span = tracing::Span::current();
self.conn
.interact(move |conn| {
let _guard = span.enter();
query(&WriteTx::new(conn))
})
.await
.map_err(|err| E::from(DatabaseError::interact(&msg, &err)))?
}
pub async fn commit(self) -> Result<(), DatabaseError> {
run_tx_stmt(&self.conn, "COMMIT").await
}
pub async fn rollback(self) -> Result<(), DatabaseError> {
run_tx_stmt(&self.conn, "ROLLBACK").await
}
}
#[cfg(test)]
mod tests {
use std::num::NonZeroUsize;
use std::path::{Path, PathBuf};
use rusqlite::Connection;
use super::Database;
use crate::DatabaseError;
struct TempDb {
path: PathBuf,
}
impl TempDb {
fn new(name: &str) -> Self {
let path = std::env::temp_dir()
.join(format!("miden-node-db-pool-{name}-{}.sqlite3", std::process::id()));
let db = Self { path };
db.remove_files();
let conn = Connection::open(&db.path).expect("create db file");
conn.execute_batch("CREATE TABLE items (id INTEGER PRIMARY KEY);")
.expect("create table");
db
}
fn path(&self) -> &Path {
&self.path
}
fn remove_files(&self) {
let _ = fs_err::remove_file(&self.path);
let _ = fs_err::remove_file(self.path.with_extension("sqlite3-wal"));
let _ = fs_err::remove_file(self.path.with_extension("sqlite3-shm"));
}
}
impl Drop for TempDb {
fn drop(&mut self) {
self.remove_files();
}
}
fn open_db(temp: &TempDb) -> Database {
Database::new_with_pool_size(temp.path(), NonZeroUsize::new(4).unwrap()).unwrap()
}
async fn count_items(db: &Database) -> i64 {
db.read::<_, DatabaseError, _>("count", |r| {
Ok(r.query("SELECT COUNT(*) FROM items", &[], |row| row.get::<i64>(0))?
.into_iter()
.next()
.unwrap_or(0))
})
.await
.unwrap()
}
async fn insert_committed(db: &Database, id: i64) {
let tx = db.begin_write().await.unwrap();
tx.run::<_, DatabaseError, _>("insert", move |w| {
w.execute("INSERT INTO items (id) VALUES (?1)", &[&id])?;
Ok(())
})
.await
.unwrap();
tx.commit().await.unwrap();
}
#[tokio::test]
async fn held_write_transaction_commits_across_awaits() {
let temp = TempDb::new("commit");
let db = open_db(&temp);
let tx = db.begin_write().await.unwrap();
tx.run::<_, DatabaseError, _>("insert-1", |w| {
w.execute("INSERT INTO items (id) VALUES (?1)", &[&1i64])?;
Ok(())
})
.await
.unwrap();
tokio::task::yield_now().await;
tx.run::<_, DatabaseError, _>("insert-2", |w| {
w.execute("INSERT INTO items (id) VALUES (?1)", &[&2i64])?;
Ok(())
})
.await
.unwrap();
tx.commit().await.unwrap();
assert_eq!(count_items(&db).await, 2);
}
#[tokio::test]
async fn dropped_write_transaction_rolls_back() {
let temp = TempDb::new("rollback");
let db = open_db(&temp);
{
let tx = db.begin_write().await.unwrap();
tx.run::<_, DatabaseError, _>("insert", |w| {
w.execute("INSERT INTO items (id) VALUES (?1)", &[&1i64])?;
Ok(())
})
.await
.unwrap();
}
insert_committed(&db, 2).await;
assert_eq!(count_items(&db).await, 1);
}
#[tokio::test]
async fn reads_proceed_while_write_transaction_is_held() {
let temp = TempDb::new("concurrent");
let db = open_db(&temp);
insert_committed(&db, 1).await;
let tx = db.begin_write().await.unwrap();
tx.run::<_, DatabaseError, _>("insert-uncommitted", |w| {
w.execute("INSERT INTO items (id) VALUES (?1)", &[&2i64])?;
Ok(())
})
.await
.unwrap();
assert_eq!(count_items(&db).await, 1);
tx.commit().await.unwrap();
assert_eq!(count_items(&db).await, 2);
}
#[tokio::test]
async fn reader_connections_are_query_only() {
let temp = TempDb::new("query_only");
let db = open_db(&temp);
let query_only = db
.read::<_, DatabaseError, _>("pragma", |r| {
Ok(r.query("PRAGMA query_only", &[], |row| row.get::<i64>(0))?
.into_iter()
.next()
.unwrap_or(0))
})
.await
.unwrap();
assert_eq!(query_only, 1, "reader connections must be query_only");
let result = db
.read::<(), DatabaseError, _>("rejected-write", |r| {
r.query("INSERT INTO items (id) VALUES (99)", &[], |_| Ok(()))?;
Ok(())
})
.await;
assert!(result.is_err(), "writes on a reader connection must fail");
}
}