1use 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#[derive(Debug, Clone, Deserialize, Serialize)]
18pub struct Config {
19 pub database_url: String,
21 #[serde(default = "default_max_connections")]
23 pub max_connections: u32,
24 #[serde(default)]
26 pub min_connections: u32,
27 #[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
40pub struct SqliteDatabase {
42 pool: Pool<Sqlite>,
43}
44
45impl SqliteDatabase {
46 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
232pub 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}