Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
132 changes: 85 additions & 47 deletions invoices/sql_migration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -543,88 +543,126 @@ func TestReconstructLegacyAMPStateMixedHTLCStates(t *testing.T) {
// NOTE: This test may need to be changed if the Invoice or any of the related
// types are modified.
func TestMigrateSingleInvoiceRapid(t *testing.T) {
// Create a shared Postgres instance for efficient testing.
pgFixture := sqldb.NewTestPgFixture(
t, sqldb.DefaultPostgresFixtureLifetime,
)
t.Cleanup(func() {
pgFixture.TearDown(t)
})
tests := []struct {
name string
sqlite bool
}{
{
name: "SQLite",
sqlite: true,
},
{
name: "Postgres",
},
}

makeSQLDB := func(t *testing.T, sqlite bool) *SQLStore {
var db *sqldb.BaseDB
if sqlite {
db = sqldb.NewTestSqliteDB(t).BaseDB
} else {
db = sqldb.NewTestPostgresDB(t, pgFixture).BaseDB
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
var pgFixture *sqldb.TestPgFixture
if !test.sqlite {
pgFixture = sqldb.NewTestPgFixture(
t, sqldb.DefaultPostgresFixtureLifetime,
)
t.Cleanup(func() {
pgFixture.TearDown(t)
})
}

executor := sqldb.NewTransactionExecutor(
db, func(tx *sql.Tx) SQLInvoiceQueries {
return db.WithTx(tx)
},
)
// Each property check uses a clean database. This keeps
// failures reproducible and isolated during shrinking.
rapid.Check(t, func(rt *rapid.T) {
var db *sqldb.BaseDB
if test.sqlite {
db = sqldb.NewTestSqliteDB(t).BaseDB
} else {
db = sqldb.NewTestPostgresDB(
t, pgFixture,
).BaseDB
}

testClock := clock.NewTestClock(time.Unix(1, 0))
// Helpers also clean up at the outer level.
// Close this handle after each check to avoid
// retaining database resources.
rt.Cleanup(func() {
require.NoError(rt, db.Close())
})

return NewSQLStore(executor, testClock)
executor := sqldb.NewTransactionExecutor(
db, func(tx *sql.Tx) SQLInvoiceQueries {
return db.WithTx(tx)
},
)
store := NewSQLStore(
executor,
clock.NewTestClock(time.Unix(1, 0)),
)

// Randomized feature flags for MPP and AMP.
mpp := rapid.Bool().Draw(rt, "mpp")
amp := rapid.Bool().Draw(rt, "amp")

testMigrateSingleInvoiceRapid(
rt, store, mpp, amp,
)
})
})
}

// Define property-based test using rapid.
rapid.Check(t, func(rt *rapid.T) {
// Randomized feature flags for MPP and AMP.
mpp := rapid.Bool().Draw(rt, "mpp")
amp := rapid.Bool().Draw(rt, "amp")

for _, sqlite := range []bool{true, false} {
store := makeSQLDB(t, sqlite)
testMigrateSingleInvoiceRapid(rt, store, mpp, amp)
}
})
}

// testMigrateSingleInvoiceRapid is the primary function for the migration of a
// single invoice with random data in a rapid-based test setup.
func testMigrateSingleInvoiceRapid(t *rapid.T, store *SQLStore, mpp bool,
amp bool) {

ctxb := t.Context()
invoices := make(map[lntypes.Hash]*Invoice)
const invoicesPerCheck = 10

for i := 0; i < 100; i++ {
type testInvoice struct {
hash lntypes.Hash
invoice *Invoice
}

ctxb := t.Context()
invoices := make([]testInvoice, 0, invoicesPerCheck)
for range invoicesPerCheck {
invoice := generateTestInvoiceRapid(t, mpp, amp)
var hash lntypes.Hash
_, err := crand.Read(hash[:])
require.NoError(t, err)

invoices[hash] = invoice
invoices = append(invoices, testInvoice{
hash: hash,
invoice: invoice,
})
}

ops := sqldb.WriteTxOpt()
err := store.db.ExecTx(ctxb, ops, func(tx SQLInvoiceQueries) error {
for hash, invoice := range invoices {
err := MigrateSingleInvoice(ctxb, tx, invoice, hash)
require.NoError(t, err)
for _, test := range invoices {
err := MigrateSingleInvoice(
ctxb, tx, test.invoice, test.hash,
)
if err != nil {
return err
}
}

return nil
}, sqldb.NoOpReset)
require.NoError(t, err)

// Fetch and compare each migrated invoice from the store with the
// original.
for hash, invoice := range invoices {
// Fetch and compare each migrated invoice with the original.
for _, test := range invoices {
sqlInvoice, err := store.LookupInvoice(
ctxb, InvoiceRefByHash(hash),
ctxb, InvoiceRefByHash(test.hash),
)
require.NoError(t, err)

invoice.AddIndex = sqlInvoice.AddIndex
test.invoice.AddIndex = sqlInvoice.AddIndex

OverrideInvoiceTimeZone(invoice)
OverrideInvoiceTimeZone(test.invoice)
OverrideInvoiceTimeZone(&sqlInvoice)

require.Equal(t, *invoice, sqlInvoice)
require.Equal(t, *test.invoice, sqlInvoice)
}
}

Expand Down
Loading