Skip to main content

rskit_database_sqlite/
store.rs

1//! `SQLite` backend implementing [`rskit_database::DatabaseClient`].
2
3use std::sync::Arc;
4use std::time::Duration;
5
6use rskit_database::{
7    DatabaseClient, DatabaseConfig, DatabaseFactory, DatabaseQuery, DatabaseRegistry,
8    DatabaseResult, DatabaseTransaction,
9};
10use rskit_errors::{AppError, AppResult, ErrorCode};
11use serde::{Deserialize, Serialize};
12use sqlx::sqlite::{SqliteConnectOptions, SqlitePoolOptions};
13use sqlx::{AssertSqlSafe, Executor, Pool, Sqlite};
14use tokio::sync::Mutex;
15
16/// Configuration for the `SQLite` database backend.
17#[derive(Debug, Clone, Deserialize, Serialize)]
18pub struct Config {
19    /// `SQLite` database URL, such as `sqlite://app.db` or `sqlite::memory:`.
20    pub database_url: String,
21    /// Maximum pooled connections.
22    #[serde(default = "default_max_connections")]
23    pub max_connections: u32,
24    /// Minimum pooled connections.
25    #[serde(default)]
26    pub min_connections: u32,
27    /// Connection acquisition timeout.
28    #[serde(default = "default_connect_timeout")]
29    pub connect_timeout: Duration,
30}
31
32const fn default_max_connections() -> u32 {
33    10
34}
35
36const fn default_connect_timeout() -> Duration {
37    Duration::from_secs(30)
38}
39
40/// `SQLite` database client.
41pub struct SqliteDatabase {
42    pool: Pool<Sqlite>,
43}
44
45impl SqliteDatabase {
46    /// Connect to `SQLite` using the provided config.
47    pub async fn connect(config: Config) -> AppResult<Self> {
48        validate_config(&config)?;
49        let options = config
50            .database_url
51            .parse::<SqliteConnectOptions>()
52            .map_err(|error| {
53                AppError::new(ErrorCode::InvalidInput, "invalid `SQLite` database URL")
54                    .with_cause(error)
55            })?
56            .create_if_missing(true);
57        let pool = SqlitePoolOptions::new()
58            .max_connections(config.max_connections)
59            .min_connections(config.min_connections)
60            .acquire_timeout(config.connect_timeout)
61            .connect_with(options)
62            .await
63            .map_err(database_error("connect `SQLite` database"))?;
64        Ok(Self { pool })
65    }
66}
67
68#[async_trait::async_trait]
69impl DatabaseClient for SqliteDatabase {
70    async fn execute(&self, query: DatabaseQuery) -> AppResult<DatabaseResult> {
71        execute_on(&self.pool, query).await
72    }
73
74    async fn begin(&self) -> AppResult<Box<dyn DatabaseTransaction>> {
75        let tx = self
76            .pool
77            .begin()
78            .await
79            .map_err(database_error("begin `SQLite` transaction"))?;
80        Ok(Box::new(SqliteTransaction { tx: Mutex::new(tx) }))
81    }
82
83    async fn ping(&self) -> AppResult<()> {
84        self.pool
85            .acquire()
86            .await
87            .map_err(database_error("ping `SQLite` database"))?;
88        Ok(())
89    }
90}
91
92struct SqliteTransaction {
93    tx: Mutex<sqlx::Transaction<'static, Sqlite>>,
94}
95
96#[async_trait::async_trait]
97impl DatabaseTransaction for SqliteTransaction {
98    #[allow(clippy::significant_drop_tightening)]
99    async fn execute(&self, query: DatabaseQuery) -> AppResult<DatabaseResult> {
100        let mut guard = self.tx.lock().await;
101        execute_on(&mut **guard, query).await
102    }
103
104    async fn commit(self: Box<Self>) -> AppResult<()> {
105        self.tx
106            .into_inner()
107            .commit()
108            .await
109            .map_err(database_error("commit `SQLite` transaction"))
110    }
111
112    async fn rollback(self: Box<Self>) -> AppResult<()> {
113        self.tx
114            .into_inner()
115            .rollback()
116            .await
117            .map_err(database_error("rollback `SQLite` transaction"))
118    }
119}
120
121async fn execute_on<'e, E>(executor: E, query: DatabaseQuery) -> AppResult<DatabaseResult>
122where
123    E: Executor<'e, Database = Sqlite>,
124{
125    if query.statement.trim().is_empty() {
126        return Err(AppError::new(
127            ErrorCode::InvalidInput,
128            "database query statement is required",
129        ));
130    }
131    let mut sql = sqlx::query(AssertSqlSafe(query.statement.as_str()));
132    for parameter in query.parameters {
133        sql = bind_json_value(sql, parameter)?;
134    }
135    let result = sql
136        .execute(executor)
137        .await
138        .map_err(database_error("execute `SQLite` statement"))?;
139    Ok(DatabaseResult {
140        rows_affected: result.rows_affected(),
141    })
142}
143
144fn bind_json_value(
145    query: sqlx::query::Query<'_, Sqlite, sqlx::sqlite::SqliteArguments>,
146    value: serde_json::Value,
147) -> AppResult<sqlx::query::Query<'_, Sqlite, sqlx::sqlite::SqliteArguments>> {
148    match value {
149        serde_json::Value::Null => Ok(query.bind(Option::<String>::None)),
150        serde_json::Value::Bool(value) => Ok(query.bind(value)),
151        serde_json::Value::Number(value) => {
152            if let Some(value) = value.as_i64() {
153                Ok(query.bind(value))
154            } else if let Some(value) = value.as_u64() {
155                let value = i64::try_from(value).map_err(|_| {
156                    AppError::new(
157                        ErrorCode::InvalidInput,
158                        "`SQLite` integer parameter exceeds i64::MAX",
159                    )
160                })?;
161                Ok(query.bind(value))
162            } else if let Some(value) = value.as_f64() {
163                Ok(query.bind(value))
164            } else {
165                Err(AppError::new(
166                    ErrorCode::InvalidInput,
167                    "`SQLite` numeric parameter is not representable",
168                ))
169            }
170        }
171        serde_json::Value::String(value) => Ok(query.bind(value)),
172        value @ (serde_json::Value::Array(_) | serde_json::Value::Object(_)) => {
173            let text = serde_json::to_string(&value).map_err(|error| {
174                AppError::new(
175                    ErrorCode::InvalidInput,
176                    "`SQLite` structured parameter is not serializable to JSON text",
177                )
178                .with_cause(error)
179            })?;
180            Ok(query.bind(text))
181        }
182    }
183}
184
185fn validate_config(config: &Config) -> AppResult<()> {
186    if config.database_url.trim().is_empty() {
187        return Err(AppError::new(
188            ErrorCode::MissingField,
189            "`SQLite` database_url is required",
190        ));
191    }
192    if config.max_connections == 0 {
193        return Err(AppError::new(
194            ErrorCode::InvalidInput,
195            "`SQLite` max_connections must be greater than zero",
196        ));
197    }
198    if config.min_connections > config.max_connections {
199        return Err(AppError::new(
200            ErrorCode::InvalidInput,
201            "`SQLite` min_connections must not exceed max_connections",
202        ));
203    }
204    if config.connect_timeout.is_zero() {
205        return Err(AppError::new(
206            ErrorCode::InvalidInput,
207            "`SQLite` connect_timeout must be greater than zero",
208        ));
209    }
210    Ok(())
211}
212
213fn database_error(operation: &'static str) -> impl FnOnce(sqlx::Error) -> AppError {
214    move |error| {
215        AppError::new(ErrorCode::DatabaseError, format!("{operation} failed")).with_cause(error)
216    }
217}
218
219struct SqliteFactory {
220    config: Config,
221}
222
223#[async_trait::async_trait]
224impl DatabaseFactory for SqliteFactory {
225    async fn create(&self, _config: &DatabaseConfig) -> AppResult<Arc<dyn DatabaseClient>> {
226        Ok(Arc::new(
227            SqliteDatabase::connect(self.config.clone()).await?,
228        ))
229    }
230}
231
232/// Explicitly register the `SQLite` database backend.
233pub fn register(registry: &mut DatabaseRegistry, config: Config) -> AppResult<()> {
234    registry.register("sqlite", Arc::new(SqliteFactory { config }))
235}
236
237#[cfg(test)]
238mod tests {
239    use super::*;
240    use rskit_database::{DatabaseClient, DatabaseConfig, DatabaseQuery, DatabaseRegistry};
241
242    fn config() -> Config {
243        Config {
244            database_url: "sqlite::memory:".into(),
245            max_connections: 1,
246            min_connections: 0,
247            connect_timeout: Duration::from_secs(5),
248        }
249    }
250
251    #[tokio::test]
252    async fn execute_binds_parameters_without_sql_concatenation() {
253        let db = SqliteDatabase::connect(config()).await.unwrap();
254        db.execute(DatabaseQuery::new("CREATE TABLE users (name TEXT)"))
255            .await
256            .unwrap();
257        db.execute(
258            DatabaseQuery::new("INSERT INTO users (name) VALUES (?)")
259                .with_parameter("Robert'); DROP TABLE users;--"),
260        )
261        .await
262        .unwrap();
263        db.execute(
264            DatabaseQuery::new("INSERT INTO users (name) VALUES (?)").with_parameter("Alice"),
265        )
266        .await
267        .unwrap();
268        assert_eq!(
269            db.execute(
270                DatabaseQuery::new("UPDATE users SET name = ? WHERE name = ?")
271                    .with_parameter("Bob")
272                    .with_parameter("Alice")
273            )
274            .await
275            .unwrap()
276            .rows_affected,
277            1
278        );
279        db.execute(
280            DatabaseQuery::new("INSERT INTO users (name) VALUES (?)")
281                .with_parameter(serde_json::Value::Null),
282        )
283        .await
284        .unwrap();
285    }
286
287    #[tokio::test]
288    async fn transaction_commit_and_rollback_are_explicit() {
289        let db = SqliteDatabase::connect(config()).await.unwrap();
290        db.execute(DatabaseQuery::new("CREATE TABLE items (name TEXT)"))
291            .await
292            .unwrap();
293        let tx = db.begin().await.unwrap();
294        tx.execute(DatabaseQuery::new("INSERT INTO items (name) VALUES (?)").with_parameter("one"))
295            .await
296            .unwrap();
297        tx.commit().await.unwrap();
298        let tx = db.begin().await.unwrap();
299        tx.execute(DatabaseQuery::new("INSERT INTO items (name) VALUES (?)").with_parameter("two"))
300            .await
301            .unwrap();
302        tx.rollback().await.unwrap();
303    }
304
305    #[tokio::test]
306    async fn structured_parameters_bind_as_json_text() {
307        let db = SqliteDatabase::connect(config()).await.unwrap();
308        db.execute(DatabaseQuery::new("CREATE TABLE docs (payload TEXT)"))
309            .await
310            .unwrap();
311        db.execute(
312            DatabaseQuery::new("INSERT INTO docs (payload) VALUES (?)")
313                .with_parameter(serde_json::json!({"key": [1, 2, 3]})),
314        )
315        .await
316        .unwrap();
317    }
318
319    #[test]
320    fn config_validation_rejects_invalid_values() {
321        assert_eq!(
322            validate_config(&Config {
323                database_url: String::new(),
324                ..config()
325            })
326            .unwrap_err()
327            .code(),
328            ErrorCode::MissingField
329        );
330        assert_eq!(
331            validate_config(&Config {
332                max_connections: 0,
333                ..config()
334            })
335            .unwrap_err()
336            .code(),
337            ErrorCode::InvalidInput
338        );
339        assert_eq!(
340            validate_config(&Config {
341                min_connections: 2,
342                ..config()
343            })
344            .unwrap_err()
345            .code(),
346            ErrorCode::InvalidInput
347        );
348    }
349
350    #[test]
351    fn config_defaults_apply_when_fields_are_omitted() {
352        let config: Config =
353            serde_json::from_value(serde_json::json!({ "database_url": "sqlite::memory:" }))
354                .unwrap();
355        assert_eq!(config.max_connections, default_max_connections());
356        assert_eq!(config.min_connections, 0);
357        assert_eq!(config.connect_timeout, default_connect_timeout());
358    }
359
360    #[test]
361    fn config_validation_rejects_zero_connect_timeout() {
362        assert_eq!(
363            validate_config(&Config {
364                connect_timeout: Duration::ZERO,
365                ..config()
366            })
367            .unwrap_err()
368            .code(),
369            ErrorCode::InvalidInput
370        );
371    }
372
373    #[tokio::test]
374    async fn connect_rejects_invalid_database_url() {
375        let result = SqliteDatabase::connect(Config {
376            database_url: "sqlite://foo?mode=bogus".into(),
377            ..config()
378        })
379        .await;
380        assert_eq!(
381            result.err().map(|error| error.code()),
382            Some(ErrorCode::InvalidInput)
383        );
384    }
385
386    #[tokio::test]
387    async fn execute_rejects_empty_statement() {
388        let db = SqliteDatabase::connect(config()).await.unwrap();
389        assert_eq!(
390            db.execute(DatabaseQuery::new("   "))
391                .await
392                .unwrap_err()
393                .code(),
394            ErrorCode::InvalidInput
395        );
396    }
397
398    #[tokio::test]
399    async fn execute_maps_sqlx_errors_to_database_error() {
400        let db = SqliteDatabase::connect(config()).await.unwrap();
401        assert_eq!(
402            db.execute(DatabaseQuery::new("SELECT * FROM missing_table"))
403                .await
404                .unwrap_err()
405                .code(),
406            ErrorCode::DatabaseError
407        );
408    }
409
410    #[tokio::test]
411    async fn ping_acquires_a_connection() {
412        let db = SqliteDatabase::connect(config()).await.unwrap();
413        db.ping().await.unwrap();
414    }
415
416    #[tokio::test]
417    async fn binds_every_json_scalar_variant() {
418        let db = SqliteDatabase::connect(config()).await.unwrap();
419        db.execute(DatabaseQuery::new(
420            "CREATE TABLE values_table (flag INTEGER, signed INTEGER, real REAL, text TEXT)",
421        ))
422        .await
423        .unwrap();
424        let result = db
425            .execute(
426                DatabaseQuery::new(
427                    "INSERT INTO values_table (flag, signed, real, text) VALUES (?, ?, ?, ?)",
428                )
429                .with_parameter(true)
430                .with_parameter(-7_i64)
431                .with_parameter(1.5_f64)
432                .with_parameter("hello"),
433            )
434            .await
435            .unwrap();
436        assert_eq!(result.rows_affected, 1);
437    }
438
439    #[tokio::test]
440    async fn unsigned_integer_above_i64_max_is_rejected() {
441        let db = SqliteDatabase::connect(config()).await.unwrap();
442        db.execute(DatabaseQuery::new("CREATE TABLE big (value INTEGER)"))
443            .await
444            .unwrap();
445        assert_eq!(
446            db.execute(
447                DatabaseQuery::new("INSERT INTO big (value) VALUES (?)")
448                    .with_parameter(serde_json::json!(u64::MAX)),
449            )
450            .await
451            .unwrap_err()
452            .code(),
453            ErrorCode::InvalidInput
454        );
455    }
456
457    #[tokio::test]
458    async fn transaction_execute_maps_errors() {
459        let db = SqliteDatabase::connect(config()).await.unwrap();
460        let tx = db.begin().await.unwrap();
461        assert_eq!(
462            tx.execute(DatabaseQuery::new("   "))
463                .await
464                .unwrap_err()
465                .code(),
466            ErrorCode::InvalidInput
467        );
468        tx.rollback().await.unwrap();
469    }
470
471    #[tokio::test]
472    async fn register_adds_backend_without_connecting() {
473        let mut registry = DatabaseRegistry::new();
474        register(&mut registry, config()).unwrap();
475        assert!(registry.contains("sqlite"));
476        let built = registry
477            .build(&DatabaseConfig {
478                backend: "sqlite".into(),
479                ..DatabaseConfig::default()
480            })
481            .await
482            .unwrap();
483        built.ping().await.unwrap();
484    }
485}