diff --git a/sweepbatcher/sweep_batch.go b/sweepbatcher/sweep_batch.go index 05561b68..b76ed3aa 100644 --- a/sweepbatcher/sweep_batch.go +++ b/sweepbatcher/sweep_batch.go @@ -197,7 +197,11 @@ type batch struct { // main event loop. callLeave chan struct{} - // quit signals that the batch must stop. + // stopped signals that the batch has stopped. + stopped chan struct{} + + // quit is owned by the parent batcher and signals that the batch must + // stop. quit chan struct{} // wallet is the wallet client used to create and publish the batch @@ -261,6 +265,7 @@ type batchKit struct { purger Purger store BatcherStore log btclog.Logger + quit chan struct{} } // scheduleNextCall schedules the next call to the batch handler's main event @@ -270,6 +275,9 @@ func (b *batch) scheduleNextCall() (func(), error) { case b.callEnter <- struct{}{}: case <-b.quit: + return func() {}, ErrBatcherShuttingDown + + case <-b.stopped: return func() {}, ErrBatchShuttingDown } @@ -293,7 +301,8 @@ func NewBatch(cfg batchConfig, bk batchKit) *batch { errChan: make(chan error, 1), callEnter: make(chan struct{}), callLeave: make(chan struct{}), - quit: make(chan struct{}), + stopped: make(chan struct{}), + quit: bk.quit, batchTxid: bk.batchTxid, wallet: bk.wallet, chainNotifier: bk.chainNotifier, @@ -320,7 +329,8 @@ func NewBatchFromDB(cfg batchConfig, bk batchKit) *batch { errChan: make(chan error, 1), callEnter: make(chan struct{}), callLeave: make(chan struct{}), - quit: make(chan struct{}), + stopped: make(chan struct{}), + quit: bk.quit, batchTxid: bk.batchTxid, batchPkScript: bk.batchPkScript, rbfCache: bk.rbfCache, @@ -447,7 +457,7 @@ func (b *batch) Run(ctx context.Context) error { runCtx, cancel := context.WithCancel(ctx) defer func() { cancel() - close(b.quit) + close(b.stopped) b.wg.Wait() }() diff --git a/sweepbatcher/sweep_batcher.go b/sweepbatcher/sweep_batcher.go index a0b254a4..8816cd16 100644 --- a/sweepbatcher/sweep_batcher.go +++ b/sweepbatcher/sweep_batcher.go @@ -216,6 +216,7 @@ func (b *Batcher) Run(ctx context.Context) error { runCtx, cancel := context.WithCancel(ctx) defer func() { cancel() + close(b.quit) for _, batch := range b.batches { batch.Wait() @@ -379,6 +380,7 @@ func (b *Batcher) spinUpBatch(ctx context.Context) (*batch, error) { verifySchnorrSig: b.VerifySchnorrSig, purger: b.AddSweep, store: b.store, + quit: b.quit, } batch := NewBatch(cfg, batchKit) @@ -461,6 +463,7 @@ func (b *Batcher) spinUpBatchFromDB(ctx context.Context, batch *batch) error { purger: b.AddSweep, store: b.store, log: batchPrefixLogger(fmt.Sprintf("%d", batch.id)), + quit: b.quit, } cfg := batchConfig{ diff --git a/sweepbatcher/sweep_batcher_test.go b/sweepbatcher/sweep_batcher_test.go index fa56ff60..4ddbb656 100644 --- a/sweepbatcher/sweep_batcher_test.go +++ b/sweepbatcher/sweep_batcher_test.go @@ -2,7 +2,7 @@ package sweepbatcher import ( "context" - "strings" + "errors" "testing" "time" @@ -43,6 +43,15 @@ var dummyNotifier = SpendNotifier{ QuitChan: make(chan bool, ntfnBufferSize), } +func checkBatcherError(t *testing.T, err error) { + if !errors.Is(err, context.Canceled) && + !errors.Is(err, ErrBatcherShuttingDown) && + !errors.Is(err, ErrBatchShuttingDown) { + + require.NoError(t, err) + } +} + // TestSweepBatcherBatchCreation tests that sweep requests enter the expected // batch based on their timeout distance. func TestSweepBatcherBatchCreation(t *testing.T) { @@ -60,9 +69,7 @@ func TestSweepBatcherBatchCreation(t *testing.T) { testMuSig2SignSweep, nil, lnd.ChainParams, batcherStore, store) go func() { err := batcher.Run(ctx) - if !strings.Contains(err.Error(), "context canceled") { - require.NoError(t, err) - } + checkBatcherError(t, err) }() // Create a sweep request. @@ -215,9 +222,7 @@ func TestSweepBatcherSimpleLifecycle(t *testing.T) { testMuSig2SignSweep, nil, lnd.ChainParams, batcherStore, store) go func() { err := batcher.Run(ctx) - if !strings.Contains(err.Error(), "context canceled") { - require.NoError(t, err) - } + checkBatcherError(t, err) }() // Create a sweep request. @@ -354,9 +359,7 @@ func TestSweepBatcherSweepReentry(t *testing.T) { testMuSig2SignSweep, nil, lnd.ChainParams, batcherStore, store) go func() { err := batcher.Run(ctx) - if !strings.Contains(err.Error(), "context canceled") { - require.NoError(t, err) - } + checkBatcherError(t, err) }() // Create some sweep requests with timeouts not too far away, in order @@ -561,9 +564,7 @@ func TestSweepBatcherNonWalletAddr(t *testing.T) { testMuSig2SignSweep, nil, lnd.ChainParams, batcherStore, store) go func() { err := batcher.Run(ctx) - if !strings.Contains(err.Error(), "context canceled") { - require.NoError(t, err) - } + checkBatcherError(t, err) }() // Create a sweep request. @@ -727,9 +728,7 @@ func TestSweepBatcherComposite(t *testing.T) { testMuSig2SignSweep, nil, lnd.ChainParams, batcherStore, store) go func() { err := batcher.Run(ctx) - if !strings.Contains(err.Error(), "context canceled") { - require.NoError(t, err) - } + checkBatcherError(t, err) }() // Create a sweep request.