lightning-terminal/accounts/sql_migration.go
2026-05-19 17:28:45 -07:00

409 lines
11 KiB
Go

package accounts
import (
"bytes"
"context"
"database/sql"
"errors"
"fmt"
"math"
"reflect"
"time"
"github.com/davecgh/go-spew/spew"
"github.com/lightninglabs/lightning-terminal/db/sqlcmig6"
"github.com/lightninglabs/lightning-terminal/db/tombstone"
"github.com/lightningnetwork/lnd/kvdb"
"github.com/lightningnetwork/lnd/lnrpc"
"github.com/lightningnetwork/lnd/lntypes"
"github.com/lightningnetwork/lnd/lnwire"
"github.com/pmezard/go-difflib/difflib"
)
var (
// ErrMigrationMismatch is returned when the migrated account does not
// match the original account.
ErrMigrationMismatch = fmt.Errorf("migrated account does not match " +
"original account")
)
const migrationProgressLogInterval = 100
// MigrateAccountStoreToSQL runs the migration of all accounts and indices from
// the KV database to the SQL database. The migration is done in a single
// transaction to ensure that all accounts are migrated or none at all.
func MigrateAccountStoreToSQL(ctx context.Context, kvStore kvdb.Backend,
tx *sqlcmig6.Queries) error {
log.Infof("Starting migration of the KV accounts store to SQL")
err := migrateAccountsToSQL(ctx, kvStore, tx)
if err != nil {
return fmt.Errorf("unsuccessful migration of accounts to "+
"SQL: %w", err)
}
err = migrateAccountsIndicesToSQL(ctx, kvStore, tx)
if err != nil {
return fmt.Errorf("unsuccessful migration of account indices "+
"to SQL: %w", err)
}
return nil
}
// migrateAccountsToSQL runs the migration of all accounts from the KV database
// to the SQL database. The migration is done in a single transaction to ensure
// that all accounts are migrated or none at all.
func migrateAccountsToSQL(ctx context.Context, kvStore kvdb.Backend,
tx *sqlcmig6.Queries) error {
log.Infof("Starting migration of accounts from KV to SQL")
kvAccounts, err := getBBoltAccounts(kvStore)
if err != nil {
return err
}
log.Infof("Collected %d accounts for KV to SQL migration",
len(kvAccounts))
for i, kvAccount := range kvAccounts {
migratedAccountID, err := migrateSingleAccountToSQL(
ctx, tx, kvAccount,
)
if err != nil {
return fmt.Errorf("unable to migrate account(%v): %w",
kvAccount.ID, err)
}
migratedAccount, err := getAndMarshalMig6Account(
ctx, tx, migratedAccountID,
)
if err != nil {
return fmt.Errorf("unable to fetch migrated "+
"account(%v): %w", kvAccount.ID, err)
}
overrideAccountTimeZone(kvAccount)
overrideAccountTimeZone(migratedAccount)
if !reflect.DeepEqual(kvAccount, migratedAccount) {
diff := difflib.UnifiedDiff{
A: difflib.SplitLines(
spew.Sdump(kvAccount),
),
B: difflib.SplitLines(
spew.Sdump(migratedAccount),
),
FromFile: "Expected",
FromDate: "",
ToFile: "Actual",
ToDate: "",
Context: 3,
}
diffText, _ := difflib.GetUnifiedDiffString(diff)
return fmt.Errorf("%w: %v.\n%v", ErrMigrationMismatch,
kvAccount.ID, diffText)
}
migratedCount := i + 1
if migratedCount%migrationProgressLogInterval == 0 {
log.Infof("Migrated %d/%d accounts from KV to SQL",
migratedCount, len(kvAccounts))
}
}
log.Infof("All accounts migrated from KV to SQL. Total number of "+
"accounts migrated: %d", len(kvAccounts))
return nil
}
// getBBoltAccounts is a helper function that fetches all accounts from the
// Bbolt store, by iterating directly over the buckets, without needing to
// use any public functions of the BoltStore struct.
func getBBoltAccounts(db kvdb.Backend) ([]*OffChainBalanceAccount, error) {
var accounts []*OffChainBalanceAccount
err := db.View(func(tx kvdb.RTx) error {
// This function will be called in the ForEach and receive
// the key and value of each account in the DB. The key, which
// is also the ID is not used because it is also marshaled into
// the value.
readFn := func(k, v []byte) error {
// Skip the two special purpose keys.
if bytes.Equal(k, lastAddIndexKey) ||
bytes.Equal(k, lastSettleIndexKey) {
return nil
}
// Also skip the kvdb deprecation marker key. We
// still want to allow rerunning the kvdb -> SQL
// migration after the SQL database has been deleted or
// downgraded, even though normal bbolt startup should
// reject the tombstoned kvdb files.
if tombstone.IsMigrationTombstoneKey(k) {
return nil
}
// There should be no sub-buckets.
if v == nil {
return fmt.Errorf("invalid bucket structure")
}
account, err := deserializeAccount(v)
if err != nil {
return err
}
accounts = append(accounts, account)
return nil
}
// We know the bucket should exist since it's created when
// the account storage is initialized.
return tx.ReadBucket(accountBucketName).ForEach(readFn)
}, func() {
accounts = nil
})
if err != nil {
return nil, err
}
return accounts, nil
}
// getAndMarshalMig6Account retrieves the account with the given ID. If the
// account cannot be found, then ErrAccNotFound is returned.
func getAndMarshalMig6Account(ctx context.Context, db *sqlcmig6.Queries,
id int64) (*OffChainBalanceAccount, error) {
dbAcct, err := db.GetAccount(ctx, id)
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrAccNotFound
} else if err != nil {
return nil, err
}
return marshalDBMig6Account(ctx, db, dbAcct)
}
// marshalDBMig6Account marshals a sqlcmig6.Account to a OffChainBalanceAccount.
func marshalDBMig6Account(ctx context.Context, db *sqlcmig6.Queries,
dbAcct sqlcmig6.Account) (*OffChainBalanceAccount, error) {
alias, err := AccountIDFromInt64(dbAcct.Alias)
if err != nil {
return nil, err
}
account := &OffChainBalanceAccount{
ID: alias,
Type: AccountType(dbAcct.Type),
InitialBalance: lnwire.MilliSatoshi(dbAcct.InitialBalanceMsat),
CurrentBalance: dbAcct.CurrentBalanceMsat,
LastUpdate: dbAcct.LastUpdated.UTC(),
ExpirationDate: dbAcct.Expiration.UTC(),
Invoices: make(AccountInvoices),
Payments: make(AccountPayments),
Label: dbAcct.Label.String,
}
invoices, err := db.ListAccountInvoices(ctx, dbAcct.ID)
if err != nil {
return nil, err
}
for _, invoice := range invoices {
var hash lntypes.Hash
copy(hash[:], invoice.Hash)
account.Invoices[hash] = struct{}{}
}
payments, err := db.ListAccountPayments(ctx, dbAcct.ID)
if err != nil {
return nil, err
}
for _, payment := range payments {
var hash lntypes.Hash
copy(hash[:], payment.Hash)
account.Payments[hash] = &PaymentEntry{
Status: lnrpc.Payment_PaymentStatus(payment.Status),
FullAmount: lnwire.MilliSatoshi(payment.FullAmountMsat),
}
}
return account, nil
}
// migrateSingleAccountToSQL runs the migration for a single account from the
// KV database to the SQL database.
func migrateSingleAccountToSQL(ctx context.Context,
tx *sqlcmig6.Queries, account *OffChainBalanceAccount) (int64, error) {
accountAlias, err := account.ID.ToInt64()
if err != nil {
return 0, err
}
insertAccountParams := sqlcmig6.InsertAccountParams{
Type: int16(account.Type),
InitialBalanceMsat: int64(account.InitialBalance),
CurrentBalanceMsat: account.CurrentBalance,
LastUpdated: account.LastUpdate.UTC(),
Alias: accountAlias,
Expiration: account.ExpirationDate.UTC(),
Label: sql.NullString{
String: account.Label,
Valid: len(account.Label) > 0,
},
}
sqlId, err := tx.InsertAccount(ctx, insertAccountParams)
if err != nil {
return 0, err
}
for hash := range account.Invoices {
addInvoiceParams := sqlcmig6.AddAccountInvoiceParams{
AccountID: sqlId,
Hash: hash[:],
}
err = tx.AddAccountInvoice(ctx, addInvoiceParams)
if err != nil {
return sqlId, err
}
}
for hash, paymentEntry := range account.Payments {
upsertPaymentParams := sqlcmig6.UpsertAccountPaymentParams{
AccountID: sqlId,
Hash: hash[:],
Status: int16(paymentEntry.Status),
FullAmountMsat: int64(paymentEntry.FullAmount),
}
err = tx.UpsertAccountPayment(ctx, upsertPaymentParams)
if err != nil {
return sqlId, err
}
}
return sqlId, nil
}
// migrateAccountsIndicesToSQL runs the migration for the account indices from
// the KV database to the SQL database.
func migrateAccountsIndicesToSQL(ctx context.Context, kvStore kvdb.Backend,
tx *sqlcmig6.Queries) error {
log.Infof("Starting migration of accounts indices from KV to SQL")
addIndex, settleIndex, err := getBBoltIndices(kvStore)
if errors.Is(err, ErrNoInvoiceIndexKnown) {
log.Infof("No indices found in KV store, skipping migration")
return nil
} else if err != nil {
return err
}
if addIndex > math.MaxInt64 {
return fmt.Errorf("%s:%v is above max int64 value",
addIndexName, addIndex)
}
if settleIndex > math.MaxInt64 {
return fmt.Errorf("%s:%v is above max int64 value",
settleIndexName, settleIndex)
}
setAddIndexParams := sqlcmig6.SetAccountIndexParams{
Name: addIndexName,
Value: int64(addIndex),
}
err = tx.SetAccountIndex(ctx, setAddIndexParams)
if err != nil {
return err
}
setSettleIndexParams := sqlcmig6.SetAccountIndexParams{
Name: settleIndexName,
Value: int64(settleIndex),
}
err = tx.SetAccountIndex(ctx, setSettleIndexParams)
if err != nil {
return err
}
log.Infof("Successfully migratated accounts indices from KV to SQL")
return nil
}
// getBBoltIndices is a helper function that fetches the índices from the
// Bbolt store, by iterating directly over the buckets, without needing to
// use any public functions of the BoltStore struct.
func getBBoltIndices(db kvdb.Backend) (uint64, uint64, error) {
var (
addValue, settleValue []byte
)
err := db.View(func(tx kvdb.RTx) error {
bucket := tx.ReadBucket(accountBucketName)
if bucket == nil {
return ErrAccountBucketNotFound
}
av := bucket.Get(lastAddIndexKey)
if len(av) == 0 {
return ErrNoInvoiceIndexKnown
}
sv := bucket.Get(lastSettleIndexKey)
if len(sv) == 0 {
return ErrNoInvoiceIndexKnown
}
// Copy values since bbolt's Get returns slices into the
// mmap'd file, only valid within the transaction.
addValue = make([]byte, len(av))
copy(addValue, av)
settleValue = make([]byte, len(sv))
copy(settleValue, sv)
return nil
}, func() {
addValue, settleValue = nil, nil
})
if err != nil {
return 0, 0, err
}
return byteOrder.Uint64(addValue), byteOrder.Uint64(settleValue), nil
}
// overrideAccountTimeZone overrides the time zone of the account to the local
// time zone and chops off the nanosecond part for comparison. This is needed
// because KV database stores times as-is which as an unwanted side effect would
// fail migration due to time comparison expecting both the original and
// migrated accounts to be in the same local time zone and in microsecond
// precision. Note that PostgresSQL stores times in microsecond precision while
// SQLite can store times in nanosecond precision if using TEXT storage class.
func overrideAccountTimeZone(account *OffChainBalanceAccount) {
fixTime := func(t time.Time) time.Time {
return t.In(time.Local).Truncate(time.Microsecond)
}
if !account.ExpirationDate.IsZero() {
account.ExpirationDate = fixTime(account.ExpirationDate)
}
if !account.LastUpdate.IsZero() {
account.LastUpdate = fixTime(account.LastUpdate)
}
}