Skip to main content

dinoco_engine/backends/sqlite/
mod.rs

1use std::sync::Arc;
2
3use anyhow::{Context, anyhow};
4use deadpool_sqlite::{Config, Hook, HookError, Pool, Runtime};
5use rusqlite::types::{FromSql, FromSqlError, FromSqlResult, ToSqlOutput, Value, ValueRef};
6
7mod compiler;
8
9use crate::{
10    CompiledTransactionCommand, CompiledTransactionStatement, DinocoAdapter, DinocoSqlite, DinocoValue,
11    RawTransactionOutput, TransactionCommandKind, TransactionResults,
12};
13
14pub struct SqliteAdapter {
15    pub path: String,
16    pub pool: Arc<Pool>,
17    with_logger: bool,
18}
19
20#[async_trait::async_trait]
21impl DinocoAdapter for SqliteAdapter {
22    async fn new(path: String) -> Result<Self, String> {
23        let path = normalize_sqlite_path(path);
24        if let Some(parent) = std::path::Path::new(&path).parent() {
25            std::fs::create_dir_all(parent).map_err(|err| err.to_string())?;
26        }
27        let cfg = Config::new(&path);
28        let pool = cfg
29            .builder(Runtime::Tokio1)
30            .map_err(|err| err.to_string())?
31            .post_create(Hook::async_fn(|connection, _| {
32                Box::pin(async move {
33                    match connection
34                        .interact(|conn| -> rusqlite::Result<bool> {
35                            conn.pragma_update(None, "foreign_keys", true)?;
36                            conn.pragma_query_value(None, "foreign_keys", |row| row.get::<_, bool>(0))
37                        })
38                        .await
39                    {
40                        Ok(Ok(true)) => Ok(()),
41                        Ok(Ok(false)) => Err(HookError::message(
42                            "SQLite did not enable foreign key enforcement for a new connection",
43                        )),
44                        Ok(Err(err)) => Err(HookError::Backend(err)),
45                        Err(err) => Err(HookError::message(format!(
46                            "failed to configure SQLite foreign key enforcement: {err}"
47                        ))),
48                    }
49                })
50            }))
51            .build()
52            .map_err(|err| err.to_string())?;
53
54        // Open one pooled connection eagerly so `connect()` guarantees that a
55        // file-backed SQLite database has been created and configured.
56        let connection = pool.get().await.map_err(|err| err.to_string())?;
57        drop(connection);
58
59        Ok(Self { path, pool: Arc::new(pool), with_logger: false })
60    }
61
62    async fn query<M>(&self, query: &str, params: &[DinocoValue]) -> anyhow::Result<Vec<M>>
63    where
64        M: DinocoSqlite,
65    {
66        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
67        let query_owned = query.to_string();
68        let params_owned = params.to_vec();
69
70        conn.interact(move |conn| -> anyhow::Result<Vec<M>> {
71            let mut stmt = conn.prepare_cached(&query_owned)?;
72            let params_refs: Vec<&dyn rusqlite::ToSql> =
73                params_owned.iter().map(|p| p as &dyn rusqlite::ToSql).collect();
74
75            let mut rows = stmt.query(params_refs.as_slice())?;
76            let mut result = Vec::new();
77
78            while let Some(row) = rows.next()? {
79                let item = M::from_sqlite_row(row).ok_or_else(|| anyhow!("Failed to parse sqlite row"))?;
80                result.push(item);
81            }
82
83            Ok(result)
84        })
85        .await
86        .map_err(|err| anyhow!(err.to_string()))?
87    }
88
89    async fn query_optional<M>(&self, query: &str, params: &[DinocoValue]) -> anyhow::Result<Vec<M>>
90    where
91        M: DinocoSqlite,
92    {
93        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
94        let query_owned = query.to_string();
95        let params_owned = params.to_vec();
96
97        conn.interact(move |conn| -> anyhow::Result<Vec<M>> {
98            let mut stmt = conn.prepare_cached(&query_owned)?;
99            let params_refs: Vec<&dyn rusqlite::ToSql> =
100                params_owned.iter().map(|p| p as &dyn rusqlite::ToSql).collect();
101
102            let mut rows = stmt.query(params_refs.as_slice())?;
103            let mut result = Vec::new();
104
105            while let Some(row) = rows.next()? {
106                if let Some(item) = M::from_sqlite_row(row) {
107                    result.push(item);
108                }
109            }
110
111            Ok(result)
112        })
113        .await
114        .map_err(|err| anyhow!(err.to_string()))?
115    }
116
117    async fn execute(&self, query: &str, params: &[DinocoValue]) -> anyhow::Result<usize> {
118        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
119        let query_owned = query.to_string();
120        let params_owned = params.to_vec();
121
122        conn.interact(move |conn| -> anyhow::Result<usize> {
123            let mut stmt = conn.prepare_cached(&query_owned)?;
124            let params_refs: Vec<&dyn rusqlite::ToSql> =
125                params_owned.iter().map(|p| p as &dyn rusqlite::ToSql).collect();
126
127            Ok(stmt.execute(params_refs.as_slice())?)
128        })
129        .await
130        .map_err(|err| anyhow!(err.to_string()))?
131    }
132}
133
134fn normalize_sqlite_path(path: String) -> String {
135    if path == ":memory:"
136        || std::path::Path::new(&path).is_absolute()
137        || path.starts_with("file:")
138        || path.starts_with("dinoco/")
139    {
140        path
141    } else {
142        format!("dinoco/{path}")
143    }
144}
145
146impl SqliteAdapter {
147    pub(crate) fn set_logger(&mut self, enabled: bool) {
148        self.with_logger = enabled;
149    }
150
151    pub(crate) fn logger_enabled(&self) -> bool {
152        self.with_logger
153    }
154
155    pub(crate) async fn execute_compiled_transaction(
156        &self,
157        commands: Vec<CompiledTransactionCommand>,
158    ) -> anyhow::Result<TransactionResults> {
159        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
160
161        conn.interact(move |conn| -> anyhow::Result<TransactionResults> {
162            let transaction = conn.transaction()?;
163            let execution = (|| {
164                let mut values = Vec::with_capacity(commands.len());
165
166                for command in commands {
167                    let raw = execute_transaction_command(&transaction, &command)?;
168                    values.push(command.finish(raw)?);
169                }
170
171                Ok(values)
172            })();
173
174            match execution {
175                Ok(values) => {
176                    transaction.commit()?;
177                    Ok(TransactionResults::new(values))
178                }
179                Err(error) => {
180                    transaction.rollback().context("Failed to roll back sqlite transaction")?;
181                    Err(error)
182                }
183            }
184        })
185        .await
186        .map_err(|err| anyhow!(err.to_string()))?
187    }
188
189    pub async fn query_count(&self, query: &str, params: &[DinocoValue]) -> anyhow::Result<i64> {
190        let conn = self.pool.get().await.context("Failed to get sqlite connection from pool")?;
191        let query_owned = query.to_string();
192        let params_owned = params.to_vec();
193
194        conn.interact(move |conn| -> anyhow::Result<i64> {
195            let mut stmt = conn.prepare_cached(&query_owned)?;
196            let params_refs: Vec<&dyn rusqlite::ToSql> =
197                params_owned.iter().map(|p| p as &dyn rusqlite::ToSql).collect();
198
199            Ok(stmt.query_row(params_refs.as_slice(), |row| row.get(0))?)
200        })
201        .await
202        .map_err(|err| anyhow!(err.to_string()))?
203    }
204}
205
206fn execute_transaction_command(
207    transaction: &rusqlite::Transaction<'_>,
208    command: &CompiledTransactionCommand,
209) -> anyhow::Result<RawTransactionOutput> {
210    let mut output = None;
211    for statement in &command.statements {
212        let raw = execute_transaction_statement(transaction, statement)?;
213        if statement.output {
214            output = Some(raw);
215        }
216    }
217
218    output.ok_or_else(|| anyhow!("Dinoco transaction command contains no output statement."))
219}
220
221fn execute_transaction_statement(
222    transaction: &rusqlite::Transaction<'_>,
223    command: &CompiledTransactionStatement,
224) -> anyhow::Result<RawTransactionOutput> {
225    if command.sql.is_empty() {
226        return match command.kind {
227            TransactionCommandKind::Rows => Ok(RawTransactionOutput::Rows(Vec::new())),
228            TransactionCommandKind::Execute => Ok(RawTransactionOutput::Affected(0)),
229            TransactionCommandKind::Count => Ok(RawTransactionOutput::Count(0)),
230        };
231    }
232
233    let params_refs = command.params.iter().map(|param| param as &dyn rusqlite::ToSql).collect::<Vec<_>>();
234
235    match command.kind {
236        TransactionCommandKind::Rows => {
237            let decoder = command
238                .decoder
239                .ok_or_else(|| anyhow!("Dinoco transaction query is missing its sqlite row decoder."))?;
240            let mut statement = transaction.prepare_cached(&command.sql)?;
241            let mut rows = statement.query(params_refs.as_slice())?;
242            let mut values = Vec::new();
243
244            while let Some(row) = rows.next()? {
245                values.push((decoder.sqlite)(row).ok_or_else(|| anyhow!("Failed to parse sqlite transaction row"))?);
246            }
247
248            Ok(RawTransactionOutput::Rows(values))
249        }
250        TransactionCommandKind::Execute => {
251            let mut statement = transaction.prepare_cached(&command.sql)?;
252            Ok(RawTransactionOutput::Affected(statement.execute(params_refs.as_slice())?))
253        }
254        TransactionCommandKind::Count => {
255            let mut statement = transaction.prepare_cached(&command.sql)?;
256            let total = statement.query_row(params_refs.as_slice(), |row| row.get(0))?;
257            Ok(RawTransactionOutput::Count(total))
258        }
259    }
260}
261
262impl rusqlite::ToSql for DinocoValue {
263    fn to_sql(&self) -> rusqlite::Result<ToSqlOutput<'_>> {
264        match self {
265            DinocoValue::Null => Ok(ToSqlOutput::Owned(Value::Null)),
266            DinocoValue::Integer(i) => Ok(ToSqlOutput::Owned(Value::Integer(*i))),
267            DinocoValue::Float(f) => Ok(ToSqlOutput::Owned(Value::Real(*f))),
268            DinocoValue::Boolean(b) => Ok(ToSqlOutput::Owned(Value::Integer(if *b { 1 } else { 0 }))),
269            DinocoValue::String(s) => Ok(ToSqlOutput::Owned(Value::Text(s.clone()))),
270            DinocoValue::Enum(_, s) => Ok(ToSqlOutput::Owned(Value::Text(s.clone()))),
271            DinocoValue::Bytes(v) => Ok(ToSqlOutput::Owned(Value::Blob(v.clone()))),
272            DinocoValue::Json(v) => Ok(ToSqlOutput::Owned(Value::Blob(v.to_string().into_bytes()))),
273            DinocoValue::DateTime(dt) => Ok(ToSqlOutput::Owned(Value::Text(dt.to_rfc3339()))),
274            DinocoValue::Date(date) => Ok(ToSqlOutput::Owned(Value::Text(date.to_string()))),
275        }
276    }
277}
278
279impl FromSql for DinocoValue {
280    fn column_result(value: ValueRef<'_>) -> FromSqlResult<Self> {
281        match value {
282            ValueRef::Null => Ok(DinocoValue::Null),
283            ValueRef::Integer(value) => Ok(DinocoValue::Integer(value)),
284            ValueRef::Real(value) => Ok(DinocoValue::Float(value)),
285            ValueRef::Text(value) => String::from_utf8(value.to_vec())
286                .map(DinocoValue::String)
287                .map_err(|err| FromSqlError::Other(Box::new(err))),
288            ValueRef::Blob(value) => Ok(DinocoValue::Bytes(value.to_vec())),
289        }
290    }
291}