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,
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())?;
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 {
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);
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(());
}
}
}
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())),
}
}
}