diff --git a/consensus.go b/consensus.go index baae75db..186fcedc 100644 --- a/consensus.go +++ b/consensus.go @@ -7,7 +7,6 @@ import ( "decred.org/dcrwallet/v2/errors" w "decred.org/dcrwallet/v2/wallet" - "github.com/decred/dcrd/chaincfg/chainhash" "github.com/decred/dcrd/chaincfg/v3" "github.com/decred/dcrd/wire" diff --git a/internal/loader/loader.go b/internal/loader/loader.go index 349ca226..c1e5c787 100644 --- a/internal/loader/loader.go +++ b/internal/loader/loader.go @@ -13,11 +13,15 @@ import ( "decred.org/dcrwallet/v2/errors" "decred.org/dcrwallet/v2/wallet" - _ "decred.org/dcrwallet/v2/wallet/drivers/bdb" // driver loaded during init + + // driver loaded during init + _ "decred.org/dcrwallet/v2/wallet/drivers/bdb" "github.com/decred/dcrd/chaincfg/v3" "github.com/decred/dcrd/dcrutil/v4" "github.com/decred/dcrd/txscript/v4/stdaddr" - _ "github.com/planetdecred/dcrlibwallet/badgerdb" // initialize badger driver + + // initialize badger driver + _ "github.com/planetdecred/dcrlibwallet/badgerdb" ) const ( diff --git a/multiwallet.go b/multiwallet.go index c2af0433..7999e632 100644 --- a/multiwallet.go +++ b/multiwallet.go @@ -23,14 +23,19 @@ import ( "golang.org/x/crypto/bcrypt" ) +// MultiWallet allows for tracking/managing multiple Wallets. +// It is safe to use concurrently from multiple go-routines. type MultiWallet struct { dbDriver string rootDir string db *storm.DB + walletsMu sync.RWMutex + wallets map[int]*Wallet + badWalletsMu sync.RWMutex + badWallets map[int]*Wallet + chainParams *chaincfg.Params - wallets map[int]*Wallet - badWallets map[int]*Wallet syncData *syncData notificationListenersMu sync.RWMutex @@ -150,9 +155,14 @@ func (mw *MultiWallet) Shutdown() { mw.CancelRescan() mw.CancelSync() - for _, wallet := range mw.wallets { + // Locking mw.walletsMu to prevent other actors from accessing wallets during shutdown. + mw.walletsMu.Lock() + for key, wallet := range mw.wallets { wallet.Shutdown() + // Wallet is no longer operational after it's been shut down. + delete(mw.wallets, key) } + mw.walletsMu.Unlock() if mw.db != nil { if err := mw.db.Close(); err != nil { @@ -270,17 +280,24 @@ func (mw *MultiWallet) OpenWallets(startupPassphrase []byte) error { return err } + // Locking mw.walletsMu to prevent other actors from accessing wallets until all are + // opened (means ready to work with). + mw.walletsMu.Lock() for _, wallet := range mw.wallets { err = wallet.openWallet() if err != nil { return err } } + mw.walletsMu.Unlock() return nil } func (mw *MultiWallet) AllWalletsAreWatchOnly() (bool, error) { + mw.walletsMu.Lock() + defer mw.walletsMu.Unlock() + if len(mw.wallets) == 0 { return false, errors.New(ErrInvalid) } @@ -424,6 +441,43 @@ func (mw *MultiWallet) LinkExistingWallet(walletName, walletDataDir, originalPub }) } +// walletsReadCopy returns a copy of mw.wallets map, copying it in concurrently-safe +// manner. This allows iterating over map without holding mutex during iterations +// (useful when iteration takes long time). +func (mw *MultiWallet) walletsReadCopy() map[int]*Wallet { + mw.walletsMu.RLock() + defer mw.walletsMu.RUnlock() + + result := make(map[int]*Wallet, len(mw.wallets)) + for key, value := range mw.wallets { + result[key] = value + } + return result +} + +// walletsUpdate is concurrently-safe way to update mw.wallets map element. +func (mw *MultiWallet) walletsUpdate(key int, newWallet *Wallet) { + mw.walletsMu.Lock() + defer mw.walletsMu.Unlock() + + mw.wallets[key] = newWallet +} + +// walletsUpdate is concurrently-safe way to get an element from mw.wallets map. +func (mw *MultiWallet) walletsGet(key int) *Wallet { + mw.walletsMu.RLock() + defer mw.walletsMu.RUnlock() + + return mw.wallets[key] +} + +// walletsUpdate is concurrently-safe way to delete from mw.wallets map. +func (mw *MultiWallet) walletsDelete(key int) { + mw.walletsMu.Lock() + delete(mw.wallets, key) + mw.walletsMu.Unlock() +} + // saveNewWallet performs the following tasks using a db batch operation to ensure // that db changes are rolled back if any of the steps below return an error. // @@ -491,7 +545,7 @@ func (mw *MultiWallet) saveNewWallet(wallet *Wallet, setupWallet func() error) ( return nil, translateError(err) } - mw.wallets[wallet.ID] = wallet + mw.walletsUpdate(wallet.ID, wallet) return wallet, nil } @@ -542,17 +596,22 @@ func (mw *MultiWallet) DeleteWallet(walletID int, privPass []byte) error { return translateError(err) } - delete(mw.wallets, walletID) + mw.walletsDelete(walletID) return nil } func (mw *MultiWallet) BadWallets() map[int]*Wallet { + mw.badWalletsMu.RLock() + defer mw.badWalletsMu.RUnlock() + return mw.badWallets } func (mw *MultiWallet) DeleteBadWallet(walletID int) error { + mw.badWalletsMu.RLock() wallet := mw.badWallets[walletID] + mw.badWalletsMu.RUnlock() if wallet == nil { return errors.New(ErrNotExist) } @@ -565,16 +624,19 @@ func (mw *MultiWallet) DeleteBadWallet(walletID int) error { } os.RemoveAll(wallet.dataDir) + + mw.badWalletsMu.Lock() delete(mw.badWallets, walletID) + mw.badWalletsMu.Unlock() return nil } func (mw *MultiWallet) WalletWithID(walletID int) *Wallet { - if wallet, ok := mw.wallets[walletID]; ok { - return wallet - } - return nil + mw.walletsMu.RLock() + defer mw.walletsMu.RUnlock() + + return mw.wallets[walletID] } // VerifySeedForWallet compares seedMnemonic with the decrypted wallet.EncryptedSeed and clears wallet.EncryptedSeed if they match. @@ -600,7 +662,7 @@ func (mw *MultiWallet) VerifySeedForWallet(walletID int, seedMnemonic string, pr // NumWalletsNeedingSeedBackup returns the number of opened wallets whose seed haven't been verified. func (mw *MultiWallet) NumWalletsNeedingSeedBackup() int32 { var backupsNeeded int32 - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { if wallet.WalletOpened() && wallet.EncryptedSeed != nil { backupsNeeded++ } @@ -609,12 +671,15 @@ func (mw *MultiWallet) NumWalletsNeedingSeedBackup() int32 { } func (mw *MultiWallet) LoadedWalletsCount() int32 { + mw.walletsMu.RLock() + defer mw.walletsMu.RUnlock() + return int32(len(mw.wallets)) } func (mw *MultiWallet) OpenedWalletIDsRaw() []int { walletIDs := make([]int, 0) - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { if wallet.WalletOpened() { walletIDs = append(walletIDs, wallet.ID) } @@ -634,7 +699,7 @@ func (mw *MultiWallet) OpenedWalletsCount() int32 { func (mw *MultiWallet) SyncedWalletsCount() int32 { var syncedWallets int32 - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { if wallet.WalletOpened() && wallet.synced { syncedWallets++ } diff --git a/multiwallet_utils.go b/multiwallet_utils.go index 9208c2be..c2411952 100644 --- a/multiwallet_utils.go +++ b/multiwallet_utils.go @@ -134,7 +134,7 @@ func (mw *MultiWallet) WalletWithXPub(xpub string) (int, error) { ctx, cancel := mw.contextWithShutdownCancel() defer cancel() - for _, w := range mw.wallets { + for _, w := range mw.walletsReadCopy() { if !w.WalletOpened() { return -1, errors.Errorf("wallet %d is not open and cannot be checked", w.ID) } @@ -171,7 +171,7 @@ func (mw *MultiWallet) WalletWithSeed(seedMnemonic string) (int, error) { return -1, err } - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { if !wallet.WalletOpened() { return -1, errors.Errorf("cannot check if seed matches unloaded wallet %d", wallet.ID) } diff --git a/sync.go b/sync.go index 18cfcd09..6f236cf6 100644 --- a/sync.go +++ b/sync.go @@ -218,7 +218,7 @@ func (mw *MultiWallet) SpvSync() error { mw.initActiveSyncData() wallets := make(map[int]*w.Wallet) - for id, wallet := range mw.wallets { + for id, wallet := range mw.walletsReadCopy() { wallets[id] = wallet.Internal() wallet.waitingForHeaders = true wallet.syncing = true @@ -288,7 +288,7 @@ func (mw *MultiWallet) CancelSync() { log.Info("Canceling sync. May take a while for sync to fully cancel.") // Stop running cspp mixers - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { if wallet.IsAccountMixerActive() { log.Infof("[%d] Stopping cspp mixer", wallet.ID) err := mw.StopAccountMixer(wallet.ID) @@ -420,7 +420,7 @@ func (mw *MultiWallet) PeerInfo() (string, error) { func (mw *MultiWallet) GetBestBlock() *BlockInfo { var bestBlock int32 = -1 var blockInfo *BlockInfo - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { if !wallet.WalletOpened() { continue } @@ -438,7 +438,7 @@ func (mw *MultiWallet) GetBestBlock() *BlockInfo { func (mw *MultiWallet) GetLowestBlock() *BlockInfo { var lowestBlock int32 = -1 var blockInfo *BlockInfo - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { if !wallet.WalletOpened() { continue } @@ -483,7 +483,7 @@ func (wallet *Wallet) GetBestBlockTimeStamp() int64 { func (mw *MultiWallet) GetLowestBlockTimestamp() int64 { var timestamp int64 = -1 - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { bestBlockTimestamp := wallet.GetBestBlockTimeStamp() if bestBlockTimestamp < timestamp || timestamp == -1 { timestamp = bestBlockTimestamp diff --git a/syncnotification.go b/syncnotification.go index 4e444437..663b9ceb 100644 --- a/syncnotification.go +++ b/syncnotification.go @@ -166,9 +166,11 @@ func (mw *MultiWallet) fetchHeadersStarted(peerInitialHeight int32) { return } + mw.walletsMu.RLock() for _, wallet := range mw.wallets { wallet.waitingForHeaders = true } + mw.walletsMu.RUnlock() lowestBlockHeight := mw.GetLowestBlock().Height @@ -200,7 +202,7 @@ func (mw *MultiWallet) fetchHeadersProgress(lastFetchedHeaderHeight int32, lastF return } - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { if wallet.waitingForHeaders { wallet.waitingForHeaders = wallet.GetBestBlock() > lastFetchedHeaderHeight } @@ -484,7 +486,7 @@ func (mw *MultiWallet) rescanProgress(walletID int, rescannedThrough int32) { return } - wallet := mw.wallets[walletID] + wallet := mw.walletsGet(walletID) totalHeadersToScan := wallet.GetBestBlock() rescanRate := float64(rescannedThrough) / float64(totalHeadersToScan) @@ -621,7 +623,7 @@ func (mw *MultiWallet) resetSyncData() { mw.syncData.activeSyncData = nil mw.syncData.mu.Unlock() - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { wallet.waitingForHeaders = true wallet.LockWallet() // lock wallet if previously unlocked to perform account discovery. } @@ -633,7 +635,7 @@ func (mw *MultiWallet) synced(walletID int, synced bool) { // begin indexing transactions after sync is completed, // syncProgressListeners.OnSynced() will be invoked after transactions are indexed var txIndexing errgroup.Group - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { txIndexing.Go(wallet.IndexTransactions) } @@ -662,7 +664,7 @@ func (mw *MultiWallet) synced(walletID int, synced bool) { return } - wallet := mw.wallets[walletID] + wallet := mw.walletsGet(walletID) wallet.synced = synced wallet.syncing = false mw.listenForTransactions(wallet.ID) diff --git a/ticket.go b/ticket.go index e1b2de3a..3e741d98 100644 --- a/ticket.go +++ b/ticket.go @@ -32,7 +32,7 @@ func (wallet *Wallet) TotalStakingRewards() (int64, error) { func (mw *MultiWallet) TotalStakingRewards() (int64, error) { var totalRewards int64 - for _, wal := range mw.wallets { + for _, wal := range mw.walletsReadCopy() { walletTotalRewards, err := wal.TotalStakingRewards() if err != nil { return 0, err @@ -94,7 +94,7 @@ func (wallet *Wallet) StakingOverview() (stOverview *StakingOverview, err error) func (mw *MultiWallet) StakingOverview() (stOverview *StakingOverview, err error) { stOverview = &StakingOverview{} - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { st, err := wallet.StakingOverview() if err != nil { return nil, err @@ -134,7 +134,7 @@ func (wallet *Wallet) TicketPrice() (*TicketPriceResponse, error) { func (mw *MultiWallet) TicketPrice() (*TicketPriceResponse, error) { bestBlock := mw.GetBestBlock() - for _, wal := range mw.wallets { + for _, wal := range mw.walletsReadCopy() { resp, err := wal.TicketPrice() if err != nil { return nil, err diff --git a/transactions.go b/transactions.go index 9bfb7684..7d1d23ac 100644 --- a/transactions.go +++ b/transactions.go @@ -125,7 +125,7 @@ func (mw *MultiWallet) GetTransactions(offset, limit, txFilter int32, newestFirs func (mw *MultiWallet) GetTransactionsRaw(offset, limit, txFilter int32, newestFirst bool) ([]Transaction, error) { transactions := make([]Transaction, 0) - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { walletTransactions, err := wallet.GetTransactionsRaw(offset, limit, txFilter, newestFirst) if err != nil { return nil, err diff --git a/txandblocknotifications.go b/txandblocknotifications.go index dd55166a..207273b7 100644 --- a/txandblocknotifications.go +++ b/txandblocknotifications.go @@ -9,7 +9,7 @@ import ( func (mw *MultiWallet) listenForTransactions(walletID int) { go func() { - wallet := mw.wallets[walletID] + wallet := mw.walletsGet(walletID) n := wallet.Internal().NtfnServer.TransactionNotifications() for { @@ -112,7 +112,7 @@ func (mw *MultiWallet) RemoveTxAndBlockNotificationListener(uniqueIdentifier str } func (mw *MultiWallet) checkWalletMixers() { - for _, wallet := range mw.wallets { + for _, wallet := range mw.walletsReadCopy() { if wallet.IsAccountMixerActive() { unmixedAccount := wallet.ReadInt32ConfigValueForKey(AccountMixerUnmixedAccount, -1) hasMixableOutput, err := wallet.accountHasMixableOutput(unmixedAccount) diff --git a/types.go b/types.go index 4d05cf9d..90b33fc8 100644 --- a/types.go +++ b/types.go @@ -6,7 +6,6 @@ import ( "net" "decred.org/dcrwallet/v2/wallet/udb" - "github.com/decred/dcrd/chaincfg/v3" "github.com/decred/dcrd/dcrutil/v4" "github.com/planetdecred/dcrlibwallet/internal/vsp" diff --git a/wallets.go b/wallets.go index c5f97cc7..80105db1 100644 --- a/wallets.go +++ b/wallets.go @@ -1,6 +1,9 @@ package dcrlibwallet func (mw *MultiWallet) AllWallets() (wallets []*Wallet) { + mw.walletsMu.RLock() + defer mw.walletsMu.RUnlock() + for _, wallet := range mw.wallets { wallets = append(wallets, wallet) }