From 841d443a2d2ffa4184a26928124429782a722fc6 Mon Sep 17 00:00:00 2001 From: Matthew Date: Thu, 13 Aug 2026 15:03:10 -0500 Subject: [PATCH] fix(wallet): avoid deadlocks in custom persister callbacks --- bdk-ffi/src/store.rs | 72 +++++- bdk-ffi/src/tests/wallet.rs | 427 +++++++++++++++++++++++++++++++++++- bdk-ffi/src/wallet.rs | 124 +++++++++-- 3 files changed, 597 insertions(+), 26 deletions(-) diff --git a/bdk-ffi/src/store.rs b/bdk-ffi/src/store.rs index f543c8dcb..eadcf1d0c 100644 --- a/bdk-ffi/src/store.rs +++ b/bdk-ffi/src/store.rs @@ -7,8 +7,9 @@ use bdk_wallet::migration::{ }; use bdk_wallet::{rusqlite::Connection as BdkConnection, WalletPersister}; -use std::ops::DerefMut; -use std::sync::{Arc, Mutex}; +use std::ops::{Deref, DerefMut}; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::{Arc, Mutex, MutexGuard}; /// Definition of a wallet persistence implementation. #[uniffi::export(with_foreign)] @@ -25,6 +26,44 @@ pub(crate) enum PersistenceType { Sql(Mutex), } +// SQL operations retain the existing mutex guard. Custom operations retain only a +// non-blocking lease so foreign callbacks never run while `Persister::inner` is locked. +pub(crate) enum PersistenceOperation<'a> { + Locked(MutexGuard<'a, PersistenceType>), + Custom { + persistence: PersistenceType, + in_progress: &'a AtomicBool, + }, +} + +impl Deref for PersistenceOperation<'_> { + type Target = PersistenceType; + + fn deref(&self) -> &Self::Target { + match self { + Self::Locked(persistence) => persistence.deref(), + Self::Custom { persistence, .. } => persistence, + } + } +} + +impl DerefMut for PersistenceOperation<'_> { + fn deref_mut(&mut self) -> &mut Self::Target { + match self { + Self::Locked(persistence) => persistence.deref_mut(), + Self::Custom { persistence, .. } => persistence, + } + } +} + +impl Drop for PersistenceOperation<'_> { + fn drop(&mut self) { + if let Self::Custom { in_progress, .. } = self { + in_progress.store(false, Ordering::Release); + } + } +} + /// `PreV1WalletKeychain` represents a structure that holds the keychain details /// and metadata required for managing a wallet's keys. #[derive(Debug, Clone, uniffi::Record)] @@ -42,6 +81,8 @@ pub struct PreV1WalletKeychain { #[derive(uniffi::Object)] pub struct Persister { pub(crate) inner: Mutex, + // Serializes the full initialize/persist operation without blocking reentrant callbacks. + custom_operation_in_progress: AtomicBool, } #[uniffi::export] @@ -52,6 +93,7 @@ impl Persister { let conn = BdkConnection::open(path)?; Ok(Self { inner: PersistenceType::Sql(conn.into()).into(), + custom_operation_in_progress: AtomicBool::new(false), }) } @@ -61,6 +103,7 @@ impl Persister { let conn = BdkConnection::open_in_memory()?; Ok(Self { inner: PersistenceType::Sql(conn.into()).into(), + custom_operation_in_progress: AtomicBool::new(false), }) } @@ -69,6 +112,7 @@ impl Persister { pub fn custom(persistence: Arc) -> Self { Self { inner: PersistenceType::Custom(persistence).into(), + custom_operation_in_progress: AtomicBool::new(false), } } @@ -89,6 +133,30 @@ impl Persister { } } +impl Persister { + pub(crate) fn begin_operation(&self) -> Result, PersistenceError> { + let lock = self.inner.lock().unwrap(); + match lock.deref() { + PersistenceType::Sql(_) => Ok(PersistenceOperation::Locked(lock)), + PersistenceType::Custom(persistence) => { + // Waiting here could deadlock if a callback delegates reentry to another thread. + self.custom_operation_in_progress + .compare_exchange(false, true, Ordering::Acquire, Ordering::Relaxed) + .map_err(|_| PersistenceError::Reason { + error_message: "custom persistence operation already in progress" + .to_string(), + })?; + let persistence = Arc::clone(persistence); + drop(lock); + Ok(PersistenceOperation::Custom { + persistence: PersistenceType::Custom(persistence), + in_progress: &self.custom_operation_in_progress, + }) + } + } + } +} + impl From for PreV1WalletKeychain { fn from(value: BdkPreV1WalletKeychain) -> Self { Self { diff --git a/bdk-ffi/src/tests/wallet.rs b/bdk-ffi/src/tests/wallet.rs index 40fe2c5f2..ff8f15302 100644 --- a/bdk-ffi/src/tests/wallet.rs +++ b/bdk-ffi/src/tests/wallet.rs @@ -1,10 +1,10 @@ use crate::bitcoin::{Amount, BlockHash, Network, NetworkKind}; use crate::descriptor::Descriptor; -use crate::error::LoadWithPersistError; +use crate::error::{LoadWithPersistError, PersistenceError}; use crate::signer::SignersContainer; -use crate::store::Persister; +use crate::store::{Persistence, Persister}; use crate::tx_builder::TxBuilder; -use crate::types::Update; +use crate::types::{ChangeSet, Update}; use crate::wallet::{CreateParams, LoadParams, Wallet}; use bdk_wallet::bitcoin::Amount as BdkAmount; @@ -12,7 +12,11 @@ use bdk_wallet::bitcoin::Transaction as BdkTransaction; use bdk_wallet::bitcoin::{absolute, transaction, TxOut as BdkTxOut}; use bdk_wallet::KeychainKind; -use std::sync::Arc; +use std::panic::{catch_unwind, AssertUnwindSafe}; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use std::sync::{mpsc, Arc, Mutex, Weak}; +use std::thread; +use std::time::Duration; const EXTERNAL_DESCRIPTOR: &str = "wpkh(tprv8ZgxMBicQKsPf2qfrEygW6fdYseJDDrVnDv26PH5BHdvSuG6ecCbHqLVof9yZcMoM31z9ur3tTYbSnr1WBqbGX97CbXcmp5H6qeMpyvx35B/84h/1h/1h/0/*)"; const INTERNAL_DESCRIPTOR: &str = "wpkh(tprv8ZgxMBicQKsPf2qfrEygW6fdYseJDDrVnDv26PH5BHdvSuG6ecCbHqLVof9yZcMoM31z9ur3tTYbSnr1WBqbGX97CbXcmp5H6qeMpyvx35B/84h/1h/1h/1/*)"; @@ -448,3 +452,418 @@ fn test_load_from_two_path_descriptor_with_params() { error => panic!("expected InvalidChangeSet error, got {:?}", error), } } + +#[test] +fn test_custom_persistence_callback_can_read_same_wallet() { + struct ReentrantReadPersistence { + wallet: Arc, + reads: AtomicUsize, + } + + impl Persistence for ReentrantReadPersistence { + fn initialize(&self) -> Result, PersistenceError> { + Ok(Arc::new(ChangeSet::new())) + } + + fn persist(&self, _changeset: Arc) -> Result<(), PersistenceError> { + let _ = self.wallet.balance(); + self.reads.fetch_add(1, Ordering::Relaxed); + Ok(()) + } + } + + let wallet = Arc::new(build_wallet()); + wallet.reveal_next_address(KeychainKind::External); + let persistence = Arc::new(ReentrantReadPersistence { + wallet: Arc::clone(&wallet), + reads: AtomicUsize::new(0), + }); + let persister = Arc::new(Persister::custom(persistence.clone())); + let (sender, receiver) = mpsc::channel(); + let wallet_for_thread = Arc::clone(&wallet); + let handle = thread::spawn(move || { + sender.send(wallet_for_thread.persist(persister)).unwrap(); + }); + + let persisted = receiver + .recv_timeout(Duration::from_secs(2)) + .expect("reentrant persistence callback should not deadlock") + .unwrap(); + handle.join().unwrap(); + + assert!(persisted); + assert_eq!(persistence.reads.load(Ordering::Relaxed), 1); + assert!(wallet.staged().is_none()); +} + +#[test] +fn test_custom_persistence_callback_mutation_remains_staged() { + struct ReentrantMutationPersistence { + wallet: Arc, + calls: AtomicUsize, + persisted: Mutex>, + } + + impl Persistence for ReentrantMutationPersistence { + fn initialize(&self) -> Result, PersistenceError> { + Ok(Arc::new(ChangeSet::new())) + } + + fn persist(&self, changeset: Arc) -> Result<(), PersistenceError> { + self.persisted + .lock() + .unwrap() + .push(changeset.as_ref().clone().into()); + if self.calls.fetch_add(1, Ordering::Relaxed) == 0 { + self.wallet.reveal_next_address(KeychainKind::External); + } + Ok(()) + } + } + + let wallet = Arc::new(build_wallet()); + wallet.reveal_next_address(KeychainKind::External); + let persistence = Arc::new(ReentrantMutationPersistence { + wallet: Arc::clone(&wallet), + calls: AtomicUsize::new(0), + persisted: Mutex::new(Vec::new()), + }); + let persister = Arc::new(Persister::custom(persistence.clone())); + let (sender, receiver) = mpsc::channel(); + let wallet_for_thread = Arc::clone(&wallet); + let persister_for_thread = Arc::clone(&persister); + let handle = thread::spawn(move || { + sender + .send(wallet_for_thread.persist(persister_for_thread)) + .unwrap(); + }); + + let first_persisted = receiver + .recv_timeout(Duration::from_secs(2)) + .expect("reentrant wallet mutation should not deadlock") + .unwrap(); + handle.join().unwrap(); + + assert!(first_persisted); + assert_eq!(wallet.derivation_index(KeychainKind::External), Some(1)); + let staged: bdk_wallet::ChangeSet = wallet.staged().unwrap().as_ref().clone().into(); + assert_eq!( + staged.indexer.last_revealed.values().copied().max(), + Some(1) + ); + + assert!(wallet.persist(Arc::clone(&persister)).unwrap()); + assert!(!wallet.persist(persister).unwrap()); + + let persisted = persistence.persisted.lock().unwrap(); + let persisted_indexes: Vec> = persisted + .iter() + .map(|changeset| changeset.indexer.last_revealed.values().copied().max()) + .collect(); + assert_eq!(persisted_indexes, vec![Some(0), Some(1)]); + assert!(wallet.staged().is_none()); +} + +#[test] +fn test_custom_persistence_error_retains_reentrant_mutation() { + struct FailAfterMutationPersistence { + wallet: Arc, + should_fail: AtomicBool, + persisted: Mutex>, + } + + impl Persistence for FailAfterMutationPersistence { + fn initialize(&self) -> Result, PersistenceError> { + Ok(Arc::new(ChangeSet::new())) + } + + fn persist(&self, changeset: Arc) -> Result<(), PersistenceError> { + self.persisted + .lock() + .unwrap() + .push(changeset.as_ref().clone().into()); + if self.should_fail.swap(false, Ordering::Relaxed) { + self.wallet.reveal_next_address(KeychainKind::External); + return Err(PersistenceError::Reason { + error_message: "write failed".to_string(), + }); + } + Ok(()) + } + } + + let wallet = Arc::new(build_wallet()); + wallet.reveal_next_address(KeychainKind::External); + let persistence = Arc::new(FailAfterMutationPersistence { + wallet: Arc::clone(&wallet), + should_fail: AtomicBool::new(true), + persisted: Mutex::new(Vec::new()), + }); + let persister = Arc::new(Persister::custom(persistence.clone())); + + let error = wallet.persist(Arc::clone(&persister)).unwrap_err(); + assert!(error.to_string().contains("write failed")); + let staged: bdk_wallet::ChangeSet = wallet.staged().unwrap().as_ref().clone().into(); + assert_eq!( + staged.indexer.last_revealed.values().copied().max(), + Some(1) + ); + + assert!(wallet.persist(persister).unwrap()); + let persisted = persistence.persisted.lock().unwrap(); + let persisted_indexes: Vec> = persisted + .iter() + .map(|changeset| changeset.indexer.last_revealed.values().copied().max()) + .collect(); + assert_eq!(persisted_indexes, vec![Some(0), Some(1)]); + assert!(wallet.staged().is_none()); +} + +#[test] +fn test_nested_persist_on_same_wallet_fails_fast() { + struct CountingPersistence(AtomicUsize); + + impl Persistence for CountingPersistence { + fn initialize(&self) -> Result, PersistenceError> { + Ok(Arc::new(ChangeSet::new())) + } + + fn persist(&self, _changeset: Arc) -> Result<(), PersistenceError> { + self.0.fetch_add(1, Ordering::Relaxed); + Ok(()) + } + } + + struct NestedWalletPersistence { + wallet: Arc, + nested_persister: Arc, + nested_result: Mutex>>, + } + + impl Persistence for NestedWalletPersistence { + fn initialize(&self) -> Result, PersistenceError> { + Ok(Arc::new(ChangeSet::new())) + } + + fn persist(&self, _changeset: Arc) -> Result<(), PersistenceError> { + let result = self.wallet.persist(Arc::clone(&self.nested_persister)); + *self.nested_result.lock().unwrap() = Some(result); + Ok(()) + } + } + + let wallet = Arc::new(build_wallet()); + wallet.reveal_next_address(KeychainKind::External); + let nested_persistence = Arc::new(CountingPersistence(AtomicUsize::new(0))); + let nested_persister = Arc::new(Persister::custom(nested_persistence.clone())); + let persistence = Arc::new(NestedWalletPersistence { + wallet: Arc::clone(&wallet), + nested_persister, + nested_result: Mutex::new(None), + }); + let persister = Arc::new(Persister::custom(persistence.clone())); + let (sender, receiver) = mpsc::channel(); + let wallet_for_thread = Arc::clone(&wallet); + let handle = thread::spawn(move || { + sender.send(wallet_for_thread.persist(persister)).unwrap(); + }); + + let persisted = receiver + .recv_timeout(Duration::from_secs(2)) + .expect("nested wallet persistence should fail instead of deadlocking") + .unwrap(); + handle.join().unwrap(); + + assert!(persisted); + match persistence.nested_result.lock().unwrap().take().unwrap() { + Err(PersistenceError::Reason { error_message }) => { + assert_eq!( + error_message, + "wallet persistence operation already in progress" + ); + } + result => panic!("expected nested persistence to fail, got {:?}", result), + } + assert_eq!(nested_persistence.0.load(Ordering::Relaxed), 0); + assert!(wallet.staged().is_none()); +} + +#[test] +fn test_reentrant_persist_with_same_persister_fails_fast() { + struct ReentrantPersisterPersistence { + nested_wallet: Arc, + persister: Mutex>>, + calls: AtomicUsize, + nested_result: Mutex>>, + } + + impl Persistence for ReentrantPersisterPersistence { + fn initialize(&self) -> Result, PersistenceError> { + Ok(Arc::new(ChangeSet::new())) + } + + fn persist(&self, _changeset: Arc) -> Result<(), PersistenceError> { + if self.calls.fetch_add(1, Ordering::Relaxed) == 0 { + let persister = self + .persister + .lock() + .unwrap() + .as_ref() + .unwrap() + .upgrade() + .unwrap(); + let result = self.nested_wallet.persist(persister); + *self.nested_result.lock().unwrap() = Some(result); + } + Ok(()) + } + } + + let wallet = Arc::new(build_wallet()); + let nested_wallet = Arc::new(build_wallet()); + wallet.reveal_next_address(KeychainKind::External); + nested_wallet.reveal_next_address(KeychainKind::External); + let persistence = Arc::new(ReentrantPersisterPersistence { + nested_wallet: Arc::clone(&nested_wallet), + persister: Mutex::new(None), + calls: AtomicUsize::new(0), + nested_result: Mutex::new(None), + }); + let persister = Arc::new(Persister::custom(persistence.clone())); + *persistence.persister.lock().unwrap() = Some(Arc::downgrade(&persister)); + let (sender, receiver) = mpsc::channel(); + let wallet_for_thread = Arc::clone(&wallet); + let persister_for_thread = Arc::clone(&persister); + let handle = thread::spawn(move || { + sender + .send(wallet_for_thread.persist(persister_for_thread)) + .unwrap(); + }); + + let persisted = receiver + .recv_timeout(Duration::from_secs(2)) + .expect("reentrant persister use should fail instead of deadlocking") + .unwrap(); + handle.join().unwrap(); + + assert!(persisted); + match persistence.nested_result.lock().unwrap().take().unwrap() { + Err(PersistenceError::Reason { error_message }) => { + assert_eq!( + error_message, + "custom persistence operation already in progress" + ); + } + result => panic!("expected reentrant persister use to fail, got {:?}", result), + } + assert!(nested_wallet.staged().is_some()); + assert!(nested_wallet.persist(Arc::clone(&persister)).unwrap()); + assert_eq!(persistence.calls.load(Ordering::Relaxed), 2); + assert!(nested_wallet.staged().is_none()); +} + +#[test] +fn test_custom_persister_guards_complete_wallet_creation() { + struct ReentrantCreatePersistence { + persister: Mutex>>, + attempted: AtomicBool, + nested_result: Mutex>>, + } + + impl Persistence for ReentrantCreatePersistence { + fn initialize(&self) -> Result, PersistenceError> { + Ok(Arc::new(ChangeSet::new())) + } + + fn persist(&self, _changeset: Arc) -> Result<(), PersistenceError> { + if !self.attempted.swap(true, Ordering::Relaxed) { + let persister = self + .persister + .lock() + .unwrap() + .as_ref() + .unwrap() + .upgrade() + .unwrap(); + let result = Wallet::new( + external_descriptor(), + internal_descriptor(), + Network::Signet, + persister, + 25, + ) + .map(|_| ()) + .map_err(|error| error.to_string()); + *self.nested_result.lock().unwrap() = Some(result); + } + Ok(()) + } + } + + let persistence = Arc::new(ReentrantCreatePersistence { + persister: Mutex::new(None), + attempted: AtomicBool::new(false), + nested_result: Mutex::new(None), + }); + let persister = Arc::new(Persister::custom(persistence.clone())); + *persistence.persister.lock().unwrap() = Some(Arc::downgrade(&persister)); + let (sender, receiver) = mpsc::channel(); + let handle = thread::spawn(move || { + let result = Wallet::new( + external_descriptor(), + internal_descriptor(), + Network::Signet, + persister, + 25, + ) + .map(|_| ()) + .map_err(|error| error.to_string()); + sender.send(result).unwrap(); + }); + + receiver + .recv_timeout(Duration::from_secs(2)) + .expect("reentrant wallet creation should fail instead of deadlocking") + .unwrap(); + handle.join().unwrap(); + + let nested_error = persistence + .nested_result + .lock() + .unwrap() + .take() + .unwrap() + .unwrap_err(); + assert!(nested_error.contains("custom persistence operation already in progress")); +} + +#[test] +fn test_custom_persistence_panic_releases_operation_guards() { + struct PanicOncePersistence(AtomicBool); + + impl Persistence for PanicOncePersistence { + fn initialize(&self) -> Result, PersistenceError> { + Ok(Arc::new(ChangeSet::new())) + } + + fn persist(&self, _changeset: Arc) -> Result<(), PersistenceError> { + if self.0.swap(false, Ordering::Relaxed) { + panic!("persistence callback panicked"); + } + Ok(()) + } + } + + let wallet = Arc::new(build_wallet()); + wallet.reveal_next_address(KeychainKind::External); + let persister = Arc::new(Persister::custom(Arc::new(PanicOncePersistence( + AtomicBool::new(true), + )))); + + let panic_result = catch_unwind(AssertUnwindSafe(|| wallet.persist(Arc::clone(&persister)))); + assert!(panic_result.is_err()); + assert!(wallet.staged().is_some()); + + assert!(wallet.persist(persister).unwrap()); + assert!(wallet.staged().is_none()); +} diff --git a/bdk-ffi/src/wallet.rs b/bdk-ffi/src/wallet.rs index 4af75cfda..c385d638c 100644 --- a/bdk-ffi/src/wallet.rs +++ b/bdk-ffi/src/wallet.rs @@ -20,10 +20,11 @@ use bdk_wallet::keys::KeyMap; use bdk_wallet::signer::SignOptions as BdkSignOptions; use bdk_wallet::{ CreateParams as BdkCreateParams, LoadParams as BdkLoadParams, PersistedWallet, - Wallet as BdkWallet, + Wallet as BdkWallet, WalletPersister, }; use std::ops::DerefMut; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex, MutexGuard}; /// A Bitcoin wallet. @@ -41,6 +42,16 @@ use std::sync::{Arc, Mutex, MutexGuard}; #[derive(uniffi::Object)] pub struct Wallet { inner_mutex: Mutex>, + // Prevents overlapping snapshots from being persisted out of order. + persistence_in_progress: AtomicBool, +} + +struct WalletPersistenceOperation<'a>(&'a AtomicBool); + +impl Drop for WalletPersistenceOperation<'_> { + fn drop(&mut self) { + self.0.store(false, Ordering::Release); + } } /// Parameters for `Wallet` creation. @@ -151,8 +162,13 @@ impl Wallet { ) -> Result { let descriptor = descriptor.to_string_with_secret(); let change_descriptor = change_descriptor.to_string_with_secret(); - let mut persist_lock = persister.inner.lock().unwrap(); - let deref = persist_lock.deref_mut(); + let mut persist_operation = + persister + .begin_operation() + .map_err(|error| CreateWithPersistError::Persist { + error_message: error.to_string(), + })?; + let deref = persist_operation.deref_mut(); let bdk_params = BdkWallet::create(descriptor, change_descriptor).network(network); let bdk_params = params.apply_to(bdk_params); @@ -163,6 +179,7 @@ impl Wallet { Ok(Wallet { inner_mutex: Mutex::new(wallet), + persistence_in_progress: AtomicBool::new(false), }) } @@ -210,8 +227,13 @@ impl Wallet { params: CreateParams, ) -> Result { let descriptor = descriptor.to_string_with_secret(); - let mut persist_lock = persister.inner.lock().unwrap(); - let deref = persist_lock.deref_mut(); + let mut persist_operation = + persister + .begin_operation() + .map_err(|error| CreateWithPersistError::Persist { + error_message: error.to_string(), + })?; + let deref = persist_operation.deref_mut(); let bdk_params = BdkWallet::create_single(descriptor).network(network); let bdk_params = params.apply_to(bdk_params); @@ -222,6 +244,7 @@ impl Wallet { Ok(Wallet { inner_mutex: Mutex::new(wallet), + persistence_in_progress: AtomicBool::new(false), }) } @@ -265,8 +288,13 @@ impl Wallet { params: CreateParams, ) -> Result { let descriptor = two_path_descriptor.to_string_with_secret(); - let mut persist_lock = persister.inner.lock().unwrap(); - let deref = persist_lock.deref_mut(); + let mut persist_operation = + persister + .begin_operation() + .map_err(|error| CreateWithPersistError::Persist { + error_message: error.to_string(), + })?; + let deref = persist_operation.deref_mut(); let bdk_params = BdkWallet::create_from_two_path_descriptor(descriptor).network(network); let bdk_params = params.apply_to(bdk_params); @@ -277,6 +305,7 @@ impl Wallet { Ok(Wallet { inner_mutex: Mutex::new(wallet), + persistence_in_progress: AtomicBool::new(false), }) } @@ -310,8 +339,13 @@ impl Wallet { ) -> Result { let descriptor = descriptor.to_string_with_secret(); let change_descriptor = change_descriptor.to_string_with_secret(); - let mut persist_lock = persister.inner.lock().unwrap(); - let deref = persist_lock.deref_mut(); + let mut persist_operation = + persister + .begin_operation() + .map_err(|error| LoadWithPersistError::Persist { + error_message: error.to_string(), + })?; + let deref = persist_operation.deref_mut(); let bdk_params = BdkWallet::load() .descriptor(KeychainKind::External, Some(descriptor)) @@ -326,6 +360,7 @@ impl Wallet { Ok(Wallet { inner_mutex: Mutex::new(wallet), + persistence_in_progress: AtomicBool::new(false), }) } @@ -362,8 +397,13 @@ impl Wallet { params: LoadParams, ) -> Result { let descriptor = two_path_descriptor.to_string(); - let mut persist_lock = persister.inner.lock().unwrap(); - let deref = persist_lock.deref_mut(); + let mut persist_operation = + persister + .begin_operation() + .map_err(|error| LoadWithPersistError::Persist { + error_message: error.to_string(), + })?; + let deref = persist_operation.deref_mut(); let bdk_params = BdkWallet::load().two_path_descriptor(descriptor); let bdk_params = params.apply_to(bdk_params); @@ -375,6 +415,7 @@ impl Wallet { Ok(Wallet { inner_mutex: Mutex::new(wallet), + persistence_in_progress: AtomicBool::new(false), }) } @@ -400,8 +441,13 @@ impl Wallet { params: LoadParams, ) -> Result { let descriptor = descriptor.to_string_with_secret(); - let mut persist_lock = persister.inner.lock().unwrap(); - let deref = persist_lock.deref_mut(); + let mut persist_operation = + persister + .begin_operation() + .map_err(|error| LoadWithPersistError::Persist { + error_message: error.to_string(), + })?; + let deref = persist_operation.deref_mut(); let bdk_params = BdkWallet::load() .descriptor(KeychainKind::External, Some(descriptor)) @@ -415,6 +461,7 @@ impl Wallet { Ok(Wallet { inner_mutex: Mutex::new(wallet), + persistence_in_progress: AtomicBool::new(false), }) } @@ -942,14 +989,40 @@ impl Wallet { /// Returns whether any new changes were persisted. /// /// If the persister errors, the staged changes will not be cleared. + /// + /// Overlapping persistence calls for the same wallet or custom persister return an error. pub fn persist(&self, persister: Arc) -> Result { - let mut persist_lock = persister.inner.lock().unwrap(); - let deref = persist_lock.deref_mut(); - self.get_wallet() - .persist(deref) - .map_err(|e| PersistenceError::Reason { - error_message: e.to_string(), - }) + let _wallet_operation = self.begin_persistence_operation()?; + let mut persist_operation = persister.begin_operation()?; + + if matches!(&*persist_operation, PersistenceType::Sql(_)) { + return self + .get_wallet() + .persist(persist_operation.deref_mut()) + .map_err(|error| PersistenceError::Reason { + error_message: error.to_string(), + }); + } + + // Keep the stage in the wallet while the callback runs so errors and reentrant + // mutations cannot discard changes that have not been persisted successfully. + let Some(staged) = self.get_wallet().staged().cloned() else { + return Ok(false); + }; + + WalletPersister::persist(persist_operation.deref_mut(), &staged).map_err(|error| { + PersistenceError::Reason { + error_message: error.to_string(), + } + })?; + + let mut wallet = self.get_wallet(); + // A callback may mutate the wallet. Clear only the exact snapshot that was + // persisted; otherwise retain the merged stage for the next persistence call. + if wallet.staged() == Some(&staged) { + let _ = wallet.take_staged(); + } + Ok(true) } /// Get a reference of the staged [`ChangeSet`] that is yet to be committed (if any). @@ -997,6 +1070,17 @@ impl Wallet { } impl Wallet { + fn begin_persistence_operation( + &self, + ) -> Result, PersistenceError> { + self.persistence_in_progress + .compare_exchange(false, true, Ordering::Acquire, Ordering::Relaxed) + .map_err(|_| PersistenceError::Reason { + error_message: "wallet persistence operation already in progress".to_string(), + })?; + Ok(WalletPersistenceOperation(&self.persistence_in_progress)) + } + pub(crate) fn get_wallet(&self) -> MutexGuard<'_, PersistedWallet> { self.inner_mutex.lock().expect("wallet") }