dinoco_engine 2.0.5

Database adapters, query execution, and migration engine components for Dinoco.
Documentation
use std::sync::Arc;

use anyhow::{Context, anyhow};
use deadpool_sqlite::{Config, Hook, HookError, Pool, Runtime};
use rusqlite::types::{FromSql, FromSqlError, FromSqlResult, ToSqlOutput, Value, ValueRef};

mod compiler;

use crate::{
    CompiledTransactionCommand, CompiledTransactionStatement, DinocoAdapter, DinocoSqlite, DinocoValue,
    LiveTransactionMessage, RawTransactionOutput, RowDecodeError, TransactionCommandKind,
};

#[derive(Clone)]
pub struct SqliteAdapter {
    pub path: String,
    pub pool: Arc<Pool>,
    with_logger: bool,
    /// Keeps an in-memory database alive while any clone of the adapter
    /// exists, even if the pool drops every connection.
    memory_anchor: Option<Arc<std::sync::Mutex<rusqlite::Connection>>>,
}

#[async_trait::async_trait]
impl DinocoAdapter for SqliteAdapter {
    async fn new(path: String) -> Result<Self, String> {
        let path = normalize_sqlite_path(path);
        if let Some(parent) = std::path::Path::new(&path).parent() {
            std::fs::create_dir_all(parent).map_err(|err| err.to_string())?;
        }
        let cfg = Config::new(&path);
        let pool = cfg
            .builder(Runtime::Tokio1)
            .map_err(|err| err.to_string())?
            .post_create(Hook::async_fn(|connection, _| {
                Box::pin(async move {
                    match connection
                        .interact(|conn| -> rusqlite::Result<bool> {
                            conn.pragma_update(None, "foreign_keys", true)?;
                            conn.pragma_query_value(None, "foreign_keys", |row| row.get::<_, bool>(0))
                        })
                        .await
                    {
                        Ok(Ok(true)) => Ok(()),
                        Ok(Ok(false)) => Err(HookError::message(
                            "SQLite did not enable foreign key enforcement for a new connection",
                        )),
                        Ok(Err(err)) => Err(HookError::Backend(err)),
                        Err(err) => Err(HookError::message(format!(
                            "failed to configure SQLite foreign key enforcement: {err}"
                        ))),
                    }
                })
            }))
            .build()
            .map_err(|err| err.to_string())?;

        // Open one pooled connection eagerly so `connect()` guarantees that a
        // file-backed SQLite database has been created and configured.
        let connection = pool.get().await.map_err(|err| err.to_string())?;
        drop(connection);

        Ok(Self { path, pool: Arc::new(pool), with_logger: false, memory_anchor: None })
    }

    async fn query<M>(&self, query: &str, params: &[DinocoValue]) -> anyhow::Result<Vec<M>>
    where
        M: DinocoSqlite,
    {
        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
        let query_owned = query.to_string();
        let params_owned = params.to_vec();

        conn.interact(move |conn| -> anyhow::Result<Vec<M>> {
            let mut stmt = conn.prepare_cached(&query_owned)?;
            let params_refs: Vec<&dyn rusqlite::ToSql> =
                params_owned.iter().map(|p| p as &dyn rusqlite::ToSql).collect();

            let mut rows = stmt.query(params_refs.as_slice())?;
            let mut result = Vec::new();

            while let Some(row) = rows.next()? {
                let item = M::from_sqlite_row(row).ok_or_else(|| RowDecodeError::new(std::any::type_name::<M>()))?;
                result.push(item);
            }

            Ok(result)
        })
        .await
        .map_err(|err| anyhow!(err.to_string()))?
    }

    async fn query_optional<M>(&self, query: &str, params: &[DinocoValue]) -> anyhow::Result<Vec<M>>
    where
        M: DinocoSqlite,
    {
        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
        let query_owned = query.to_string();
        let params_owned = params.to_vec();

        conn.interact(move |conn| -> anyhow::Result<Vec<M>> {
            let mut stmt = conn.prepare_cached(&query_owned)?;
            let params_refs: Vec<&dyn rusqlite::ToSql> =
                params_owned.iter().map(|p| p as &dyn rusqlite::ToSql).collect();

            let mut rows = stmt.query(params_refs.as_slice())?;
            let mut result = Vec::new();

            while let Some(row) = rows.next()? {
                if let Some(item) = M::from_sqlite_row(row) {
                    result.push(item);
                }
            }

            Ok(result)
        })
        .await
        .map_err(|err| anyhow!(err.to_string()))?
    }

    async fn execute(&self, query: &str, params: &[DinocoValue]) -> anyhow::Result<usize> {
        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
        let query_owned = query.to_string();
        let params_owned = params.to_vec();

        conn.interact(move |conn| -> anyhow::Result<usize> {
            let mut stmt = conn.prepare_cached(&query_owned)?;
            let params_refs: Vec<&dyn rusqlite::ToSql> =
                params_owned.iter().map(|p| p as &dyn rusqlite::ToSql).collect();

            Ok(stmt.execute(params_refs.as_slice())?)
        })
        .await
        .map_err(|err| anyhow!(err.to_string()))?
    }
}

fn normalize_sqlite_path(path: String) -> String {
    if path == ":memory:"
        || std::path::Path::new(&path).is_absolute()
        || path.starts_with("file:")
        || path.starts_with("dinoco/")
    {
        path
    } else {
        format!("dinoco/{path}")
    }
}

impl SqliteAdapter {
    /// Opens a new, empty in-memory database that every pooled connection
    /// (and every transaction) of this adapter shares. Each call gets its own
    /// database, so parallel tests never see each other's rows. The database
    /// is dropped with the last clone of the adapter.
    pub async fn memory() -> anyhow::Result<Self> {
        static NEXT_DATABASE: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);

        let id = NEXT_DATABASE.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
        // The memdb VFS shares one in-memory database between connections by
        // name and, unlike `cache=shared`, keeps SQLite's regular locking.
        let path = format!("file:/dinoco-memory-{}-{id}?vfs=memdb", std::process::id());
        let anchor = rusqlite::Connection::open(&path).context("failed to open the in-memory SQLite database")?;
        let mut adapter = <Self as DinocoAdapter>::new(path).await.map_err(anyhow::Error::msg)?;
        adapter.memory_anchor = Some(Arc::new(std::sync::Mutex::new(anchor)));

        Ok(adapter)
    }

    pub(crate) fn set_logger(&mut self, enabled: bool) {
        self.with_logger = enabled;
    }

    pub(crate) fn logger_enabled(&self) -> bool {
        self.with_logger
    }

    pub(crate) async fn begin_live_transaction(
        &self,
    ) -> anyhow::Result<tokio::sync::mpsc::Sender<LiveTransactionMessage>> {
        let connection = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
        let compiler = self.clone();
        let (sender, mut receiver) = tokio::sync::mpsc::channel(1);
        let (ready_sender, ready_receiver) = tokio::sync::oneshot::channel();

        tokio::spawn(async move {
            let _ = connection
                .interact(move |connection| -> anyhow::Result<()> {
                    let transaction = match connection.transaction() {
                        Ok(transaction) => transaction,
                        Err(error) => {
                            let _ = ready_sender.send(Err(anyhow::Error::from(error)));
                            return Ok(());
                        }
                    };
                    let _ = ready_sender.send(Ok(()));

                    while let Some(message) = receiver.blocking_recv() {
                        match message {
                            LiveTransactionMessage::Execute { command, reply } => {
                                let result = command.compile(&compiler).and_then(|command| {
                                    execute_transaction_command(&transaction, &command)
                                        .and_then(|raw| command.finish(raw))
                                });
                                let _ = reply.send(result);
                            }
                            LiveTransactionMessage::Commit { reply } => {
                                let result = transaction.commit().map_err(anyhow::Error::from);
                                let _ = reply.send(result);
                                return Ok(());
                            }
                            LiveTransactionMessage::Rollback { reply } => {
                                let result = transaction.rollback().map_err(anyhow::Error::from);
                                let _ = reply.send(result);
                                return Ok(());
                            }
                        }
                    }

                    // Dropping an unfinished rusqlite transaction rolls it back.
                    Ok(())
                })
                .await;
        });

        ready_receiver.await.map_err(|_| anyhow!("sqlite transaction worker stopped during BEGIN"))??;
        Ok(sender)
    }

    pub async fn query_count(&self, query: &str, params: &[DinocoValue]) -> anyhow::Result<i64> {
        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
        let query_owned = query.to_string();
        let params_owned = params.to_vec();

        conn.interact(move |conn| -> anyhow::Result<i64> {
            let mut stmt = conn.prepare_cached(&query_owned)?;
            let params_refs: Vec<&dyn rusqlite::ToSql> =
                params_owned.iter().map(|p| p as &dyn rusqlite::ToSql).collect();

            Ok(stmt.query_row(params_refs.as_slice(), |row| row.get(0))?)
        })
        .await
        .map_err(|err| anyhow!(err.to_string()))?
    }

    pub async fn query_find_batch(
        &self,
        query: &str,
        params: &[DinocoValue],
        column_count: usize,
    ) -> anyhow::Result<Vec<Vec<crate::serde_json::Value>>> {
        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
        let query_owned = query.to_string();
        let params_owned = params.to_vec();

        conn.interact(move |conn| -> anyhow::Result<Vec<Vec<crate::serde_json::Value>>> {
            let mut stmt = conn.prepare_cached(&query_owned)?;
            let params_refs: Vec<&dyn rusqlite::ToSql> =
                params_owned.iter().map(|p| p as &dyn rusqlite::ToSql).collect();

            let columns = stmt.query_row(params_refs.as_slice(), |row| {
                (0..column_count).map(|index| row.get::<_, String>(index)).collect::<rusqlite::Result<Vec<String>>>()
            })?;

            columns
                .into_iter()
                .map(|text| crate::serde_json::from_str::<Vec<crate::serde_json::Value>>(&text).map_err(anyhow::Error::from))
                .collect()
        })
        .await
        .map_err(|err| anyhow!(err.to_string()))?
    }

    pub async fn query_exists(&self, query: &str, params: &[DinocoValue]) -> anyhow::Result<bool> {
        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
        let query_owned = query.to_string();
        let params_owned = params.to_vec();

        conn.interact(move |conn| -> anyhow::Result<bool> {
            let mut stmt = conn.prepare_cached(&query_owned)?;
            let params_refs: Vec<&dyn rusqlite::ToSql> =
                params_owned.iter().map(|p| p as &dyn rusqlite::ToSql).collect();

            Ok(stmt.query_row(params_refs.as_slice(), |row| row.get(0))?)
        })
        .await
        .map_err(|err| anyhow!(err.to_string()))?
    }
}

fn execute_transaction_command(
    transaction: &rusqlite::Transaction<'_>,
    command: &CompiledTransactionCommand,
) -> anyhow::Result<RawTransactionOutput> {
    let mut output = None;
    for statement in &command.statements {
        let raw = execute_transaction_statement(transaction, statement)?;
        if statement.output {
            output = Some(raw);
        }
    }

    output.ok_or_else(|| anyhow!("Dinoco transaction command contains no output statement."))
}

fn execute_transaction_statement(
    transaction: &rusqlite::Transaction<'_>,
    command: &CompiledTransactionStatement,
) -> anyhow::Result<RawTransactionOutput> {
    if command.sql.is_empty() {
        return match command.kind {
            TransactionCommandKind::Rows => Ok(RawTransactionOutput::Rows(Vec::new())),
            TransactionCommandKind::Execute => Ok(RawTransactionOutput::Affected(0)),
        };
    }

    let params_refs = command.params.iter().map(|param| param as &dyn rusqlite::ToSql).collect::<Vec<_>>();

    match command.kind {
        TransactionCommandKind::Rows => {
            let decoder = command
                .decoder
                .ok_or_else(|| anyhow!("Dinoco transaction query is missing its sqlite row decoder."))?;
            let mut statement = transaction.prepare_cached(&command.sql)?;
            let mut rows = statement.query(params_refs.as_slice())?;
            let mut values = Vec::new();

            while let Some(row) = rows.next()? {
                values.push((decoder.sqlite)(row).ok_or_else(|| RowDecodeError::new("transaction result"))?);
            }

            Ok(RawTransactionOutput::Rows(values))
        }
        TransactionCommandKind::Execute => {
            let mut statement = transaction.prepare_cached(&command.sql)?;
            Ok(RawTransactionOutput::Affected(statement.execute(params_refs.as_slice())?))
        }
    }
}

impl rusqlite::ToSql for DinocoValue {
    fn to_sql(&self) -> rusqlite::Result<ToSqlOutput<'_>> {
        match self {
            DinocoValue::Null => Ok(ToSqlOutput::Owned(Value::Null)),
            DinocoValue::Integer(i) => Ok(ToSqlOutput::Owned(Value::Integer(*i))),
            DinocoValue::Float(f) => Ok(ToSqlOutput::Owned(Value::Real(*f))),
            DinocoValue::Boolean(b) => Ok(ToSqlOutput::Owned(Value::Integer(if *b { 1 } else { 0 }))),
            DinocoValue::String(s) => Ok(ToSqlOutput::Owned(Value::Text(s.clone()))),
            DinocoValue::Enum(_, s) => Ok(ToSqlOutput::Owned(Value::Text(s.clone()))),
            DinocoValue::Bytes(v) => Ok(ToSqlOutput::Owned(Value::Blob(v.clone()))),
            DinocoValue::Json(v) => Ok(ToSqlOutput::Owned(Value::Blob(v.to_string().into_bytes()))),
            DinocoValue::DateTime(dt) => Ok(ToSqlOutput::Owned(Value::Text(dt.to_rfc3339()))),
            DinocoValue::Date(date) => Ok(ToSqlOutput::Owned(Value::Text(date.to_string()))),
        }
    }
}

impl FromSql for DinocoValue {
    fn column_result(value: ValueRef<'_>) -> FromSqlResult<Self> {
        match value {
            ValueRef::Null => Ok(DinocoValue::Null),
            ValueRef::Integer(value) => Ok(DinocoValue::Integer(value)),
            ValueRef::Real(value) => Ok(DinocoValue::Float(value)),
            ValueRef::Text(value) => String::from_utf8(value.to_vec())
                .map(DinocoValue::String)
                .map_err(|err| FromSqlError::Other(Box::new(err))),
            ValueRef::Blob(value) => Ok(DinocoValue::Bytes(value.to_vec())),
        }
    }
}