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 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 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}