walletdb.rs 16 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508
  1. use std::{fs::create_dir_all, path::Path, str::FromStr, time::Duration};
  2. use async_std::sync::Arc;
  3. use incrementalmerkletree::bridgetree::BridgeTree;
  4. use log::{debug, error, info, LevelFilter};
  5. use rand::rngs::OsRng;
  6. use sqlx::{
  7. sqlite::{SqliteConnectOptions, SqliteJournalMode},
  8. ConnectOptions, Row, SqlitePool,
  9. };
  10. use super::WalletError;
  11. use crate::{
  12. crypto::{
  13. coin::Coin,
  14. keypair::{Keypair, PublicKey, SecretKey},
  15. merkle_node::MerkleNode,
  16. note::Note,
  17. nullifier::Nullifier,
  18. types::DrkTokenId,
  19. OwnCoin, OwnCoins,
  20. },
  21. util::serial::serialize,
  22. Result,
  23. };
  24. use super::wallet_api::WalletApi;
  25. pub type WalletPtr = Arc<WalletDb>;
  26. #[derive(Clone, Debug)]
  27. pub struct Balance {
  28. pub token_id: DrkTokenId,
  29. pub value: u64,
  30. pub nullifier: Nullifier,
  31. }
  32. #[derive(Clone, Debug)]
  33. pub struct Balances {
  34. pub list: Vec<Balance>,
  35. }
  36. pub struct WalletDb {
  37. pub conn: SqlitePool,
  38. }
  39. impl WalletApi for WalletDb {}
  40. impl WalletDb {
  41. pub async fn new(path: &str, password: &str) -> Result<WalletPtr> {
  42. if password.trim().is_empty() {
  43. error!("Password is empty. You must set a password to use the wallet.");
  44. return Err(WalletError::EmptyPassword.into())
  45. }
  46. if path != "sqlite::memory:" {
  47. let p = Path::new(path.strip_prefix("sqlite://").unwrap());
  48. if let Some(dirname) = p.parent() {
  49. info!("Creating path to database: {}", dirname.display());
  50. create_dir_all(&dirname)?;
  51. }
  52. }
  53. let mut connect_opts = SqliteConnectOptions::from_str(path)?
  54. .pragma("key", password.to_string())
  55. .create_if_missing(true)
  56. .journal_mode(SqliteJournalMode::Off);
  57. connect_opts.log_statements(LevelFilter::Trace);
  58. connect_opts.log_slow_statements(LevelFilter::Trace, Duration::from_micros(10));
  59. let conn = SqlitePool::connect_with(connect_opts).await?;
  60. info!("Opened connection at path {}", path);
  61. Ok(Arc::new(WalletDb { conn }))
  62. }
  63. pub async fn init_db(&self) -> Result<()> {
  64. info!("Initializing wallet database");
  65. let tree = include_str!("../../script/sql/tree.sql");
  66. let keys = include_str!("../../script/sql/keys.sql");
  67. let coins = include_str!("../../script/sql/coins.sql");
  68. let mut conn = self.conn.acquire().await?;
  69. debug!("Initializing merkle tree table");
  70. sqlx::query(tree).execute(&mut conn).await?;
  71. debug!("Initializing keys table");
  72. sqlx::query(keys).execute(&mut conn).await?;
  73. debug!("Initializing coins table");
  74. sqlx::query(coins).execute(&mut conn).await?;
  75. Ok(())
  76. }
  77. pub async fn keygen(&self) -> Result<Keypair> {
  78. debug!("Attempting to generate keypairs");
  79. let keypair = Keypair::random(&mut OsRng);
  80. self.put_keypair(&keypair).await?;
  81. Ok(keypair)
  82. }
  83. pub async fn put_keypair(&self, keypair: &Keypair) -> Result<()> {
  84. debug!("Writing keypair into the wallet database");
  85. let pubkey = serialize(&keypair.public);
  86. let secret = serialize(&keypair.secret);
  87. let is_default = 0;
  88. let mut conn = self.conn.acquire().await?;
  89. sqlx::query("INSERT INTO keys(public, secret, is_default) VALUES (?1, ?2, ?3)")
  90. .bind(pubkey)
  91. .bind(secret)
  92. .bind(is_default)
  93. .execute(&mut conn)
  94. .await?;
  95. Ok(())
  96. }
  97. pub async fn set_default_keypair(&self, public: &PublicKey) -> Result<()> {
  98. debug!("Set default keypair");
  99. let mut conn = self.conn.acquire().await?;
  100. let pubkey = serialize(public);
  101. // unset previous default keypair
  102. sqlx::query("UPDATE keys SET is_default = 0;").execute(&mut conn).await?;
  103. // set new default keypair
  104. sqlx::query("UPDATE keys SET is_default = 1 WHERE public = ?1;")
  105. .bind(pubkey)
  106. .execute(&mut conn)
  107. .await?;
  108. Ok(())
  109. }
  110. pub async fn get_default_keypair(&self) -> Result<Keypair> {
  111. debug!("Returning default keypair");
  112. let mut conn = self.conn.acquire().await?;
  113. let is_default = 1;
  114. let row = sqlx::query("SELECT * FROM keys WHERE is_default = ?1;")
  115. .bind(is_default)
  116. .fetch_one(&mut conn)
  117. .await?;
  118. let public: PublicKey = self.get_value_deserialized(row.get("public"))?;
  119. let secret: SecretKey = self.get_value_deserialized(row.get("secret"))?;
  120. Ok(Keypair { secret, public })
  121. }
  122. pub async fn get_keypairs(&self) -> Result<Vec<Keypair>> {
  123. debug!("Returning keypairs");
  124. let mut conn = self.conn.acquire().await?;
  125. let mut keypairs = vec![];
  126. for row in sqlx::query("SELECT * FROM keys").fetch_all(&mut conn).await? {
  127. let public: PublicKey = self.get_value_deserialized(row.get("public"))?;
  128. let secret: SecretKey = self.get_value_deserialized(row.get("secret"))?;
  129. keypairs.push(Keypair { public, secret });
  130. }
  131. Ok(keypairs)
  132. }
  133. pub async fn tree_gen(&self) -> Result<BridgeTree<MerkleNode, 32>> {
  134. debug!("Attempting to generate merkle tree");
  135. let mut conn = self.conn.acquire().await?;
  136. match sqlx::query("SELECT * FROM tree").fetch_one(&mut conn).await {
  137. Ok(_) => {
  138. error!("Merkle tree already exists");
  139. Err(WalletError::TreeExists.into())
  140. }
  141. Err(_) => {
  142. let tree = BridgeTree::<MerkleNode, 32>::new(100);
  143. self.put_tree(&tree).await?;
  144. Ok(tree)
  145. }
  146. }
  147. }
  148. pub async fn get_tree(&self) -> Result<BridgeTree<MerkleNode, 32>> {
  149. debug!("Getting merkle tree");
  150. let mut conn = self.conn.acquire().await?;
  151. let row = sqlx::query("SELECT * FROM tree").fetch_one(&mut conn).await?;
  152. let tree: BridgeTree<MerkleNode, 32> = bincode::deserialize(row.get("tree"))?;
  153. Ok(tree)
  154. }
  155. pub async fn put_tree(&self, tree: &BridgeTree<MerkleNode, 32>) -> Result<()> {
  156. debug!("Attempting to write merkle tree");
  157. let mut conn = self.conn.acquire().await?;
  158. let tree_bytes = bincode::serialize(tree)?;
  159. debug!("Deleting old row");
  160. sqlx::query("DELETE FROM tree;").execute(&mut conn).await?;
  161. debug!("Inserting new tree");
  162. sqlx::query("INSERT INTO tree (tree) VALUES (?1);")
  163. .bind(tree_bytes)
  164. .execute(&mut conn)
  165. .await?;
  166. Ok(())
  167. }
  168. pub async fn get_own_coins(&self) -> Result<OwnCoins> {
  169. debug!("Finding own coins");
  170. let is_spent = 0;
  171. let mut conn = self.conn.acquire().await?;
  172. let rows = sqlx::query("SELECT * FROM coins WHERE is_spent = ?1;")
  173. .bind(is_spent)
  174. .fetch_all(&mut conn)
  175. .await?;
  176. let mut own_coins = vec![];
  177. for row in rows {
  178. let coin = self.get_value_deserialized(row.get("coin"))?;
  179. // Note
  180. let serial = self.get_value_deserialized(row.get("serial"))?;
  181. let coin_blind = self.get_value_deserialized(row.get("coin_blind"))?;
  182. let value_blind = self.get_value_deserialized(row.get("valcom_blind"))?;
  183. // TODO: FIXME:
  184. let value_bytes: Vec<u8> = row.get("value");
  185. let value = u64::from_le_bytes(value_bytes.try_into().unwrap());
  186. let token_id = self.get_value_deserialized(row.get("token_id"))?;
  187. let note = Note { serial, value, token_id, coin_blind, value_blind };
  188. let secret = self.get_value_deserialized(row.get("secret"))?;
  189. let nullifier = self.get_value_deserialized(row.get("nullifier"))?;
  190. let leaf_position = self.get_value_deserialized(row.get("leaf_position"))?;
  191. let oc = OwnCoin { coin, note, secret, nullifier, leaf_position };
  192. own_coins.push(oc);
  193. }
  194. Ok(own_coins)
  195. }
  196. pub async fn put_own_coin(&self, own_coin: OwnCoin) -> Result<()> {
  197. debug!("Putting own coin into wallet database");
  198. let coin = self.get_value_serialized(&own_coin.coin.to_bytes())?;
  199. let serial = self.get_value_serialized(&own_coin.note.serial)?;
  200. let coin_blind = self.get_value_serialized(&own_coin.note.coin_blind)?;
  201. let value_blind = self.get_value_serialized(&own_coin.note.value_blind)?;
  202. let value = own_coin.note.value.to_le_bytes();
  203. let token_id = self.get_value_serialized(&own_coin.note.token_id)?;
  204. let secret = self.get_value_serialized(&own_coin.secret)?;
  205. let is_spent = 0;
  206. let nullifier = self.get_value_serialized(&own_coin.nullifier)?;
  207. let leaf_position = self.get_value_serialized(&own_coin.leaf_position)?;
  208. let mut conn = self.conn.acquire().await?;
  209. sqlx::query(
  210. "INSERT OR REPLACE INTO coins
  211. (coin, serial, value, token_id, coin_blind,
  212. valcom_blind, secret, is_spent, nullifier,
  213. leaf_position)
  214. VALUES
  215. (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10);",
  216. )
  217. .bind(coin)
  218. .bind(serial)
  219. .bind(value.to_vec())
  220. .bind(token_id)
  221. .bind(coin_blind)
  222. .bind(value_blind)
  223. .bind(secret)
  224. .bind(is_spent)
  225. .bind(nullifier)
  226. .bind(leaf_position)
  227. .execute(&mut conn)
  228. .await?;
  229. Ok(())
  230. }
  231. pub async fn remove_own_coins(&self) -> Result<()> {
  232. debug!("Removing own coins from wallet database");
  233. let mut conn = self.conn.acquire().await?;
  234. sqlx::query("DROP TABLE coins;").execute(&mut conn).await?;
  235. Ok(())
  236. }
  237. pub async fn confirm_spend_coin(&self, coin: &Coin) -> Result<()> {
  238. debug!("Confirm spend coin");
  239. let is_spent = 1;
  240. let coin = self.get_value_serialized(coin)?;
  241. let mut conn = self.conn.acquire().await?;
  242. sqlx::query("UPDATE coins SET is_spent = ?1 WHERE coin = ?2;")
  243. .bind(is_spent)
  244. .bind(coin)
  245. .execute(&mut conn)
  246. .await?;
  247. Ok(())
  248. }
  249. pub async fn get_balances(&self) -> Result<Balances> {
  250. debug!("Getting tokens and balances");
  251. let is_spent = 0;
  252. let mut conn = self.conn.acquire().await?;
  253. let rows = sqlx::query("SELECT value, token_id, nullifier FROM coins WHERE is_spent = ?1;")
  254. .bind(is_spent)
  255. .fetch_all(&mut conn)
  256. .await?;
  257. debug!("Found {} rows", rows.len());
  258. let mut list = vec![];
  259. for row in rows {
  260. // TODO: FIXME:
  261. let value_bytes: Vec<u8> = row.get("value");
  262. let value = u64::from_le_bytes(value_bytes.try_into().unwrap());
  263. let token_id = self.get_value_deserialized(row.get("token_id"))?;
  264. let nullifier = self.get_value_deserialized(row.get("nullifier"))?;
  265. list.push(Balance { token_id, value, nullifier });
  266. }
  267. Ok(Balances { list })
  268. }
  269. pub async fn get_token_id(&self) -> Result<Vec<DrkTokenId>> {
  270. debug!("Getting token ID");
  271. let is_spent = 0;
  272. let mut conn = self.conn.acquire().await?;
  273. let rows = sqlx::query("SELECT token_id FROM coins WHERE is_spent = ?1;")
  274. .bind(is_spent)
  275. .fetch_all(&mut conn)
  276. .await?;
  277. let mut token_ids = vec![];
  278. for row in rows {
  279. let token_id = self.get_value_deserialized(row.get("token_id"))?;
  280. token_ids.push(token_id);
  281. }
  282. Ok(token_ids)
  283. }
  284. pub async fn token_id_exists(&self, token_id: DrkTokenId) -> Result<bool> {
  285. debug!("Checking if token ID exists");
  286. let is_spent = 0;
  287. let id = self.get_value_serialized(&token_id)?;
  288. let mut conn = self.conn.acquire().await?;
  289. let id_check = sqlx::query("SELECT * FROM coins WHERE token_id = ?1 AND is_spent = ?2;")
  290. .bind(id)
  291. .bind(is_spent)
  292. .fetch_optional(&mut conn)
  293. .await?;
  294. Ok(id_check.is_some())
  295. }
  296. pub async fn test_wallet(&self) -> Result<()> {
  297. debug!("Testing wallet");
  298. let mut conn = self.conn.acquire().await?;
  299. let _row = sqlx::query("SELECT * FROM keys").fetch_one(&mut conn).await?;
  300. Ok(())
  301. }
  302. }
  303. #[cfg(test)]
  304. mod tests {
  305. use super::*;
  306. use crate::crypto::{
  307. merkle_node::MerkleNode,
  308. types::{DrkCoinBlind, DrkSerial, DrkValueBlind},
  309. };
  310. use group::ff::Field;
  311. use incrementalmerkletree::{Frontier, Tree};
  312. use pasta_curves::pallas;
  313. use rand::rngs::OsRng;
  314. const WPASS: &str = "darkfi";
  315. fn dummy_coin(s: &SecretKey, v: u64, t: &DrkTokenId) -> OwnCoin {
  316. let serial = DrkSerial::random(&mut OsRng);
  317. let note = Note {
  318. serial,
  319. value: v,
  320. token_id: *t,
  321. coin_blind: DrkCoinBlind::random(&mut OsRng),
  322. value_blind: DrkValueBlind::random(&mut OsRng),
  323. };
  324. let coin = Coin(pallas::Base::random(&mut OsRng));
  325. let nullifier = Nullifier::new(*s, serial);
  326. let leaf_position = 0;
  327. OwnCoin { coin, note, secret: *s, nullifier, leaf_position }
  328. }
  329. #[async_std::test]
  330. async fn test_walletdb() -> Result<()> {
  331. let wallet = WalletDb::new("sqlite::memory:", WPASS).await?;
  332. let keypair = Keypair::random(&mut OsRng);
  333. // init_db()
  334. wallet.init_db().await?;
  335. // tree_gen()
  336. let mut tree1 = wallet.tree_gen().await?;
  337. // put_keypair()
  338. wallet.put_keypair(&keypair).await?;
  339. let token_id = DrkTokenId::random(&mut OsRng);
  340. let c0 = dummy_coin(&keypair.secret, 69, &token_id);
  341. let c1 = dummy_coin(&keypair.secret, 420, &token_id);
  342. let c2 = dummy_coin(&keypair.secret, 42, &token_id);
  343. let c3 = dummy_coin(&keypair.secret, 11, &token_id);
  344. // put_own_coin()
  345. wallet.put_own_coin(c0).await?;
  346. tree1.append(&MerkleNode::from_coin(&c0.coin));
  347. tree1.witness();
  348. wallet.put_own_coin(c1).await?;
  349. tree1.append(&MerkleNode::from_coin(&c1.coin));
  350. tree1.witness();
  351. wallet.put_own_coin(c2).await?;
  352. tree1.append(&MerkleNode::from_coin(&c2.coin));
  353. tree1.witness();
  354. wallet.put_own_coin(c3).await?;
  355. tree1.append(&MerkleNode::from_coin(&c3.coin));
  356. tree1.witness();
  357. // We'll check this merkle root corresponds to the one we'll retrieve.
  358. let root1 = tree1.root();
  359. // put_tree()
  360. wallet.put_tree(&tree1).await?;
  361. // get_token_id()
  362. let id = wallet.get_token_id().await?;
  363. assert_eq!(id.len(), 4);
  364. for i in id {
  365. assert_eq!(i, token_id);
  366. assert!(wallet.token_id_exists(i).await?);
  367. }
  368. // get_balances()
  369. let balances = wallet.get_balances().await?;
  370. assert_eq!(balances.list.len(), 4);
  371. assert_eq!(balances.list[1].value, 420);
  372. assert_eq!(balances.list[2].value, 42);
  373. assert_eq!(balances.list[3].token_id, token_id);
  374. /////////////////
  375. //// keypair ////
  376. /////////////////
  377. let keypair2 = Keypair::random(&mut OsRng);
  378. // add new keypair
  379. wallet.put_keypair(&keypair2).await?;
  380. // get all keypairs
  381. let keypairs = wallet.get_keypairs().await?;
  382. assert_eq!(keypair, keypairs[0]);
  383. assert_eq!(keypair2, keypairs[1]);
  384. // set the keypair at index 1 as the default keypair
  385. wallet.set_default_keypair(&keypair2.public).await?;
  386. // get default keypair
  387. assert_eq!(keypair2, wallet.get_default_keypair().await?);
  388. // get_own_coins()
  389. let own_coins = wallet.get_own_coins().await?;
  390. assert_eq!(own_coins.len(), 4);
  391. assert_eq!(own_coins[0], c0);
  392. assert_eq!(own_coins[1], c1);
  393. assert_eq!(own_coins[2], c2);
  394. assert_eq!(own_coins[3], c3);
  395. // get_tree()
  396. let tree2 = wallet.get_tree().await?;
  397. let root2 = tree2.root();
  398. assert_eq!(root1, root2);
  399. // Let's try it once more to test sql replacing.
  400. wallet.put_tree(&tree2).await?;
  401. let tree3 = wallet.get_tree().await?;
  402. let root3 = tree3.root();
  403. assert_eq!(root2, root3);
  404. Ok(())
  405. }
  406. }