Skip to main content

a3s_orm/drivers/sqlite/
executor.rs

1use std::future::Future;
2use std::path::Path;
3use std::pin::Pin;
4use std::sync::Arc;
5
6use async_trait::async_trait;
7use tokio_rusqlite::rusqlite;
8use tokio_rusqlite::rusqlite::types::{Value as SqliteValue, ValueRef};
9
10use crate::{
11    CompiledQuery, ExecuteResult, Executor, QueryResult, Transaction, TransactionManager, Value,
12};
13
14use super::{SqliteError, SqliteOptions, SqliteRow, SqliteTransaction, SqliteTransactionError};
15
16#[derive(Clone)]
17pub struct SqliteExecutor {
18    pub(super) connection: tokio_rusqlite::Connection,
19    pub(super) transaction_lock: Arc<tokio::sync::Mutex<()>>,
20}
21
22impl SqliteExecutor {
23    pub async fn open(path: impl AsRef<Path>) -> Result<Self, SqliteError> {
24        Self::open_with_options(path, SqliteOptions::default()).await
25    }
26
27    pub async fn open_with_options(
28        path: impl AsRef<Path>,
29        options: SqliteOptions,
30    ) -> Result<Self, SqliteError> {
31        let executor = Self {
32            connection: tokio_rusqlite::Connection::open(path).await?,
33            transaction_lock: Arc::new(tokio::sync::Mutex::new(())),
34        };
35        executor.configure(options).await?;
36        Ok(executor)
37    }
38
39    pub async fn open_in_memory() -> Result<Self, SqliteError> {
40        let executor = Self {
41            connection: tokio_rusqlite::Connection::open_in_memory().await?,
42            transaction_lock: Arc::new(tokio::sync::Mutex::new(())),
43        };
44        executor.configure(SqliteOptions::in_memory()).await?;
45        Ok(executor)
46    }
47
48    pub fn connection(&self) -> &tokio_rusqlite::Connection {
49        &self.connection
50    }
51
52    async fn configure(&self, options: SqliteOptions) -> Result<(), SqliteError> {
53        self.connection
54            .call(move |connection| {
55                connection.busy_timeout(options.busy_timeout)?;
56                connection.pragma_update(
57                    None,
58                    "foreign_keys",
59                    if options.foreign_keys { "ON" } else { "OFF" },
60                )?;
61                connection.pragma_update(None, "journal_mode", options.journal_mode.as_sql())?;
62                Ok(())
63            })
64            .await?;
65        Ok(())
66    }
67
68    /// Run an operation inside a transaction and always complete it.
69    ///
70    /// The operation is committed on success and rolled back on error. If the
71    /// calling task is cancelled while the operation is running, dropping the
72    /// transaction schedules a rollback while retaining the connection gate.
73    pub async fn transaction<T, E, F>(&self, operation: F) -> Result<T, SqliteTransactionError<E>>
74    where
75        T: Send,
76        E: std::error::Error + Send + Sync + 'static,
77        F: for<'a> FnOnce(
78            &'a SqliteTransaction,
79        ) -> Pin<Box<dyn Future<Output = Result<T, E>> + Send + 'a>>,
80    {
81        let transaction = self.begin().await.map_err(SqliteTransactionError::Begin)?;
82        match operation(&transaction).await {
83            Ok(value) => {
84                transaction
85                    .commit()
86                    .await
87                    .map_err(SqliteTransactionError::Commit)?;
88                Ok(value)
89            }
90            Err(operation) => match transaction.rollback().await {
91                Ok(()) => Err(SqliteTransactionError::Operation(operation)),
92                Err(rollback) => Err(SqliteTransactionError::OperationAndRollback {
93                    operation,
94                    rollback,
95                }),
96            },
97        }
98    }
99
100    /// Execute trusted schema SQL. Application values should use typed queries.
101    pub async fn execute_schema(&self, sql: impl Into<String>) -> Result<(), SqliteError> {
102        let _guard = self.transaction_lock.lock().await;
103        let sql = sql.into();
104        self.connection
105            .call(move |connection| connection.execute_batch(&sql))
106            .await?;
107        Ok(())
108    }
109
110    pub(crate) async fn execute_unlocked(
111        &self,
112        query: &CompiledQuery,
113    ) -> Result<ExecuteResult, SqliteError> {
114        let sql = query.sql.clone();
115        let parameters = sqlite_parameters(&query.parameters)?;
116        let rows_affected = self
117            .connection
118            .call(move |connection| {
119                connection.execute(&sql, rusqlite::params_from_iter(parameters))
120            })
121            .await?;
122        Ok(ExecuteResult {
123            rows_affected: rows_affected as u64,
124        })
125    }
126
127    pub(crate) async fn fetch_all_unlocked(
128        &self,
129        query: &CompiledQuery,
130    ) -> Result<QueryResult<SqliteRow>, SqliteError> {
131        let sql = query.sql.clone();
132        let parameters = sqlite_parameters(&query.parameters)?;
133        let rows = self
134            .connection
135            .call(move |connection| {
136                let mut statement = connection.prepare(&sql)?;
137                let column_count = statement.column_count();
138                let mut cursor = statement.query(rusqlite::params_from_iter(parameters))?;
139                let mut rows = Vec::new();
140                while let Some(row) = cursor.next()? {
141                    let mut values = Vec::with_capacity(column_count);
142                    for index in 0..column_count {
143                        values.push(value_from_ref(row.get_ref(index)?)?);
144                    }
145                    rows.push(SqliteRow::new(values));
146                }
147                Ok(rows)
148            })
149            .await?;
150        Ok(QueryResult { rows })
151    }
152
153    pub(crate) async fn execute_control(&self, sql: impl Into<String>) -> Result<(), SqliteError> {
154        let sql = sql.into();
155        self.connection
156            .call(move |connection| connection.execute_batch(&sql))
157            .await?;
158        Ok(())
159    }
160}
161
162#[async_trait]
163impl Executor for SqliteExecutor {
164    type Row = SqliteRow;
165    type Error = SqliteError;
166
167    async fn execute(&self, query: &CompiledQuery) -> Result<ExecuteResult, Self::Error> {
168        let _guard = self.transaction_lock.lock().await;
169        self.execute_unlocked(query).await
170    }
171
172    async fn fetch_all(
173        &self,
174        query: &CompiledQuery,
175    ) -> Result<QueryResult<Self::Row>, Self::Error> {
176        let _guard = self.transaction_lock.lock().await;
177        self.fetch_all_unlocked(query).await
178    }
179}
180
181#[async_trait]
182impl TransactionManager for SqliteExecutor {
183    type Transaction = SqliteTransaction;
184
185    async fn begin(&self) -> Result<Self::Transaction, Self::Error> {
186        let guard = self.transaction_lock.clone().lock_owned().await;
187        self.execute_control("BEGIN IMMEDIATE").await?;
188        Ok(SqliteTransaction::new(self.clone(), guard))
189    }
190}
191
192fn sqlite_parameters(values: &[Value]) -> Result<Vec<SqliteValue>, SqliteError> {
193    values.iter().map(value_to_sqlite).collect()
194}
195
196fn value_to_sqlite(value: &Value) -> Result<SqliteValue, SqliteError> {
197    Ok(match value {
198        Value::Null => SqliteValue::Null,
199        Value::Bool(value) => SqliteValue::Integer(i64::from(*value)),
200        Value::I64(value) => SqliteValue::Integer(*value),
201        Value::U64(value) => SqliteValue::Integer(
202            i64::try_from(*value).map_err(|_| SqliteError::UnsignedOverflow(*value))?,
203        ),
204        Value::F64(value) => SqliteValue::Real(*value),
205        Value::String(value) => SqliteValue::Text(value.clone()),
206        Value::Bytes(value) => SqliteValue::Blob(value.clone()),
207        Value::Array(_) => return Err(SqliteError::UnsupportedParameter("array")),
208        #[cfg(feature = "uuid")]
209        Value::Uuid(value) => SqliteValue::Text(value.to_string()),
210        #[cfg(feature = "json")]
211        Value::Json(value) => SqliteValue::Text(value.to_string()),
212        #[cfg(feature = "chrono")]
213        Value::Date(value) => SqliteValue::Text(value.to_string()),
214        #[cfg(feature = "chrono")]
215        Value::Time(value) => SqliteValue::Text(value.to_string()),
216        #[cfg(feature = "chrono")]
217        Value::DateTime(value) => SqliteValue::Text(value.to_string()),
218        #[cfg(feature = "chrono")]
219        Value::DateTimeUtc(value) => SqliteValue::Text(value.to_rfc3339()),
220        #[cfg(feature = "decimal")]
221        Value::Decimal(value) => SqliteValue::Text(value.to_string()),
222    })
223}
224
225fn value_from_ref(value: ValueRef<'_>) -> rusqlite::Result<Value> {
226    Ok(match value {
227        ValueRef::Null => Value::Null,
228        ValueRef::Integer(value) => Value::I64(value),
229        ValueRef::Real(value) => Value::F64(value),
230        ValueRef::Text(value) => Value::String(String::from_utf8_lossy(value).into_owned()),
231        ValueRef::Blob(value) => Value::Bytes(value.to_vec()),
232    })
233}