walletdb.rs 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389
  1. /* This file is part of DarkFi (https://dark.fi)
  2. *
  3. * Copyright (C) 2020-2024 Dyne.org foundation
  4. *
  5. * This program is free software: you can redistribute it and/or modify
  6. * it under the terms of the GNU Affero General Public License as
  7. * published by the Free Software Foundation, either version 3 of the
  8. * License, or (at your option) any later version.
  9. *
  10. * This program is distributed in the hope that it will be useful,
  11. * but WITHOUT ANY WARRANTY; without even the implied warranty of
  12. * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
  13. * GNU Affero General Public License for more details.
  14. *
  15. * You should have received a copy of the GNU Affero General Public License
  16. * along with this program. If not, see <https://www.gnu.org/licenses/>.
  17. */
  18. use std::{path::PathBuf, sync::Arc};
  19. use log::{debug, error};
  20. use rusqlite::{
  21. types::{ToSql, Value},
  22. Connection,
  23. };
  24. use smol::lock::Mutex;
  25. use crate::error::{WalletDbError, WalletDbResult};
  26. pub type WalletPtr = Arc<WalletDb>;
  27. /// Structure representing base wallet database operations.
  28. pub struct WalletDb {
  29. /// Connection to the SQLite database
  30. pub conn: Mutex<Connection>,
  31. }
  32. impl WalletDb {
  33. /// Create a new wallet database handler. If `path` is `None`, create it in memory.
  34. pub fn new(path: Option<PathBuf>, password: Option<&str>) -> WalletDbResult<WalletPtr> {
  35. let Ok(conn) = (match path.clone() {
  36. Some(p) => Connection::open(p),
  37. None => Connection::open_in_memory(),
  38. }) else {
  39. return Err(WalletDbError::ConnectionFailed);
  40. };
  41. if let Some(password) = password {
  42. if let Err(e) = conn.pragma_update(None, "key", password) {
  43. error!(target: "walletdb::new", "[WalletDb] Pragma update failed: {e}");
  44. return Err(WalletDbError::PragmaUpdateError);
  45. };
  46. }
  47. if let Err(e) = conn.pragma_update(None, "foreign_keys", "ON") {
  48. error!(target: "walletdb::new", "[WalletDb] Pragma update failed: {e}");
  49. return Err(WalletDbError::PragmaUpdateError);
  50. };
  51. debug!(target: "walletdb::new", "[WalletDb] Opened Sqlite connection at \"{path:?}\"");
  52. Ok(Arc::new(Self { conn: Mutex::new(conn) }))
  53. }
  54. /// This function executes a given SQL query that contains multiple SQL statements,
  55. /// that don't contain any parameters.
  56. pub async fn exec_batch_sql(&self, query: &str) -> WalletDbResult<()> {
  57. debug!(target: "walletdb::exec_batch_sql", "[WalletDb] Executing batch SQL query:\n{query}");
  58. if let Err(e) = self.conn.lock().await.execute_batch(query) {
  59. error!(target: "walletdb::exec_batch_sql", "[WalletDb] Query failed: {e}");
  60. return Err(WalletDbError::QueryExecutionFailed)
  61. };
  62. Ok(())
  63. }
  64. /// This function executes a given SQL query, but isn't able to return anything.
  65. /// Therefore it's best to use it for initializing a table or similar things.
  66. pub async fn exec_sql(&self, query: &str, params: &[&dyn ToSql]) -> WalletDbResult<()> {
  67. debug!(target: "walletdb::exec_sql", "[WalletDb] Executing SQL query:\n{query}");
  68. // If no params are provided, execute directly
  69. if params.is_empty() {
  70. if let Err(e) = self.conn.lock().await.execute(query, ()) {
  71. error!(target: "walletdb::exec_sql", "[WalletDb] Query failed: {e}");
  72. return Err(WalletDbError::QueryExecutionFailed)
  73. };
  74. return Ok(())
  75. }
  76. // First we prepare the query
  77. let conn = self.conn.lock().await;
  78. let Ok(mut stmt) = conn.prepare(query) else {
  79. return Err(WalletDbError::QueryPreparationFailed)
  80. };
  81. // Execute the query using provided params
  82. if let Err(e) = stmt.execute(params) {
  83. error!(target: "walletdb::exec_sql", "[WalletDb] Query failed: {e}");
  84. return Err(WalletDbError::QueryExecutionFailed)
  85. };
  86. // Finalize query and drop connection lock
  87. if let Err(e) = stmt.finalize() {
  88. error!(target: "walletdb::exec_sql", "[WalletDb] Query finalization failed: {e}");
  89. return Err(WalletDbError::QueryFinalizationFailed)
  90. };
  91. drop(conn);
  92. Ok(())
  93. }
  94. /// Generate a `SELECT` query for provided table from selected column names and
  95. /// provided `WHERE` clauses. Named parameters are supported in the `WHERE` clauses,
  96. /// assuming they follow the normal formatting ":{column_name}".
  97. fn generate_select_query(
  98. &self,
  99. table: &str,
  100. col_names: &[&str],
  101. params: &[(&str, &dyn ToSql)],
  102. ) -> String {
  103. let mut query = if col_names.is_empty() {
  104. format!("SELECT * FROM {}", table)
  105. } else {
  106. format!("SELECT {} FROM {}", col_names.join(", "), table)
  107. };
  108. if params.is_empty() {
  109. return query
  110. }
  111. let mut where_str = Vec::with_capacity(params.len());
  112. for (k, _) in params {
  113. let col = &k[1..];
  114. where_str.push(format!("{col} = {k}"));
  115. }
  116. query.push_str(&format!(" WHERE {}", where_str.join(" AND ")));
  117. query
  118. }
  119. /// Query provided table from selected column names and provided `WHERE` clauses,
  120. /// for a single row.
  121. pub async fn query_single(
  122. &self,
  123. table: &str,
  124. col_names: &[&str],
  125. params: &[(&str, &dyn ToSql)],
  126. ) -> WalletDbResult<Vec<Value>> {
  127. // Generate `SELECT` query
  128. let query = self.generate_select_query(table, col_names, params);
  129. debug!(target: "walletdb::query_single", "[WalletDb] Executing SQL query:\n{query}");
  130. // First we prepare the query
  131. let conn = self.conn.lock().await;
  132. let Ok(mut stmt) = conn.prepare(&query) else {
  133. return Err(WalletDbError::QueryPreparationFailed)
  134. };
  135. // Execute the query using provided params
  136. let Ok(mut rows) = stmt.query(params) else {
  137. return Err(WalletDbError::QueryExecutionFailed)
  138. };
  139. // Check if row exists
  140. let Ok(next) = rows.next() else { return Err(WalletDbError::QueryExecutionFailed) };
  141. let row = match next {
  142. Some(row_result) => row_result,
  143. None => return Err(WalletDbError::RowNotFound),
  144. };
  145. // Grab returned values
  146. let mut result = vec![];
  147. if col_names.is_empty() {
  148. let mut idx = 0;
  149. loop {
  150. let Ok(value) = row.get(idx) else { break };
  151. result.push(value);
  152. idx += 1;
  153. }
  154. } else {
  155. for col in col_names {
  156. let Ok(value) = row.get(*col) else {
  157. return Err(WalletDbError::ParseColumnValueError)
  158. };
  159. result.push(value);
  160. }
  161. }
  162. Ok(result)
  163. }
  164. /// Query provided table from selected column names and provided `WHERE` clauses,
  165. /// for multiple rows.
  166. pub async fn query_multiple(
  167. &self,
  168. table: &str,
  169. col_names: &[&str],
  170. params: &[(&str, &dyn ToSql)],
  171. ) -> WalletDbResult<Vec<Vec<Value>>> {
  172. // Generate `SELECT` query
  173. let query = self.generate_select_query(table, col_names, params);
  174. debug!(target: "walletdb::multiple", "[WalletDb] Executing SQL query:\n{query}");
  175. // First we prepare the query
  176. let conn = self.conn.lock().await;
  177. let Ok(mut stmt) = conn.prepare(&query) else {
  178. return Err(WalletDbError::QueryPreparationFailed)
  179. };
  180. // Execute the query using provided converted params
  181. let Ok(mut rows) = stmt.query(params) else {
  182. if let Err(e) = stmt.query(params) {
  183. println!("eeer: {e:?}");
  184. }
  185. return Err(WalletDbError::QueryExecutionFailed)
  186. };
  187. // Loop over returned rows and parse them
  188. let mut result = vec![];
  189. loop {
  190. // Check if an error occured
  191. let row = match rows.next() {
  192. Ok(r) => r,
  193. Err(_) => return Err(WalletDbError::QueryExecutionFailed),
  194. };
  195. // Check if no row was returned
  196. let row = match row {
  197. Some(r) => r,
  198. None => break,
  199. };
  200. // Grab row returned values
  201. let mut row_values = vec![];
  202. if col_names.is_empty() {
  203. let mut idx = 0;
  204. loop {
  205. let Ok(value) = row.get(idx) else { break };
  206. row_values.push(value);
  207. idx += 1;
  208. }
  209. } else {
  210. for col in col_names {
  211. let Ok(value) = row.get(*col) else {
  212. return Err(WalletDbError::ParseColumnValueError)
  213. };
  214. row_values.push(value);
  215. }
  216. }
  217. result.push(row_values);
  218. }
  219. Ok(result)
  220. }
  221. }
  222. /// Custom implementation of rusqlite::named_params! to use `expr` instead of `literal` as `$param_name`,
  223. /// and append the ":" named parameters prefix.
  224. #[macro_export]
  225. macro_rules! convert_named_params {
  226. () => {
  227. &[] as &[(&str, &dyn rusqlite::types::ToSql)]
  228. };
  229. ($(($param_name:expr, $param_val:expr)),+ $(,)?) => {
  230. &[$((format!(":{}", $param_name).as_str(), &$param_val as &dyn rusqlite::types::ToSql)),+] as &[(&str, &dyn rusqlite::types::ToSql)]
  231. };
  232. }
  233. #[cfg(test)]
  234. mod tests {
  235. use rusqlite::types::Value;
  236. use crate::walletdb::WalletDb;
  237. #[test]
  238. fn test_mem_wallet() {
  239. smol::block_on(async {
  240. let wallet = WalletDb::new(None, Some("foobar")).unwrap();
  241. wallet.exec_sql("CREATE TABLE mista ( numba INTEGER );", &[]).await.unwrap();
  242. wallet.exec_sql("INSERT INTO mista ( numba ) VALUES ( 42 );", &[]).await.unwrap();
  243. let ret = wallet.query_single("mista", &["numba"], &[]).await.unwrap();
  244. assert_eq!(ret.len(), 1);
  245. let numba: i64 = if let Value::Integer(numba) = ret[0] { numba } else { -1 };
  246. assert_eq!(numba, 42);
  247. });
  248. }
  249. #[test]
  250. fn test_query_single() {
  251. smol::block_on(async {
  252. let wallet = WalletDb::new(None, None).unwrap();
  253. wallet
  254. .exec_sql(
  255. "CREATE TABLE mista ( why INTEGER, are TEXT, you INTEGER, gae BLOB );",
  256. &[],
  257. )
  258. .await
  259. .unwrap();
  260. let why = 42;
  261. let are = "are".to_string();
  262. let you = 69;
  263. let gae = vec![42u8; 32];
  264. wallet
  265. .exec_sql(
  266. "INSERT INTO mista ( why, are, you, gae ) VALUES (?1, ?2, ?3, ?4);",
  267. rusqlite::params![why, are, you, gae],
  268. )
  269. .await
  270. .unwrap();
  271. let ret =
  272. wallet.query_single("mista", &["why", "are", "you", "gae"], &[]).await.unwrap();
  273. assert_eq!(ret.len(), 4);
  274. assert_eq!(ret[0], Value::Integer(why));
  275. assert_eq!(ret[1], Value::Text(are.clone()));
  276. assert_eq!(ret[2], Value::Integer(you));
  277. assert_eq!(ret[3], Value::Blob(gae.clone()));
  278. let ret = wallet
  279. .query_single(
  280. "mista",
  281. &["gae"],
  282. rusqlite::named_params! {":why": why, ":are": are, ":you": you},
  283. )
  284. .await
  285. .unwrap();
  286. assert_eq!(ret.len(), 1);
  287. assert_eq!(ret[0], Value::Blob(gae));
  288. });
  289. }
  290. #[test]
  291. fn test_query_multi() {
  292. smol::block_on(async {
  293. let wallet = WalletDb::new(None, None).unwrap();
  294. wallet
  295. .exec_sql(
  296. "CREATE TABLE mista ( why INTEGER, are TEXT, you INTEGER, gae BLOB );",
  297. &[],
  298. )
  299. .await
  300. .unwrap();
  301. let why = 42;
  302. let are = "are".to_string();
  303. let you = 69;
  304. let gae = vec![42u8; 32];
  305. wallet
  306. .exec_sql(
  307. "INSERT INTO mista ( why, are, you, gae ) VALUES (?1, ?2, ?3, ?4);",
  308. rusqlite::params![why, are, you, gae],
  309. )
  310. .await
  311. .unwrap();
  312. wallet
  313. .exec_sql(
  314. "INSERT INTO mista ( why, are, you, gae ) VALUES (?1, ?2, ?3, ?4);",
  315. rusqlite::params![why, are, you, gae],
  316. )
  317. .await
  318. .unwrap();
  319. let ret = wallet.query_multiple("mista", &[], &[]).await.unwrap();
  320. assert_eq!(ret.len(), 2);
  321. for row in ret {
  322. assert_eq!(row.len(), 4);
  323. assert_eq!(row[0], Value::Integer(why));
  324. assert_eq!(row[1], Value::Text(are.clone()));
  325. assert_eq!(row[2], Value::Integer(you));
  326. assert_eq!(row[3], Value::Blob(gae.clone()));
  327. }
  328. let ret = wallet
  329. .query_multiple(
  330. "mista",
  331. &["gae"],
  332. convert_named_params! {("why", why), ("are", are), ("you", you)},
  333. )
  334. .await
  335. .unwrap();
  336. assert_eq!(ret.len(), 2);
  337. for row in ret {
  338. assert_eq!(row.len(), 1);
  339. assert_eq!(row[0], Value::Blob(gae.clone()));
  340. }
  341. });
  342. }
  343. }