Sfoglia il codice sorgente

wallet: Use a global SQL connection rather than connecting for every method.

parazyd 4 anni fa
parent
commit
1b3f638e14
3 ha cambiato i file con 88 aggiunte e 177 eliminazioni
  1. 41 112
      src/wallet/cashierdb.rs
  2. 0 5
      src/wallet/wallet_api.rs
  3. 47 60
      src/wallet/walletdb.rs

+ 41 - 112
src/wallet/cashierdb.rs

@@ -1,7 +1,7 @@
-use std::path::{Path, PathBuf};
+use std::path::Path;
 
 use async_std::sync::{Arc, Mutex};
-use log::debug;
+use log::{debug, error, info};
 use rusqlite::{named_params, params, Connection};
 
 use super::{Keypair, WalletApi};
@@ -12,8 +12,7 @@ use crate::{types::*, Error, Result};
 pub type CashierDbPtr = Arc<CashierDb>;
 
 pub struct CashierDb {
-    pub path: PathBuf,
-    pub password: String,
+    pub conn: Connection,
     pub initialized: Mutex<bool>,
 }
 
@@ -37,59 +36,43 @@ pub struct DepositToken {
     pub mint_address: String,
 }
 
-impl WalletApi for CashierDb {
-    fn get_password(&self) -> String {
-        self.password.to_owned()
-    }
-    fn get_path(&self) -> PathBuf {
-        self.path.to_owned()
-    }
-}
+impl WalletApi for CashierDb {}
 
 impl CashierDb {
     pub fn new(path: &Path, password: String) -> Result<CashierDbPtr> {
         debug!(target: "CASHIERDB", "new() Constructor called");
+        if password.trim().is_empty() {
+            error!(target: "CASHIERDB", "Password is empty. You must set a password to use the wallet.");
+            return Err(Error::from(ClientFailed::EmptyPassword));
+        }
+
+        let conn = Connection::open(path)?;
+        conn.pragma_update(None, "key", &password)?;
+        info!(target: "CASHIERDB", "Opened connection at path: {:?}", path);
+
         Ok(Arc::new(Self {
-            path: path.to_owned(),
-            password,
+            conn,
             initialized: Mutex::new(false),
         }))
     }
 
     pub async fn init_db(&self) -> Result<()> {
         if !*self.initialized.lock().await {
-            if !self.password.trim().is_empty() {
-                let contents = include_str!("../../sql/cashier.sql");
-                let conn = Connection::open(&self.path)?;
-                debug!(target: "CASHIERDB", "Opened connection at path {:?}", self.path);
-                conn.pragma_update(None, "key", &self.password)?;
-                conn.execute_batch(contents)?;
-                *self.initialized.lock().await = true;
-            } else {
-                debug!(
-                    target: "CASHIERDB",
-                    "Password is empty. You must set a password to use the wallet."
-                );
-                return Err(Error::from(ClientFailed::EmptyPassword));
-            }
-        } else {
-            debug!(target: "WALLETDB", "Wallet already initialized.");
-            return Err(Error::from(ClientFailed::WalletInitialized));
+            let contents = include_str!("../../sql/cashier.sql");
+            self.conn.execute_batch(contents)?;
+            *self.initialized.lock().await = true;
+            return Ok(());
         }
-        Ok(())
+
+        error!(target: "WALLETDB", "Wallet already initialized.");
+        Err(Error::from(ClientFailed::WalletInitialized))
     }
 
     pub fn put_main_keys(&self, token_key: &TokenKey, network: &NetworkName) -> Result<()> {
         debug!(target: "CASHIERDB", "Put main keys");
-
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
         let network = self.get_value_serialized(network)?;
 
-        conn.execute(
+        self.conn.execute(
             "INSERT INTO main_keypairs
             (token_key_private, token_key_public, network)
             VALUES
@@ -100,23 +83,20 @@ impl CashierDb {
                 ":network": &network,
             },
         )?;
+
         Ok(())
     }
 
     pub fn get_main_keys(&self, network: &NetworkName) -> Result<Vec<TokenKey>> {
         debug!(target: "CASHIERDB", "Get main keys");
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
         let network = self.get_value_serialized(network)?;
 
-        let mut stmt = conn.prepare(
+        let mut stmt = self.conn.prepare(
             "SELECT token_key_private, token_key_public
             FROM main_keypairs
             WHERE network = :network ;",
         )?;
+
         let keys_iter = stmt
             .query_map::<(Vec<u8>, Vec<u8>), _, _>(&[(":network", &network)], |row| {
                 Ok((row.get(0)?, row.get(1)?))
@@ -137,14 +117,8 @@ impl CashierDb {
 
     pub fn remove_withdraw_and_deposit_keys(&self) -> Result<()> {
         debug!(target: "CASHIERDB", "Remove withdraw and deposit keys");
-
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
-        conn.execute("DROP TABLE deposit_keypairs;", [])?;
-        conn.execute("DROP TABLE withdraw_keypairs;", [])?;
+        self.conn.execute("DROP TABLE deposit_keypairs;", [])?;
+        self.conn.execute("DROP TABLE withdraw_keypairs;", [])?;
         Ok(())
     }
 
@@ -166,12 +140,7 @@ impl CashierDb {
         let confirm = self.get_value_serialized(&false)?;
         let mint_address = self.get_value_serialized(&mint_address)?;
 
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
-        conn.execute(
+        self.conn.execute(
             "INSERT INTO withdraw_keypairs
             (token_key_public, d_key_private, d_key_public, network,  token_id, mint_address, confirm)
             VALUES
@@ -186,6 +155,7 @@ impl CashierDb {
                 ":confirm": confirm,
             },
         )?;
+
         Ok(())
     }
 
@@ -200,19 +170,13 @@ impl CashierDb {
     ) -> Result<()> {
         debug!(target: "CASHIERDB", "Put exchange keys");
 
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
         let d_key_public = self.get_value_serialized(d_key_public)?;
         let token_id = self.get_value_serialized(token_id)?;
         let network = self.get_value_serialized(network)?;
         let confirm = self.get_value_serialized(&false)?;
-
         let mint_address = self.get_value_serialized(&mint_address)?;
 
-        conn.execute(
+        self.conn.execute(
             "INSERT INTO deposit_keypairs
             (d_key_public, token_key_private, token_key_public, network, token_id, mint_address, confirm)
             VALUES
@@ -227,19 +191,15 @@ impl CashierDb {
                 ":confirm": &confirm,
             },
         )?;
+
         Ok(())
     }
 
     pub fn get_withdraw_private_keys(&self) -> Result<Vec<DrkSecretKey>> {
         debug!(target: "CASHIERDB", "Get withdraw private keys");
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
         let confirm = self.get_value_serialized(&false)?;
 
-        let mut stmt = conn.prepare(
+        let mut stmt = self.conn.prepare(
             "SELECT d_key_private
                 FROM withdraw_keypairs
                 WHERE confirm = :confirm",
@@ -262,20 +222,15 @@ impl CashierDb {
         pub_key: &DrkPublicKey,
     ) -> Result<Option<WithdrawToken>> {
         debug!(target: "CASHIERDB", "Get token address by pub_key");
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
         let d_key_public = self.get_value_serialized(pub_key)?;
-
         let confirm = self.get_value_serialized(&false)?;
 
-        let mut stmt = conn.prepare(
+        let mut stmt = self.conn.prepare(
             "SELECT token_key_public, network, token_id, mint_address
             FROM withdraw_keypairs
             WHERE d_key_public = :d_key_public AND confirm = :confirm;",
         )?;
+
         let addr_iter = stmt.query_map(
             &[(":d_key_public", &d_key_public), (":confirm", &confirm)],
             |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?)),
@@ -307,21 +262,17 @@ impl CashierDb {
     ) -> Result<Vec<TokenKey>> {
         debug!(target: "CASHIERDB", "Check for existing dkey");
         let d_key_public = self.get_value_serialized(d_key_public)?;
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
         let network = self.get_value_serialized(network)?;
         let confirm = self.get_value_serialized(&false)?;
 
-        let mut stmt = conn.prepare(
+        let mut stmt = self.conn.prepare(
             "SELECT token_key_private, token_key_public
             FROM deposit_keypairs
             WHERE d_key_public = :d_key_public
             AND network = :network
             AND confirm = :confirm ;",
         )?;
+
         let keys_iter = stmt.query_map::<(Vec<u8>, Vec<u8>), _, _>(
             &[
                 (":d_key_public", &d_key_public),
@@ -349,20 +300,16 @@ impl CashierDb {
         network: &NetworkName,
     ) -> Result<Vec<DepositToken>> {
         debug!(target: "CASHIERDB", "Check for existing dkey");
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
         let network = self.get_value_serialized(network)?;
         let confirm = self.get_value_serialized(&false)?;
 
-        let mut stmt = conn.prepare(
+        let mut stmt = self.conn.prepare(
             "SELECT d_key_public, token_key_private, token_key_public, token_id, mint_address
             FROM deposit_keypairs
             WHERE network = :network
             AND confirm = :confirm ;",
         )?;
+
         let keys_iter =
             stmt.query_map(&[(":network", &network), (":confirm", &confirm)], |row| {
                 Ok((
@@ -403,16 +350,10 @@ impl CashierDb {
         network: &NetworkName,
     ) -> Result<Option<Keypair>> {
         debug!(target: "CASHIERDB", "Check for existing token address");
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
         let confirm = self.get_value_serialized(&false)?;
-
         let network = self.get_value_serialized(network)?;
 
-        let mut stmt = conn.prepare(
+        let mut stmt = self.conn.prepare(
             "SELECT d_key_private, d_key_public FROM withdraw_keypairs
                 WHERE token_key_public = :token_key_public
                 AND network = :network
@@ -447,16 +388,10 @@ impl CashierDb {
         network: &NetworkName,
     ) -> Result<()> {
         debug!(target: "CASHIERDB", "Confirm withdraw keys");
-
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
         let network = self.get_value_serialized(network)?;
         let confirm = self.get_value_serialized(&true)?;
 
-        conn.execute(
+        self.conn.execute(
             "UPDATE withdraw_keypairs
             SET confirm = ?1
             WHERE token_key_public = ?2
@@ -473,17 +408,11 @@ impl CashierDb {
         network: &NetworkName,
     ) -> Result<()> {
         debug!(target: "CASHIERDB", "Confirm withdraw keys");
-
-        // open connection
-        let conn = Connection::open(&self.path)?;
-        // unlock database
-        conn.pragma_update(None, "key", &self.password)?;
-
         let network = self.get_value_serialized(network)?;
         let confirm = self.get_value_serialized(&true)?;
         let d_key_public = self.get_value_serialized(d_key_public)?;
 
-        conn.execute(
+        self.conn.execute(
             "UPDATE deposit_keypairs
             SET confirm = ?1
             WHERE d_key_public = ?2

+ 0 - 5
src/wallet/wallet_api.rs

@@ -1,12 +1,7 @@
-use std::path::PathBuf;
-
 use crate::serial::{deserialize, serialize, Decodable, Encodable};
 use crate::Result;
 
 pub trait WalletApi {
-    fn get_password(&self) -> String;
-    fn get_path(&self) -> PathBuf;
-
     fn get_value_serialized<T: Encodable>(&self, data: &T) -> Result<Vec<u8>> {
         let v = serialize(data);
         Ok(v)

+ 47 - 60
src/wallet/walletdb.rs

@@ -1,8 +1,8 @@
 use std::collections::HashMap;
-use std::path::{Path, PathBuf};
+use std::path::Path;
 use std::sync::Arc;
 
-use log::{debug, error};
+use log::{debug, error, info};
 use pasta_curves::arithmetic::Field;
 use rand::rngs::OsRng;
 use rusqlite::{named_params, params, Connection};
@@ -45,68 +45,55 @@ impl Balances {
 }
 
 pub struct WalletDb {
-    pub path: PathBuf,
-    pub password: String,
+    pub conn: Connection,
 }
 
-impl WalletApi for WalletDb {
-    fn get_password(&self) -> String {
-        self.password.to_owned()
-    }
-    fn get_path(&self) -> PathBuf {
-        self.path.to_owned()
-    }
-}
+impl WalletApi for WalletDb {}
 
 impl WalletDb {
     pub fn new(path: &Path, password: String) -> Result<WalletPtr> {
         debug!(target: "WALLETDB", "new() Constructor called");
-        Ok(Arc::new(Self {
-            path: path.to_owned(),
-            password,
-        }))
-    }
-
-    fn connect(&self) -> Result<Connection> {
-        if self.password.trim().is_empty() {
+        if password.trim().is_empty() {
             error!(target: "WALLETDB", "Password is empty. You must set a password to use the wallet.");
             return Err(Error::from(ClientFailed::EmptyPassword));
         }
-        let conn = Connection::open(&self.path)?;
-        debug!(target: "WALLETDB", "OPENED CONNECTION AT PATH {:?}", self.path);
-        conn.pragma_update(None, "key", &self.password)?;
-        Ok(conn)
+
+        let conn = Connection::open(path)?;
+        conn.pragma_update(None, "key", &password)?;
+        info!(target: "WALLETDB", "Opened connection at path: {:?}", path);
+
+        Ok(Arc::new(Self { conn }))
     }
 
     pub fn init_db(&self) -> Result<()> {
         debug!(target: "WALLETDB", "Initialize...");
         let contents = include_str!("../../sql/schema.sql");
-        let conn = self.connect()?;
-        Ok(conn.execute_batch(contents)?)
+        Ok(self.conn.execute_batch(contents)?)
     }
 
     pub fn key_gen(&self) -> Result<()> {
         debug!(target: "WALLETDB", "Attempting to generate keys...");
-        let conn = self.connect()?;
-        let mut stmt = conn.prepare("SELECT * FROM keys WHERE key_id > ?")?;
+        let mut stmt = self.conn.prepare("SELECT * FROM keys WHERE key_id > ?")?;
+
         let key_check = stmt.exists(params!["0"])?;
+
         if !key_check {
             let secret = DrkSecretKey::random(&mut OsRng);
             let public = derive_publickey(secret);
             self.put_keypair(&public, &secret)?;
-        } else {
-            debug!(target: "WALLETDB", "Keys already exist.");
-            return Err(Error::from(ClientFailed::KeyExists));
+            return Ok(());
         }
-        Ok(())
+
+        error!(target: "WALLETDB", "Keys already exist.");
+        Err(Error::from(ClientFailed::KeyExists))
     }
 
     pub fn put_keypair(&self, key_public: &DrkPublicKey, key_private: &DrkSecretKey) -> Result<()> {
-        let conn = self.connect()?;
+        debug!(target: "WALLETDB", "put_keypair()");
         let key_public = serial::serialize(key_public);
         let key_private = serial::serialize(key_private);
 
-        conn.execute(
+        self.conn.execute(
             "INSERT INTO keys(key_public, key_private) VALUES (?1, ?2)",
             params![key_public, key_private],
         )?;
@@ -116,8 +103,8 @@ impl WalletDb {
 
     pub fn get_keypairs(&self) -> Result<Vec<Keypair>> {
         debug!(target: "WALLETDB", "Returning keypairs...");
-        let conn = self.connect()?;
-        let mut stmt = conn.prepare("SELECT * FROM keys")?;
+        let mut stmt = self.conn.prepare("SELECT * FROM keys")?;
+
         // this just gets the first key. maybe we should randomize this
         let key_iter = stmt.query_map([], |row| Ok((row.get(1)?, row.get(2)?)))?;
         let mut keypairs = Vec::new();
@@ -136,9 +123,12 @@ impl WalletDb {
 
     pub fn get_own_coins(&self) -> Result<OwnCoins> {
         debug!(target: "WALLETDB", "Get own coins");
-        let conn = self.connect()?;
         let is_spent = 0;
-        let mut coins = conn.prepare("SELECT * FROM coins WHERE is_spent = :is_spent ;")?;
+
+        let mut coins = self
+            .conn
+            .prepare("SELECT * FROM coins WHERE is_spent = :is_spent ;")?;
+
         let rows = coins.query_map(&[(":is_spent", &is_spent)], |row| {
             Ok((
                 row.get(0)?,
@@ -197,10 +187,7 @@ impl WalletDb {
 
     pub fn put_own_coins(&self, own_coin: OwnCoin) -> Result<()> {
         debug!(target: "WALLETDB", "Put own coins");
-        let conn = self.connect()?;
-
         let coin = self.get_value_serialized(&own_coin.coin.to_bytes())?;
-
         let serial = self.get_value_serialized(&own_coin.note.serial)?;
         let coin_blind = self.get_value_serialized(&own_coin.note.coin_blind)?;
         let value_blind = self.get_value_serialized(&own_coin.note.value_blind)?;
@@ -211,7 +198,7 @@ impl WalletDb {
         let is_spent = 0;
         let nullifier = self.get_value_serialized(&own_coin.nullifier)?;
 
-        conn.execute(
+        self.conn.execute(
             "INSERT OR REPLACE INTO coins
             (coin, serial, value, token_id, coin_blind,
             valcom_blind, witness, secret, is_spent, nullifier)
@@ -231,22 +218,22 @@ impl WalletDb {
                 ":nullifier": nullifier,
             },
         )?;
+
         Ok(())
     }
 
     pub fn remove_own_coins(&self) -> Result<()> {
         debug!(target: "WALLETDB", "Remove own coins");
-        let conn = self.connect()?;
-        conn.execute("DROP TABLE coins;", [])?;
+        let _rows = self.conn.execute("DROP TABLE coins;", [])?;
         Ok(())
     }
 
     pub fn confirm_spend_coin(&self, coin: &Coin) -> Result<()> {
         debug!(target: "WALLETDB", "Confirm spend coin");
-        let coin = self.get_value_serialized(coin)?;
-        let conn = self.connect()?;
         let is_spent = 1;
-        conn.execute(
+        let coin = self.get_value_serialized(coin)?;
+
+        self.conn.execute(
             "UPDATE coins
             SET is_spent = ?1
             WHERE coin = ?2 ;",
@@ -306,13 +293,12 @@ impl WalletDb {
 
     pub fn get_balances(&self) -> Result<Balances> {
         debug!(target: "WALLETDB", "Get token and balances...");
-        let conn = self.connect()?;
-
         let is_spent = 0;
 
-        let mut stmt = conn.prepare(
+        let mut stmt = self.conn.prepare(
             "SELECT value, token_id, nullifier FROM coins  WHERE is_spent = :is_spent ;",
         )?;
+
         let rows = stmt.query_map(&[(":is_spent", &is_spent)], |row| {
             Ok((row.get(0)?, row.get(1)?, row.get(2)?))
         })?;
@@ -336,11 +322,12 @@ impl WalletDb {
 
     pub fn get_token_id(&self) -> Result<Vec<DrkTokenId>> {
         debug!(target: "WALLETDB", "Get token ID...");
-        let conn = self.connect()?;
-
         let is_spent = 0;
 
-        let mut stmt = conn.prepare("SELECT token_id FROM coins WHERE is_spent = :is_spent ;")?;
+        let mut stmt = self
+            .conn
+            .prepare("SELECT token_id FROM coins WHERE is_spent = :is_spent ;")?;
+
         let rows = stmt.query_map(&[(":is_spent", &is_spent)], |row| row.get(0))?;
 
         let mut token_ids = Vec::new();
@@ -356,20 +343,20 @@ impl WalletDb {
 
     pub fn token_id_exists(&self, token_id: &DrkTokenId) -> Result<bool> {
         debug!(target: "WALLETDB", "Check tokenID exists");
-        let conn = self.connect()?;
-
-        let id = self.get_value_serialized(token_id)?;
         let is_spent = 0;
+        let id = self.get_value_serialized(token_id)?;
+
+        let mut stmt = self
+            .conn
+            .prepare("SELECT * FROM coins WHERE token_id = ? AND is_spent = ? ;")?;
 
-        let mut stmt = conn.prepare("SELECT * FROM coins WHERE token_id = ? AND is_spent = ? ;")?;
         let id_check = stmt.exists(params![id, is_spent])?;
+
         Ok(id_check)
     }
 
     pub fn test_wallet(&self) -> Result<()> {
-        let conn = Connection::open(&self.path)?;
-        conn.pragma_update(None, "key", &self.password)?;
-        let mut stmt = conn.prepare("SELECT * FROM keys")?;
+        let mut stmt = self.conn.prepare("SELECT * FROM keys")?;
         let _rows = stmt.query([])?;
         Ok(())
     }