diff --git a/sweepbatcher/sweep_batch.go b/sweepbatcher/sweep_batch.go index 1206b066..b4e6c65f 100644 --- a/sweepbatcher/sweep_batch.go +++ b/sweepbatcher/sweep_batch.go @@ -215,6 +215,12 @@ type batch struct { // reorgChan is the channel over which reorg notifications are received. reorgChan chan struct{} + // testReqs is a channel where test requests are received. + // This is used only in unit tests! The reason to have this is to + // avoid data races in require.Eventually calls running in parallel + // to the event loop. See method testRunInEventLoop(). + testReqs chan *testRequest + // errChan is the channel over which errors are received. errChan chan error @@ -352,6 +358,7 @@ func NewBatch(cfg batchConfig, bk batchKit) *batch { spendChan: make(chan *chainntnfs.SpendDetail), confChan: make(chan *chainntnfs.TxConfirmation, 1), reorgChan: make(chan struct{}, 1), + testReqs: make(chan *testRequest), errChan: make(chan error, 1), callEnter: make(chan struct{}), callLeave: make(chan struct{}), @@ -396,6 +403,7 @@ func NewBatchFromDB(cfg batchConfig, bk batchKit) (*batch, error) { spendChan: make(chan *chainntnfs.SpendDetail), confChan: make(chan *chainntnfs.TxConfirmation, 1), reorgChan: make(chan struct{}, 1), + testReqs: make(chan *testRequest), errChan: make(chan error, 1), callEnter: make(chan struct{}), callLeave: make(chan struct{}), @@ -756,6 +764,10 @@ func (b *batch) Run(ctx context.Context) error { return err } + case testReq := <-b.testReqs: + testReq.handler() + close(testReq.quit) + case err := <-blockErrChan: return err @@ -768,6 +780,36 @@ func (b *batch) Run(ctx context.Context) error { } } +// testRunInEventLoop runs a function in the event loop blocking until +// the function returns. For unit tests only! +func (b *batch) testRunInEventLoop(ctx context.Context, handler func()) { + // If the event loop is finished, run the function. + select { + case <-b.stopping: + handler() + + return + default: + } + + quit := make(chan struct{}) + req := &testRequest{ + handler: handler, + quit: quit, + } + + select { + case b.testReqs <- req: + case <-ctx.Done(): + return + } + + select { + case <-quit: + case <-ctx.Done(): + } +} + // timeout returns minimum timeout as block height among sweeps of the batch. // If the batch is empty, return -1. func (b *batch) timeout() int32 { diff --git a/sweepbatcher/sweep_batcher.go b/sweepbatcher/sweep_batcher.go index 21cb5fcb..7e3e95e9 100644 --- a/sweepbatcher/sweep_batcher.go +++ b/sweepbatcher/sweep_batcher.go @@ -225,6 +225,16 @@ var ( ErrBatcherShuttingDown = errors.New("batcher shutting down") ) +// testRequest is a function passed to an event loop and a channel used to +// wait until the function is executed. This is used in unit tests only! +type testRequest struct { + // handler is the function to an event loop. + handler func() + + // quit is closed when the handler completes. + quit chan struct{} +} + // Batcher is a system that is responsible for accepting sweep requests and // placing them in appropriate batches. It will spin up new batches as needed. type Batcher struct { @@ -234,6 +244,12 @@ type Batcher struct { // sweepReqs is a channel where sweep requests are received. sweepReqs chan SweepRequest + // testReqs is a channel where test requests are received. + // This is used only in unit tests! The reason to have this is to + // avoid data races in require.Eventually calls running in parallel + // to the event loop. See method testRunInEventLoop(). + testReqs chan *testRequest + // errChan is a channel where errors are received. errChan chan error @@ -461,6 +477,7 @@ func NewBatcher(wallet lndclient.WalletKitClient, return &Batcher{ batches: make(map[int32]*batch), sweepReqs: make(chan SweepRequest), + testReqs: make(chan *testRequest), errChan: make(chan error, 1), quit: make(chan struct{}), initDone: make(chan struct{}), @@ -528,6 +545,10 @@ func (b *Batcher) Run(ctx context.Context) error { return err } + case testReq := <-b.testReqs: + testReq.handler() + close(testReq.quit) + case err := <-b.errChan: log.Warnf("Batcher received an error: %v.", err) return err @@ -551,6 +572,36 @@ func (b *Batcher) AddSweep(sweepReq *SweepRequest) error { } } +// testRunInEventLoop runs a function in the event loop blocking until +// the function returns. For unit tests only! +func (b *Batcher) testRunInEventLoop(ctx context.Context, handler func()) { + // If the event loop is finished, run the function. + select { + case <-b.quit: + handler() + + return + default: + } + + quit := make(chan struct{}) + req := &testRequest{ + handler: handler, + quit: quit, + } + + select { + case b.testReqs <- req: + case <-ctx.Done(): + return + } + + select { + case <-quit: + case <-ctx.Done(): + } +} + // handleSweep handles a sweep request by either placing it in an existing // batch, or by spinning up a new batch for it. func (b *Batcher) handleSweep(ctx context.Context, sweep *sweep, diff --git a/sweepbatcher/sweep_batcher_test.go b/sweepbatcher/sweep_batcher_test.go index 1dcf2e35..d0ae4706 100644 --- a/sweepbatcher/sweep_batcher_test.go +++ b/sweepbatcher/sweep_batcher_test.go @@ -39,6 +39,8 @@ const ( eventuallyCheckFrequency = 100 * time.Millisecond ntfnBufferSize = 1024 + + confTarget = 123 ) // destAddr is a dummy p2wkh address to use as the destination address for @@ -109,18 +111,82 @@ func checkBatcherError(t *testing.T, err error) { } } +// getBatches returns batches in thread-safe way. +func getBatches(ctx context.Context, batcher *Batcher) []*batch { + var batches []*batch + batcher.testRunInEventLoop(ctx, func() { + for _, batch := range batcher.batches { + batches = append(batches, batch) + } + }) + + return batches +} + +// tryGetOnlyBatch returns a single batch if there is exactly one batch, or nil. +func tryGetOnlyBatch(ctx context.Context, batcher *Batcher) *batch { + batches := getBatches(ctx, batcher) + + if len(batches) == 1 { + return batches[0] + } else { + return nil + } +} + // getOnlyBatch makes sure the batcher has exactly one batch and returns it. -func getOnlyBatch(batcher *Batcher) *batch { - if len(batcher.batches) != 1 { - panic(fmt.Sprintf("getOnlyBatch called on a batcher having "+ - "%d batches", len(batcher.batches))) - } +func getOnlyBatch(t *testing.T, ctx context.Context, batcher *Batcher) *batch { + batches := getBatches(ctx, batcher) + require.Len(t, batches, 1) - for _, batch := range batcher.batches { - return batch - } + return batches[0] +} - panic("unreachable") +// numBatches returns the number of batches in the batcher. +func (b *Batcher) numBatches(ctx context.Context) int { + return len(getBatches(ctx, b)) +} + +// numSweeps returns the number of sweeps in the batch. +func (b *batch) numSweeps(ctx context.Context) int { + var numSweeps int + b.testRunInEventLoop(ctx, func() { + numSweeps = len(b.sweeps) + }) + + return numSweeps +} + +// snapshot returns the snapshot of the batch. It is safe to read in parallel +// with the event loop running. +func (b *batch) snapshot(ctx context.Context) *batch { + var snapshot *batch + b.testRunInEventLoop(ctx, func() { + // Deep copy sweeps. + sweeps := make(map[lntypes.Hash]sweep, len(b.sweeps)) + for h, s := range b.sweeps { + sweeps[h] = s + } + + // Deep copy cfg. + cfg := *b.cfg + + // Deep copy the batch, only data fields. + snapshot = &batch{ + id: b.id, + state: b.state, + primarySweepID: b.primarySweepID, + sweeps: sweeps, + currentHeight: b.currentHeight, + batchTxid: b.batchTxid, + batchPkScript: b.batchPkScript, + batchAddress: b.batchAddress, + rbfCache: b.rbfCache, + cfg: &cfg, + } + }) + + return snapshot } // testSweepBatcherBatchCreation tests that sweep requests enter the expected @@ -186,7 +252,7 @@ func testSweepBatcherBatchCreation(t *testing.T, store testStore, // Once batcher receives sweep request it will eventually spin up a // batch. require.Eventually(t, func() bool { - return len(batcher.batches) == 1 + return batcher.numBatches(ctx) == 1 }, test.Timeout, eventuallyCheckFrequency) // Wait for tx to be published. @@ -236,7 +302,7 @@ func testSweepBatcherBatchCreation(t *testing.T, store testStore, // Batcher should not create a second batch as timeout distance is small // enough. require.Eventually(t, func() bool { - return len(batcher.batches) == 1 + return batcher.numBatches(ctx) == 1 }, test.Timeout, eventuallyCheckFrequency) // Create a third sweep request that has more timeout distance than @@ -273,23 +339,26 @@ func testSweepBatcherBatchCreation(t *testing.T, store testStore, require.NoError(t, batcher.AddSweep(&sweepReq3)) - // Batcher should create a second batch as timeout distance is greater - // than the threshold - require.Eventually(t, func() bool { - return len(batcher.batches) == 2 - }, test.Timeout, eventuallyCheckFrequency) - // Since the second batch got created we check that it registered its // primary sweep's spend. <-lnd.RegisterSpendChannel + // Batcher should create a second batch as timeout distance is greater + // than the threshold + require.Eventually(t, func() bool { + return batcher.numBatches(ctx) == 2 + }, test.Timeout, eventuallyCheckFrequency) + // Wait for tx to be published. <-lnd.TxPublishChannel require.Eventually(t, func() bool { // Verify that each batch has the correct number of sweeps // in it. - for _, batch := range batcher.batches { + batches := getBatches(ctx, batcher) + + for _, batch := range batches { + batch := batch.snapshot(ctx) switch batch.primarySweepID { case sweepReq1.SwapHash: if len(batch.sweeps) != 2 { @@ -480,28 +549,30 @@ func testTxLabeler(t *testing.T, store testStore, // Deliver sweep request to batcher. require.NoError(t, batcher.AddSweep(&sweepReq1)) - // Eventually request will be consumed and a new batch will spin up. - require.Eventually(t, func() bool { - return len(batcher.batches) == 1 - }, test.Timeout, eventuallyCheckFrequency) - // When batch is successfully created it will execute it's first step, // which leads to a spend monitor of the primary sweep. <-lnd.RegisterSpendChannel + // Eventually request will be consumed and a new batch will spin up. + require.Eventually(t, func() bool { + return batcher.numBatches(ctx) == 1 + }, test.Timeout, eventuallyCheckFrequency) + // Wait for tx to be published. <-lnd.TxPublishChannel // Find the batch and assign it to a local variable for easier access. - var theBatch *batch - for _, btch := range batcher.batches { + var wantLabel string + for _, btch := range getBatches(ctx, batcher) { + btch := btch.snapshot(ctx) if btch.primarySweepID == sweepReq1.SwapHash { - theBatch = btch + wantLabel = fmt.Sprintf( + "BatchOutSweepSuccess -- %d", btch.id, + ) } } // Now test the label. - wantLabel := fmt.Sprintf("BatchOutSweepSuccess -- %d", theBatch.id) require.Equal(t, wantLabel, walletKit.lastLabel) // Now make the batcher quit by canceling the context. @@ -632,15 +703,15 @@ func testPublishErrorHandler(t *testing.T, store testStore, // Deliver sweep request to batcher. require.NoError(t, batcher.AddSweep(&sweepReq1)) - // Eventually request will be consumed and a new batch will spin up. - require.Eventually(t, func() bool { - return len(batcher.batches) == 1 - }, test.Timeout, eventuallyCheckFrequency) - // When batch is successfully created it will execute it's first step, // which leads to a spend monitor of the primary sweep. <-lnd.RegisterSpendChannel + // Eventually request will be consumed and a new batch will spin up. + require.Eventually(t, func() bool { + return batcher.numBatches(ctx) == 1 + }, test.Timeout, eventuallyCheckFrequency) + // The first attempt to publish the batch tx is expected to fail. require.ErrorIs(t, <-publishErrorChan, testPublishError) @@ -710,26 +781,28 @@ func testSweepBatcherSimpleLifecycle(t *testing.T, store testStore, // Deliver sweep request to batcher. require.NoError(t, batcher.AddSweep(&sweepReq1)) - // Eventually request will be consumed and a new batch will spin up. - require.Eventually(t, func() bool { - return len(batcher.batches) == 1 - }, test.Timeout, eventuallyCheckFrequency) - // When batch is successfully created it will execute it's first step, // which leads to a spend monitor of the primary sweep. <-lnd.RegisterSpendChannel + // Eventually request will be consumed and a new batch will spin up. + require.Eventually(t, func() bool { + return batcher.numBatches(ctx) == 1 + }, test.Timeout, eventuallyCheckFrequency) + // Find the batch and assign it to a local variable for easier access. batch := &batch{} - for _, btch := range batcher.batches { - if btch.primarySweepID == sweepReq1.SwapHash { - batch = btch - } + for _, btch := range getBatches(ctx, batcher) { + btch.testRunInEventLoop(ctx, func() { + if btch.primarySweepID == sweepReq1.SwapHash { + batch = btch + } + }) } require.Eventually(t, func() bool { // Batch should have the sweep stored. - return len(batch.sweeps) == 1 + return batch.numSweeps(ctx) == 1 }, test.Timeout, eventuallyCheckFrequency) // The primary sweep id should be that of the first inserted sweep. @@ -744,6 +817,8 @@ func testSweepBatcherSimpleLifecycle(t *testing.T, store testStore, // After receiving a height notification the batch will step again, // leading to a new spend monitoring. require.Eventually(t, func() bool { + batch := batch.snapshot(ctx) + return batch.currentHeight == 601 }, test.Timeout, eventuallyCheckFrequency) @@ -788,6 +863,8 @@ func testSweepBatcherSimpleLifecycle(t *testing.T, store testStore, // The batch should eventually read the spend notification and progress // its state to closed. require.Eventually(t, func() bool { + batch := batch.snapshot(ctx) + return batch.state == Closed }, test.Timeout, eventuallyCheckFrequency) @@ -895,7 +972,7 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { DestAddr: destAddr, SwapInvoice: swapInvoice, - SweepConfTarget: 123, + SweepConfTarget: confTarget, } err = store.CreateLoopOut(ctx, sweepReq.SwapHash, swap) @@ -938,16 +1015,14 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { // Eventually the batch is launched. require.Eventually(t, func() bool { - return len(batcher.batches) == 1 + return batcher.numBatches(ctx) == 1 }, test.Timeout, eventuallyCheckFrequency) // Replace the logger in the batch with wrappedLogger to watch messages. - var batch1 *batch - for _, batch := range batcher.batches { - batch1 = batch + batch1 := getOnlyBatch(t, ctx, batcher) + testLogger := &wrappedLogger{ + Logger: batch1.log(), } - require.NotNil(t, batch1) - testLogger := &wrappedLogger{Logger: batch1.log()} batch1.setLog(testLogger) // Advance the clock to publishDelay. It will trigger the publishDelay @@ -986,16 +1061,13 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { return false } - // Make sure there is exactly one active batch. - if len(batcher.batches) != 1 { + batch := tryGetOnlyBatch(ctx, batcher) + if batch == nil { return false } - // Get the batch. - batch := getOnlyBatch(batcher) - // Make sure the batch has one sweep. - return len(batch.sweeps) == 1 + return batch.numSweeps(ctx) == 1 }, test.Timeout, eventuallyCheckFrequency) // Make sure we have stored the batch. @@ -1031,25 +1103,6 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { // Wait for the batcher to be initialized. <-batcher.initDone - // Wait for batch to load. - require.Eventually(t, func() bool { - // Make sure that the sweep was stored - if !batcherStore.AssertSweepStored(sweepReq.SwapHash) { - return false - } - - // Make sure there is exactly one active batch. - if len(batcher.batches) != 1 { - return false - } - - // Get the batch. - batch := getOnlyBatch(batcher) - - // Make sure the batch has one sweep. - return len(batch.sweeps) == 1 - }, test.Timeout, eventuallyCheckFrequency) - // Expect a timer to be set: 0 (instead of publishDelay), and // RegisterSpend to be called. The order is not determined, so catch // these actions from two separate goroutines. @@ -1062,6 +1115,9 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { // Since a batch was created we check that it registered for its // primary sweep's spend. <-lnd.RegisterSpendChannel + + // Wait for tx to be published. + <-lnd.TxPublishChannel }() wg3.Add(1) @@ -1076,6 +1132,22 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { // Wait for RegisterSpend and for timer registration. wg3.Wait() + // Wait for batch to load. + require.Eventually(t, func() bool { + // Make sure that the sweep was stored + if !batcherStore.AssertSweepStored(sweepReq.SwapHash) { + return false + } + + batch := tryGetOnlyBatch(ctx, batcher) + if batch == nil { + return false + } + + // Make sure the batch has one sweep. + return batch.numSweeps(ctx) == 1 + }, test.Timeout, eventuallyCheckFrequency) + // Expect one timer: publishDelay (0). wantDelays = []time.Duration{0} require.Equal(t, wantDelays, delays) @@ -1084,9 +1156,6 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { now = now.Add(time.Millisecond) testClock.SetTime(now) - // Wait for tx to be published. - <-lnd.TxPublishChannel - // Tick tock next block. err = lnd.NotifyHeight(601) require.NoError(t, err) @@ -1195,7 +1264,7 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { DestAddr: destAddr, SwapInvoice: swapInvoice, - SweepConfTarget: 123, + SweepConfTarget: confTarget, } err = store.CreateLoopOut(ctx, sweepReq2.SwapHash, swap2) @@ -1237,15 +1306,16 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { require.Equal(t, wantDelays, delays) // Replace the logger in the batch with wrappedLogger to watch messages. - var batch2 *batch - for _, batch := range batcher.batches { + var testLogger2 *wrappedLogger + for _, batch := range getBatches(ctx, batcher) { if batch.id != batch1.id { - batch2 = batch + testLogger2 = &wrappedLogger{ + Logger: batch.log(), + } + batch.setLog(testLogger2) } } - require.NotNil(t, batch2) - testLogger2 := &wrappedLogger{Logger: batch2.log()} - batch2.setLog(testLogger2) + require.NotNil(t, testLogger2) // Add another sweep which is urgent. It will go to the same batch // to make sure minimum timeout is calculated properly. @@ -1273,7 +1343,7 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { DestAddr: destAddr, SwapInvoice: swapInvoice, - SweepConfTarget: 123, + SweepConfTarget: confTarget, } err = store.CreateLoopOut(ctx, sweepReq3.SwapHash, swap3) @@ -1297,7 +1367,7 @@ func testDelays(t *testing.T, store testStore, batcherStore testBatcherStore) { // Wait for tx to be published. tx := <-lnd.TxPublishChannel - require.Equal(t, 2, len(tx.TxIn)) + require.Len(t, tx.TxIn, 2) // Now make the batcher quit by canceling the context. cancel() @@ -1394,7 +1464,7 @@ func testMaxSweepsPerBatch(t *testing.T, store testStore, DestAddr: destAddr, SwapInvoice: swapInvoice, - SweepConfTarget: 123, + SweepConfTarget: confTarget, } err = store.CreateLoopOut(ctx, swapHash, swap) @@ -1413,15 +1483,17 @@ func testMaxSweepsPerBatch(t *testing.T, store testStore, // Eventually the batches are launched and all the sweeps are added. require.Eventually(t, func() bool { // Make sure all the batches have started. - if len(batcher.batches) != expectedBatches { + batches := getBatches(ctx, batcher) + if len(batches) != expectedBatches { return false } // Make sure all the sweeps were added. sweepsNum := 0 - for _, batch := range batcher.batches { - sweepsNum += len(batch.sweeps) + for _, batch := range batches { + sweepsNum += batch.numSweeps(ctx) } + return sweepsNum == swapsNum }, test.Timeout, eventuallyCheckFrequency) @@ -1602,20 +1674,22 @@ func testSweepBatcherSweepReentry(t *testing.T, store testStore, // Batcher should create a batch for the sweeps. require.Eventually(t, func() bool { - return len(batcher.batches) == 1 + return batcher.numBatches(ctx) == 1 }, test.Timeout, eventuallyCheckFrequency) // Find the batch and store it in a local variable for easier access. b := &batch{} - for _, btch := range batcher.batches { - if btch.primarySweepID == sweepReq1.SwapHash { - b = btch - } + for _, btch := range getBatches(ctx, batcher) { + btch.testRunInEventLoop(ctx, func() { + if btch.primarySweepID == sweepReq1.SwapHash { + b = btch + } + }) } // Batcher should contain all sweeps. require.Eventually(t, func() bool { - return len(b.sweeps) == 3 + return b.numSweeps(ctx) == 3 }, test.Timeout, eventuallyCheckFrequency) // Verify that the batch has a primary sweep id that matches the first @@ -1664,20 +1738,22 @@ func testSweepBatcherSweepReentry(t *testing.T, store testStore, // Eventually the batch reads the notification and proceeds to a closed // state. require.Eventually(t, func() bool { - return b.state == Closed - }, test.Timeout, eventuallyCheckFrequency) + b := b.snapshot(ctx) - // While handling the spend notification the batch should detect that - // some sweeps did not appear in the spending tx, therefore it redirects - // them back to the batcher and the batcher inserts them in a new batch. - require.Eventually(t, func() bool { - return len(batcher.batches) == 2 + return b.state == Closed }, test.Timeout, eventuallyCheckFrequency) // Since second batch was created we check that it registered for its // primary sweep's spend. <-lnd.RegisterSpendChannel + // While handling the spend notification the batch should detect that + // some sweeps did not appear in the spending tx, therefore it redirects + // them back to the batcher and the batcher inserts them in a new batch. + require.Eventually(t, func() bool { + return batcher.numBatches(ctx) == 2 + }, test.Timeout, eventuallyCheckFrequency) + // We mock the confirmation notification. lnd.ConfChannel <- &chainntnfs.TxConfirmation{ Tx: spendingTx, @@ -1692,26 +1768,28 @@ func testSweepBatcherSweepReentry(t *testing.T, store testStore, // confirmation forever. <-lnd.TxPublishChannel + // Re-add one of remaining sweeps to trigger removing the completed + // batch from the batcher. + require.NoError(t, batcher.AddSweep(&sweepReq3)) + // Eventually the batch receives the confirmation notification, // gracefully exits and the batcher deletes it. require.Eventually(t, func() bool { - return len(batcher.batches) == 1 + return batcher.numBatches(ctx) == 1 }, test.Timeout, eventuallyCheckFrequency) // Find the other batch, which includes the sweeps that did not appear // in the spending tx. - b = &batch{} - for _, btch := range batcher.batches { - b = btch - } + b = getOnlyBatch(t, ctx, batcher) // After all the sweeps enter, it should contain 2 sweeps. require.Eventually(t, func() bool { - return len(b.sweeps) == 2 + return b.numSweeps(ctx) == 2 }, test.Timeout, eventuallyCheckFrequency) // The batch should be in an open state. - require.Equal(t, b.state, Open) + b1 := b.snapshot(ctx) + require.Equal(t, b1.state, Open) } // testSweepBatcherNonWalletAddr tests that sweep requests that sweep to a non @@ -1767,16 +1845,16 @@ func testSweepBatcherNonWalletAddr(t *testing.T, store testStore, // Deliver sweep request to batcher. require.NoError(t, batcher.AddSweep(&sweepReq1)) - // Once batcher receives sweep request it will eventually spin up a - // batch. - require.Eventually(t, func() bool { - return len(batcher.batches) == 1 - }, test.Timeout, eventuallyCheckFrequency) - // Since a batch was created we check that it registered for its primary // sweep's spend. <-lnd.RegisterSpendChannel + // Once batcher receives sweep request it will eventually spin up a + // batch. + require.Eventually(t, func() bool { + return batcher.numBatches(ctx) == 1 + }, test.Timeout, eventuallyCheckFrequency) + // Wait for tx to be published. <-lnd.TxPublishChannel @@ -1817,16 +1895,16 @@ func testSweepBatcherNonWalletAddr(t *testing.T, store testStore, require.NoError(t, batcher.AddSweep(&sweepReq2)) - // Batcher should create a second batch as first batch is a non wallet - // addr batch. - require.Eventually(t, func() bool { - return len(batcher.batches) == 2 - }, test.Timeout, eventuallyCheckFrequency) - // Since a batch was created we check that it registered for its primary // sweep's spend. <-lnd.RegisterSpendChannel + // Batcher should create a second batch as first batch is a non wallet + // addr batch. + require.Eventually(t, func() bool { + return batcher.numBatches(ctx) == 2 + }, test.Timeout, eventuallyCheckFrequency) + // Wait for second batch to be published. <-lnd.TxPublishChannel @@ -1864,23 +1942,25 @@ func testSweepBatcherNonWalletAddr(t *testing.T, store testStore, require.NoError(t, batcher.AddSweep(&sweepReq3)) - // Batcher should create a new batch as timeout distance is greater than - // the threshold - require.Eventually(t, func() bool { - return len(batcher.batches) == 3 - }, test.Timeout, eventuallyCheckFrequency) - // Since a batch was created we check that it registered for its primary // sweep's spend. <-lnd.RegisterSpendChannel + // Batcher should create a new batch as timeout distance is greater than + // the threshold + require.Eventually(t, func() bool { + return batcher.numBatches(ctx) == 3 + }, test.Timeout, eventuallyCheckFrequency) + // Wait for tx to be published for 3rd batch. <-lnd.TxPublishChannel require.Eventually(t, func() bool { // Verify that each batch has the correct number of sweeps // in it. - for _, batch := range batcher.batches { + batches := getBatches(ctx, batcher) + for _, batch := range batches { + batch := batch.snapshot(ctx) switch batch.primarySweepID { case sweepReq1.SwapHash: if len(batch.sweeps) != 1 { @@ -2117,16 +2197,16 @@ func testSweepBatcherComposite(t *testing.T, store testStore, // Deliver sweep request to batcher. require.NoError(t, batcher.AddSweep(&sweepReq1)) - // Once batcher receives sweep request it will eventually spin up a - // batch. - require.Eventually(t, func() bool { - return len(batcher.batches) == 1 - }, test.Timeout, eventuallyCheckFrequency) - // Since a batch was created we check that it registered for its primary // sweep's spend. <-lnd.RegisterSpendChannel + // Once batcher receives sweep request it will eventually spin up a + // batch. + require.Eventually(t, func() bool { + return batcher.numBatches(ctx) == 1 + }, test.Timeout, eventuallyCheckFrequency) + // Wait for tx to be published. <-lnd.TxPublishChannel @@ -2138,7 +2218,7 @@ func testSweepBatcherComposite(t *testing.T, store testStore, // Batcher should not create a second batch as timeout distance is small // enough. require.Eventually(t, func() bool { - return len(batcher.batches) == 1 + return batcher.numBatches(ctx) == 1 }, test.Timeout, eventuallyCheckFrequency) // Publish a block to trigger batch 1 republishing. @@ -2147,39 +2227,39 @@ func testSweepBatcherComposite(t *testing.T, store testStore, // Wait for tx for the first batch to be published (2 sweeps). tx := <-lnd.TxPublishChannel - require.Equal(t, 2, len(tx.TxIn)) + require.Len(t, tx.TxIn, 2) require.NoError(t, batcher.AddSweep(&sweepReq3)) + // Since a batch was created we check that it registered for its primary + // sweep's spend. + <-lnd.RegisterSpendChannel + // Batcher should create a second batch as this sweep pays to a non // wallet address. require.Eventually(t, func() bool { - return len(batcher.batches) == 2 + return batcher.numBatches(ctx) == 2 }, test.Timeout, eventuallyCheckFrequency) + // Wait for tx for the second batch to be published (1 sweep). + tx = <-lnd.TxPublishChannel + require.Len(t, tx.TxIn, 1) + + require.NoError(t, batcher.AddSweep(&sweepReq4)) + // Since a batch was created we check that it registered for its primary // sweep's spend. <-lnd.RegisterSpendChannel - // Wait for tx for the second batch to be published (1 sweep). - tx = <-lnd.TxPublishChannel - require.Equal(t, 1, len(tx.TxIn)) - - require.NoError(t, batcher.AddSweep(&sweepReq4)) - // Batcher should create a third batch as timeout distance is greater // than the threshold. require.Eventually(t, func() bool { - return len(batcher.batches) == 3 + return batcher.numBatches(ctx) == 3 }, test.Timeout, eventuallyCheckFrequency) - // Since a batch was created we check that it registered for its primary - // sweep's spend. - <-lnd.RegisterSpendChannel - // Wait for tx for the third batch to be published (1 sweep). tx = <-lnd.TxPublishChannel - require.Equal(t, 1, len(tx.TxIn)) + require.Len(t, tx.TxIn, 1) require.NoError(t, batcher.AddSweep(&sweepReq5)) @@ -2195,29 +2275,31 @@ func testSweepBatcherComposite(t *testing.T, store testStore, // Batcher should not create a fourth batch as timeout distance is small // enough for it to join the last batch. require.Eventually(t, func() bool { - return len(batcher.batches) == 3 + return batcher.numBatches(ctx) == 3 }, test.Timeout, eventuallyCheckFrequency) require.NoError(t, batcher.AddSweep(&sweepReq6)) - // Batcher should create a fourth batch as this sweep pays to a non - // wallet address. - require.Eventually(t, func() bool { - return len(batcher.batches) == 4 - }, test.Timeout, eventuallyCheckFrequency) - // Since a batch was created we check that it registered for its primary // sweep's spend. <-lnd.RegisterSpendChannel + // Batcher should create a fourth batch as this sweep pays to a non + // wallet address. + require.Eventually(t, func() bool { + return batcher.numBatches(ctx) == 4 + }, test.Timeout, eventuallyCheckFrequency) + // Wait for tx for the 4th batch to be published (1 sweep). tx = <-lnd.TxPublishChannel - require.Equal(t, 1, len(tx.TxIn)) + require.Len(t, tx.TxIn, 1) require.Eventually(t, func() bool { // Verify that each batch has the correct number of sweeps in // it. - for _, batch := range batcher.batches { + batches := getBatches(ctx, batcher) + for _, batch := range batches { + batch := batch.snapshot(ctx) switch batch.primarySweepID { case sweepReq1.SwapHash: if len(batch.sweeps) != 2 { @@ -2374,8 +2456,11 @@ func testRestoringEmptyBatch(t *testing.T, store testStore, require.Eventually(t, func() bool { // Make sure that the sweep was stored and we have exactly one // active batch. - return batcherStore.AssertSweepStored(sweepReq.SwapHash) && - len(batcher.batches) == 1 + if !batcherStore.AssertSweepStored(sweepReq.SwapHash) { + return false + } + + return batcher.numBatches(ctx) == 1 }, test.Timeout, eventuallyCheckFrequency) // Make sure we have only one batch stored (as we dropped the dormant @@ -2591,9 +2676,14 @@ func testHandleSweepTwice(t *testing.T, backend testStore, require.Eventually(t, func() bool { // Make sure that the sweep was stored and we have exactly one // active batch. - return batcherStore.AssertSweepStored(sweepReq1.SwapHash) && - batcherStore.AssertSweepStored(sweepReq2.SwapHash) && - len(batcher.batches) == 2 + if !batcherStore.AssertSweepStored(sweepReq1.SwapHash) { + return false + } + if !batcherStore.AssertSweepStored(sweepReq2.SwapHash) { + return false + } + + return batcher.numBatches(ctx) == 2 }, test.Timeout, eventuallyCheckFrequency) // Change the second sweep so that it can be added to the first batch. @@ -2622,7 +2712,8 @@ func testHandleSweepTwice(t *testing.T, backend testStore, require.Eventually(t, func() bool { // Make sure there are two batches. - batches := batcher.batches + batches := getBatches(ctx, batcher) + if len(batches) != 2 { return false } @@ -2636,9 +2727,10 @@ func testHandleSweepTwice(t *testing.T, backend testStore, secondBatch = batch } } + snapshot := secondBatch.snapshot(ctx) // Make sure the second batch has the second sweep. - sweep2, has := secondBatch.sweeps[sweepReq2.SwapHash] + sweep2, has := snapshot.sweeps[sweepReq2.SwapHash] if !has { return false } @@ -2649,8 +2741,10 @@ func testHandleSweepTwice(t *testing.T, backend testStore, // Make sure each batch has one sweep. If the second sweep was added to // both batches, the following check won't pass. - for _, batch := range batcher.batches { - require.Equal(t, 1, len(batch.sweeps)) + batches := getBatches(ctx, batcher) + for _, batch := range batches { + // Make sure the batch has one sweep. + require.Equal(t, 1, batch.numSweeps(ctx)) } // Publish a block to trigger batch 2 republishing. @@ -2718,7 +2812,7 @@ func testRestoringPreservesConfTarget(t *testing.T, store testStore, DestAddr: destAddr, SwapInvoice: swapInvoice, - SweepConfTarget: 123, + SweepConfTarget: confTarget, } err = store.CreateLoopOut(ctx, sweepReq.SwapHash, swap) @@ -2743,21 +2837,21 @@ func testRestoringPreservesConfTarget(t *testing.T, store testStore, return false } - // Make sure there is exactly one active batch. - if len(batcher.batches) != 1 { + batch := tryGetOnlyBatch(ctx, batcher) + if batch == nil { return false } - // Get the batch. - batch := getOnlyBatch(batcher) + // Make sure the batch has one sweep. + snapshot := batch.snapshot(ctx) // Make sure the batch has one sweep. - if len(batch.sweeps) != 1 { + if len(snapshot.sweeps) != 1 { return false } // Make sure the batch has proper batchConfTarget. - return batch.cfg.batchConfTarget == 123 + return snapshot.cfg.batchConfTarget == confTarget }, test.Timeout, eventuallyCheckFrequency) // Make sure we have stored the batch. @@ -2786,6 +2880,12 @@ func testRestoringPreservesConfTarget(t *testing.T, store testStore, // Wait for the batcher to be initialized. <-batcher.initDone + // Expect registration for spend notification. + <-lnd.RegisterSpendChannel + + // Wait for tx to be published. + <-lnd.TxPublishChannel + // Wait for batch to load. require.Eventually(t, func() bool { // Make sure that the sweep was stored @@ -2793,26 +2893,18 @@ func testRestoringPreservesConfTarget(t *testing.T, store testStore, return false } - // Make sure there is exactly one active batch. - if len(batcher.batches) != 1 { + batch := tryGetOnlyBatch(ctx, batcher) + if batch == nil { return false } - // Get the batch. - batch := getOnlyBatch(batcher) - // Make sure the batch has one sweep. - return len(batch.sweeps) == 1 + return batch.numSweeps(ctx) == 1 }, test.Timeout, eventuallyCheckFrequency) // Make sure batchConfTarget was preserved. - require.Equal(t, 123, int(getOnlyBatch(batcher).cfg.batchConfTarget)) - - // Expect registration for spend notification. - <-lnd.RegisterSpendChannel - - // Wait for tx to be published. - <-lnd.TxPublishChannel + batch := getOnlyBatch(t, ctx, batcher).snapshot(ctx) + require.Equal(t, int32(confTarget), batch.cfg.batchConfTarget) // Now make the batcher quit by canceling the context. cancel() @@ -2884,7 +2976,7 @@ func testSweepFetcher(t *testing.T, store testStore, require.NoError(t, err) sweepInfo := &SweepInfo{ - ConfTarget: 123, + ConfTarget: confTarget, Timeout: 111, SwapInvoicePaymentAddr: *swapPaymentAddr, ProtocolVersion: loopdb.ProtocolVersionMuSig2, @@ -2957,21 +3049,20 @@ func testSweepFetcher(t *testing.T, store testStore, return false } - // Make sure there is exactly one active batch. - if len(batcher.batches) != 1 { + // Try to get the batch. + batch := tryGetOnlyBatch(ctx, batcher) + if batch == nil { return false } - // Get the batch. - batch := getOnlyBatch(batcher) - // Make sure the batch has one sweep. - if len(batch.sweeps) != 1 { + snapshot := batch.snapshot(ctx) + if len(snapshot.sweeps) != 1 { return false } // Make sure the batch has proper batchConfTarget. - return batch.cfg.batchConfTarget == 123 + return snapshot.cfg.batchConfTarget == confTarget }, test.Timeout, eventuallyCheckFrequency) // Get the published transaction and check the fee rate. @@ -3291,7 +3382,7 @@ func testWithMixedBatch(t *testing.T, store testStore, sweepInfo := &SweepInfo{ Preimage: preimages[i], - ConfTarget: 123, + ConfTarget: confTarget, Timeout: 111, SwapInvoicePaymentAddr: *swapPaymentAddr, ProtocolVersion: loopdb.ProtocolVersionMuSig2, @@ -3330,7 +3421,7 @@ func testWithMixedBatch(t *testing.T, store testStore, // A transaction is published. tx := <-lnd.TxPublishChannel - require.Equal(t, i+1, len(tx.TxIn)) + require.Len(t, tx.TxIn, i+1) // Check types of inputs. var witnessSizes []int @@ -3462,7 +3553,7 @@ func testWithMixedBatchCustom(t *testing.T, store testStore, Preimage: preimages[i], NonCoopHint: nonCoopHints[i], - ConfTarget: 123, + ConfTarget: confTarget, Timeout: 111, SwapInvoicePaymentAddr: *swapPaymentAddr, ProtocolVersion: loopdb.ProtocolVersionMuSig2, @@ -3499,7 +3590,7 @@ func testWithMixedBatchCustom(t *testing.T, store testStore, // A transaction is published. tx := <-lnd.TxPublishChannel - require.Equal(t, len(preimages), len(tx.TxIn)) + require.Len(t, tx.TxIn, len(preimages)) // Check types of inputs. var witnessSizes []int @@ -3819,9 +3910,10 @@ func testFeeRateGrows(t *testing.T, store testStore, <-lnd.TxPublishChannel // Make sure the fee rate is feeRateMedium. - batch := getOnlyBatch(batcher) - require.Len(t, batch.sweeps, 1) - require.Equal(t, feeRateMedium, batch.rbfCache.FeeRate) + batch := getOnlyBatch(t, ctx, batcher) + snapshot := batch.snapshot(ctx) + require.Len(t, snapshot.sweeps, 1) + require.Equal(t, feeRateMedium, snapshot.rbfCache.FeeRate) // Now decrease the fee of sweep1. setFeeRate(swapHash1, feeRateLow) @@ -3835,7 +3927,8 @@ func testFeeRateGrows(t *testing.T, store testStore, <-lnd.TxPublishChannel // Make sure the fee rate is still feeRateMedium. - require.Equal(t, feeRateMedium, batch.rbfCache.FeeRate) + snapshot = batch.snapshot(ctx) + require.Equal(t, feeRateMedium, snapshot.rbfCache.FeeRate) // Add sweep2, with feeRateMedium. swapHash2 := lntypes.Hash{2, 2, 2} @@ -3881,8 +3974,9 @@ func testFeeRateGrows(t *testing.T, store testStore, <-lnd.TxPublishChannel // Make sure the fee rate is still feeRateMedium. - require.Len(t, batch.sweeps, 2) - require.Equal(t, feeRateMedium, batch.rbfCache.FeeRate) + snapshot = batch.snapshot(ctx) + require.Len(t, snapshot.sweeps, 2) + require.Equal(t, feeRateMedium, snapshot.rbfCache.FeeRate) // Now update fee rate of second sweep (which is not primary) to // feeRateHigh. Fee rate of sweep 1 is still feeRateLow. @@ -3898,7 +3992,8 @@ func testFeeRateGrows(t *testing.T, store testStore, <-lnd.TxPublishChannel // Make sure the fee rate increased to feeRateHigh. - require.Equal(t, feeRateHigh, batch.rbfCache.FeeRate) + snapshot = batch.snapshot(ctx) + require.Equal(t, feeRateHigh, snapshot.rbfCache.FeeRate) } // TestSweepBatcherBatchCreation tests that sweep requests enter the expected