walletdb.rs 17 KB

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