walletdb.rs 18 KB

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