From 5f4e52bd884cd5f6e356e9ce0585c77cd64f591d Mon Sep 17 00:00:00 2001 From: Roland <33993199+rolznz@users.noreply.github.com> Date: Sat, 8 Aug 2026 13:38:28 +0700 Subject: [PATCH] fix: publish transaction events only after the database transaction commits (#2520) * fix: publish transaction events only after the database transaction commits markTransactionSettled and markPaymentFailed published nwc_payment_sent / nwc_payment_received / nwc_payment_failed (and checkBudgetUsage published nwc_budget_warning) while still inside the caller's database transaction, so connected apps and the Alby API could be notified of a payment whose row was never committed, and subscribers reading the database in response to an event could race with the commit. Every function that writes transaction state now owns its own database transaction and publishes its events only after the commit succeeds: - markTransactionSettled and markPaymentFailed open their own transaction; callers no longer wrap them in db.Transaction - new createSettledTransactionFromNotification inserts transactions reported by LNClient notifications for payments the hub has no record of (external payments, received keysends) directly in their settled state, removing the transient PENDING row and the zombie row left behind on duplicate events - markPaymentFailed now refuses to mark a settled transaction as failed, replacing CancelHoldInvoice's in-transaction ACCEPTED re-check and also protecting the SendPaymentSync error path from a racing settle - checkBudgetUsage returns the budget warning event instead of publishing it Closes #2506 Co-Authored-By: Claude Fable 5 * fix: serialize payment failure with settlement and propagate lock errors Address review findings on the previous commit: - markPaymentFailed now takes the same payment-hash row lock as settlement (postgres), so the settled-state guard cannot be bypassed by a concurrent settle between the state check and the update; it also returns not-found instead of publishing an event when the transaction row no longer exists, and reports whether this call transitioned the row so CancelHoldInvoice only publishes nwc_hold_invoice_canceled when it performed the cancellation - findSettledTransaction propagates errors from the lock query and the settled-transaction lookup instead of treating a failed lookup as "no settled transaction exists", which could defeat the dedup guard - TestMarkSettled_Twice no longer shares one transaction struct between concurrent goroutines and collects errors instead of asserting inside them Co-Authored-By: Claude Fable 5 * fix: mark failed keysend payments via markPaymentFailed The SendKeysend failure path updated the transaction directly, which never zeroed the fee reserve, recorded no failure reason, published no nwc_payment_failed event, and had no guard against overwriting a concurrently settled payment. Route it through markPaymentFailed like SendPaymentSync, and allow MockLn keysends to fail so the path is testable. Co-Authored-By: Claude Fable 5 --------- Co-authored-by: Claude Fable 5 --- tests/mock_ln_client.go | 9 + transactions/app_payments_test.go | 43 ++ transactions/keysend_test.go | 29 ++ transactions/payments_test.go | 84 ++-- transactions/transactions_service.go | 617 ++++++++++++++++----------- 5 files changed, 501 insertions(+), 281 deletions(-) diff --git a/tests/mock_ln_client.go b/tests/mock_ln_client.go index 35c775d2..e935b385 100644 --- a/tests/mock_ln_client.go +++ b/tests/mock_ln_client.go @@ -81,6 +81,8 @@ type MockLn struct { MakeInvoiceErrors []error PayInvoiceResponses []*lnclient.PayInvoiceResponse PayInvoiceErrors []error + PayKeysendResponses []*lnclient.PayKeysendResponse + PayKeysendErrors []error PaymentDelay *time.Duration Pubkey string MockTransaction *lnclient.Transaction @@ -110,6 +112,13 @@ func (mln *MockLn) SendPaymentSync(payReq string, amountMsat *uint64) (*lnclient } func (mln *MockLn) SendKeysend(amountMsat uint64, destination string, custom_records []lnclient.TLVRecord, preimage string) (*lnclient.PayKeysendResponse, error) { + if len(mln.PayKeysendResponses) > 0 { + response := mln.PayKeysendResponses[0] + err := mln.PayKeysendErrors[0] + mln.PayKeysendResponses = mln.PayKeysendResponses[1:] + mln.PayKeysendErrors = mln.PayKeysendErrors[1:] + return response, err + } return &lnclient.PayKeysendResponse{ FeeMsat: 1, }, nil diff --git a/transactions/app_payments_test.go b/transactions/app_payments_test.go index 552ef6dc..b989407a 100644 --- a/transactions/app_payments_test.go +++ b/transactions/app_payments_test.go @@ -62,6 +62,49 @@ func TestSendPaymentSync_App_WithPermission(t *testing.T) { assert.Equal(t, dbRequestEvent.ID, *transaction.RequestEventId) } +func TestMarkSettled_App_BudgetWarning(t *testing.T) { + svc, err := tests.CreateTestService(t) + require.NoError(t, err) + defer svc.Remove() + + app, _, err := tests.CreateApp(svc) + assert.NoError(t, err) + + appPermission := &db.AppPermission{ + AppId: app.ID, + App: *app, + Scope: constants.PAY_INVOICE_SCOPE, + MaxAmountSat: 100, + } + err = svc.DB.Create(appPermission).Error + assert.NoError(t, err) + + // settling this payment pushes the app over 80% of its 100 sat budget + dbTransaction := db.Transaction{ + AppId: &app.ID, + State: constants.TRANSACTION_STATE_PENDING, + Type: constants.TRANSACTION_TYPE_OUTGOING, + PaymentHash: tests.MockLNClientTransaction.PaymentHash, + AmountMsat: 90000, + } + svc.DB.Create(&dbTransaction) + + mockEventConsumer := tests.NewMockEventConsumer() + svc.EventPublisher.RegisterSubscriber(mockEventConsumer) + transactionsService := NewTransactionsService(svc.DB, svc.EventPublisher) + _, err = transactionsService.markTransactionSettled(&dbTransaction, "test", 0, false) + + assert.NoError(t, err) + consumedEvents := mockEventConsumer.GetConsumedEvents() + assert.Equal(t, 2, len(consumedEvents)) + eventNames := []string{} + for _, consumedEvent := range consumedEvents { + eventNames = append(eventNames, consumedEvent.Event) + } + assert.Contains(t, eventNames, "nwc_payment_sent") + assert.Contains(t, eventNames, "nwc_budget_warning") +} + func TestSendPaymentSync_App_BudgetExceeded(t *testing.T) { svc, err := tests.CreateTestService(t) require.NoError(t, err) diff --git a/transactions/keysend_test.go b/transactions/keysend_test.go index 42f21d86..ba50db49 100644 --- a/transactions/keysend_test.go +++ b/transactions/keysend_test.go @@ -4,6 +4,7 @@ import ( "context" "encoding/hex" "encoding/json" + "errors" "strconv" "testing" @@ -47,6 +48,34 @@ func TestSendKeysend(t *testing.T) { settledTransaction := mockEventConsumer.GetConsumedEvents()[0].Properties.(*db.Transaction) assert.Equal(t, transaction, settledTransaction) } +func TestSendKeysend_FailedRemovesFeeReserve(t *testing.T) { + svc, err := tests.CreateTestService(t) + require.NoError(t, err) + defer svc.Remove() + + svc.LNClient.(*tests.MockLn).PayKeysendErrors = append(svc.LNClient.(*tests.MockLn).PayKeysendErrors, errors.New("Some error")) + svc.LNClient.(*tests.MockLn).PayKeysendResponses = append(svc.LNClient.(*tests.MockLn).PayKeysendResponses, nil) + + mockEventConsumer := tests.NewMockEventConsumer() + svc.EventPublisher.RegisterSubscriber(mockEventConsumer) + + transactionsService := NewTransactionsService(svc.DB, svc.EventPublisher) + transaction, err := transactionsService.SendKeysend(uint64(1000), "fake destination", nil, "", svc.LNClient, nil, nil) + + assert.Error(t, err) + assert.Nil(t, transaction) + + failedTransaction := db.Transaction{} + require.NoError(t, svc.DB.Where("type = ? AND state = ?", constants.TRANSACTION_TYPE_OUTGOING, constants.TRANSACTION_STATE_FAILED).First(&failedTransaction).Error) + + assert.Equal(t, uint64(1000), failedTransaction.AmountMsat) + assert.Zero(t, failedTransaction.FeeReserveMsat) + assert.Equal(t, "Some error", failedTransaction.FailureReason) + + assert.Equal(t, 1, len(mockEventConsumer.GetConsumedEvents())) + assert.Equal(t, "nwc_payment_failed", mockEventConsumer.GetConsumedEvents()[0].Event) +} + func TestSendKeysend_CustomPreimage(t *testing.T) { svc, err := tests.CreateTestService(t) require.NoError(t, err) diff --git a/transactions/payments_test.go b/transactions/payments_test.go index ef2a675f..ef8a5037 100644 --- a/transactions/payments_test.go +++ b/transactions/payments_test.go @@ -12,7 +12,6 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "gorm.io/gorm" "github.com/getAlby/hub/constants" "github.com/getAlby/hub/db" @@ -181,10 +180,7 @@ func TestMarkSettled_Sent(t *testing.T) { mockEventConsumer := tests.NewMockEventConsumer() svc.EventPublisher.RegisterSubscriber(mockEventConsumer) transactionsService := NewTransactionsService(svc.DB, svc.EventPublisher) - err = svc.DB.Transaction(func(tx *gorm.DB) error { - _, err = transactionsService.markTransactionSettled(tx, &dbTransaction, "test", 0, false) - return err - }) + _, err = transactionsService.markTransactionSettled(&dbTransaction, "test", 0, false) assert.NoError(t, err) assert.Equal(t, constants.TRANSACTION_STATE_SETTLED, dbTransaction.State) @@ -212,29 +208,36 @@ func TestMarkSettled_Twice(t *testing.T) { transactionsService := NewTransactionsService(svc.DB, svc.EventPublisher) var wg sync.WaitGroup n := 10 + markErrors := make([]error, n) wg.Add(n) - for range n { + for i := range n { go func() { defer wg.Done() - err = svc.DB.Transaction(func(tx *gorm.DB) error { - time.Sleep(time.Duration(n) * 10 * time.Millisecond) - _, err = transactionsService.markTransactionSettled(tx, &dbTransaction, "test", 0, false) - time.Sleep(time.Duration(n) * 10 * time.Millisecond) - return err - }) - require.NoError(t, err) + // load an independent copy so goroutines don't share the struct + var transactionCopy db.Transaction + if err := svc.DB.First(&transactionCopy, dbTransaction.ID).Error; err != nil { + markErrors[i] = err + return + } + _, markErrors[i] = transactionsService.markTransactionSettled(&transactionCopy, "test", 0, false) }() } wg.Wait() + for _, markError := range markErrors { + assert.NoError(t, markError) + } + // ensure we only mark transaction settled once and only fire // settled notifications once - assert.NoError(t, err) - assert.Equal(t, constants.TRANSACTION_STATE_SETTLED, dbTransaction.State) + var reloadedTransaction db.Transaction + require.NoError(t, svc.DB.First(&reloadedTransaction, dbTransaction.ID).Error) + assert.Equal(t, constants.TRANSACTION_STATE_SETTLED, reloadedTransaction.State) assert.Equal(t, 1, len(mockEventConsumer.GetConsumedEvents())) assert.Equal(t, "nwc_payment_sent", mockEventConsumer.GetConsumedEvents()[0].Event) settledTransaction := mockEventConsumer.GetConsumedEvents()[0].Properties.(*db.Transaction) - assert.Equal(t, &dbTransaction, settledTransaction) + assert.Equal(t, constants.TRANSACTION_STATE_SETTLED, settledTransaction.State) + assert.Equal(t, dbTransaction.PaymentHash, settledTransaction.PaymentHash) } func TestMarkSettled_Received(t *testing.T) { @@ -253,10 +256,7 @@ func TestMarkSettled_Received(t *testing.T) { mockEventConsumer := tests.NewMockEventConsumer() svc.EventPublisher.RegisterSubscriber(mockEventConsumer) transactionsService := NewTransactionsService(svc.DB, svc.EventPublisher) - err = svc.DB.Transaction(func(tx *gorm.DB) error { - _, err = transactionsService.markTransactionSettled(tx, &dbTransaction, "test", 0, false) - return err - }) + _, err = transactionsService.markTransactionSettled(&dbTransaction, "test", 0, false) assert.NoError(t, err) assert.Equal(t, constants.TRANSACTION_STATE_SETTLED, dbTransaction.State) @@ -284,10 +284,7 @@ func TestDoNotMarkSettledTwice(t *testing.T) { mockEventConsumer := tests.NewMockEventConsumer() svc.EventPublisher.RegisterSubscriber(mockEventConsumer) transactionsService := NewTransactionsService(svc.DB, svc.EventPublisher) - err = svc.DB.Transaction(func(tx *gorm.DB) error { - _, err = transactionsService.markTransactionSettled(tx, &dbTransaction, "test", 0, false) - return err - }) + _, err = transactionsService.markTransactionSettled(&dbTransaction, "test", 0, false) assert.NoError(t, err) assert.Zero(t, len(mockEventConsumer.GetConsumedEvents())) @@ -309,11 +306,10 @@ func TestMarkFailed(t *testing.T) { mockEventConsumer := tests.NewMockEventConsumer() svc.EventPublisher.RegisterSubscriber(mockEventConsumer) transactionsService := NewTransactionsService(svc.DB, svc.EventPublisher) - err = svc.DB.Transaction(func(tx *gorm.DB) error { - return transactionsService.markPaymentFailed(tx, &dbTransaction, "some routing error") - }) + markedFailed, err := transactionsService.markPaymentFailed(&dbTransaction, "some routing error") assert.NoError(t, err) + assert.True(t, markedFailed) assert.Equal(t, constants.TRANSACTION_STATE_FAILED, dbTransaction.State) assert.Equal(t, 1, len(mockEventConsumer.GetConsumedEvents())) assert.Equal(t, "nwc_payment_failed", mockEventConsumer.GetConsumedEvents()[0].Event) @@ -340,15 +336,43 @@ func TestDoNotMarkFailedTwice(t *testing.T) { mockEventConsumer := tests.NewMockEventConsumer() svc.EventPublisher.RegisterSubscriber(mockEventConsumer) transactionsService := NewTransactionsService(svc.DB, svc.EventPublisher) - err = svc.DB.Transaction(func(tx *gorm.DB) error { - return transactionsService.markPaymentFailed(tx, &dbTransaction, "some routing error") - }) + markedFailed, err := transactionsService.markPaymentFailed(&dbTransaction, "some routing error") assert.NoError(t, err) + assert.False(t, markedFailed) assert.Equal(t, updatedAt, dbTransaction.UpdatedAt) assert.Zero(t, len(mockEventConsumer.GetConsumedEvents())) } +func TestDoNotMarkSettledPaymentFailed(t *testing.T) { + svc, err := tests.CreateTestService(t) + require.NoError(t, err) + defer svc.Remove() + + settledAt := time.Now() + dbTransaction := db.Transaction{ + State: constants.TRANSACTION_STATE_SETTLED, + Type: constants.TRANSACTION_TYPE_OUTGOING, + PaymentHash: tests.MockLNClientTransaction.PaymentHash, + AmountMsat: 123000, + SettledAt: &settledAt, + } + svc.DB.Create(&dbTransaction) + + mockEventConsumer := tests.NewMockEventConsumer() + svc.EventPublisher.RegisterSubscriber(mockEventConsumer) + transactionsService := NewTransactionsService(svc.DB, svc.EventPublisher) + markedFailed, err := transactionsService.markPaymentFailed(&dbTransaction, "some routing error") + + assert.Error(t, err) + assert.False(t, markedFailed) + + var reloadedTransaction db.Transaction + require.NoError(t, svc.DB.First(&reloadedTransaction, dbTransaction.ID).Error) + assert.Equal(t, constants.TRANSACTION_STATE_SETTLED, reloadedTransaction.State) + assert.Zero(t, len(mockEventConsumer.GetConsumedEvents())) +} + func TestSendPaymentSync_FailedRemovesFeeReserve(t *testing.T) { svc, err := tests.CreateTestService(t) require.NoError(t, err) diff --git a/transactions/transactions_service.go b/transactions/transactions_service.go index 407d07d7..d0497676 100644 --- a/transactions/transactions_service.go +++ b/transactions/transactions_service.go @@ -444,19 +444,17 @@ func (svc *transactionsService) SendPaymentSync(payReq string, amountMsat *uint6 "bolt11": payReq, }).WithError(err).Error("Failed to send payment") - svc.db.Transaction(func(tx *gorm.DB) error { - return svc.markPaymentFailed(tx, &dbTransaction, err.Error()) - }) + if _, markFailedErr := svc.markPaymentFailed(&dbTransaction, err.Error()); markFailedErr != nil { + logger.Logger.WithFields(logrus.Fields{ + "bolt11": payReq, + }).WithError(markFailedErr).Error("Failed to mark payment as failed") + } return nil, err } // the payment definitely succeeded - var settledTransaction *db.Transaction - err = svc.db.Transaction(func(tx *gorm.DB) error { - settledTransaction, err = svc.markTransactionSettled(tx, &dbTransaction, response.Preimage, response.FeeMsat, selfPayment) - return err - }) + settledTransaction, err := svc.markTransactionSettled(&dbTransaction, response.Preimage, response.FeeMsat, selfPayment) if err != nil { return nil, err } @@ -579,27 +577,18 @@ func (svc *transactionsService) SendKeysend(amountMsat uint64, destination strin "amount_msat": amountMsat, }).WithError(err).Error("Failed to send payment") - dbErr := svc.db.Model(&dbTransaction).Updates(&db.Transaction{ - PaymentHash: paymentHash, - State: constants.TRANSACTION_STATE_FAILED, - }).Error - if dbErr != nil { + if _, markFailedErr := svc.markPaymentFailed(&dbTransaction, err.Error()); markFailedErr != nil { logger.Logger.WithFields(logrus.Fields{ "destination": destination, "amount_msat": amountMsat, - }).WithError(dbErr).Error("Failed to update DB transaction") + }).WithError(markFailedErr).Error("Failed to mark payment as failed") } return nil, err } // the payment definitely succeeded - var settledTransaction *db.Transaction - err = svc.db.Transaction(func(tx *gorm.DB) error { - settledTransaction, err = svc.markTransactionSettled(tx, &dbTransaction, preimage, payKeysendResponse.FeeMsat, selfPayment) - return err - }) - + settledTransaction, err := svc.markTransactionSettled(&dbTransaction, preimage, payKeysendResponse.FeeMsat, selfPayment) if err != nil { return nil, err } @@ -795,11 +784,7 @@ func (svc *transactionsService) checkUnsettledTransaction(ctx context.Context, t } // update transaction state if lnClientTransaction.SettledAt != nil { - err = svc.db.Transaction(func(tx *gorm.DB) error { - _, err = svc.markTransactionSettled(tx, transaction, lnClientTransaction.Preimage, uint64(lnClientTransaction.FeesPaidMsat), false) - return err - }) - + _, err = svc.markTransactionSettled(transaction, lnClientTransaction.Preimage, uint64(lnClientTransaction.FeesPaidMsat), false) if err != nil { logger.Logger.WithError(err).Error("Failed to mark payment sent when checking unsettled transaction") } @@ -816,73 +801,72 @@ func (svc *transactionsService) ConsumeEvent(ctx context.Context, event *events. } var dbTransaction db.Transaction - err := svc.db.Transaction(func(tx *gorm.DB) error { - - result := tx.Limit(1).Find(&dbTransaction, &db.Transaction{ - Type: constants.TRANSACTION_TYPE_INCOMING, - PaymentHash: lnClientTransaction.PaymentHash, - }) - - if result.RowsAffected == 0 { - var appId *uint - description := lnClientTransaction.Description - var metadataBytes []byte - var boostagramBytes []byte - if lnClientTransaction.Metadata != nil { - var err error - metadataBytes, err = json.Marshal(lnClientTransaction.Metadata) - if err != nil { - logger.Logger.WithError(err).Error("Failed to serialize transaction metadata") - return err - } - - var customRecords []lnclient.TLVRecord - customRecords, _ = lnClientTransaction.Metadata["tlv_records"].([]lnclient.TLVRecord) - boostagramBytes = svc.getBoostagramBytesFromCustomRecords(customRecords) - extractedDescription := svc.getDescriptionFromCustomRecords(customRecords) - if extractedDescription != "" { - description = extractedDescription - } - // find app by custom key/value records - appId = svc.getAppIdFromCustomRecords(customRecords, tx) - } - var expiresAt *time.Time - if lnClientTransaction.ExpiresAt != nil { - expiresAtValue := time.Unix(*lnClientTransaction.ExpiresAt, 0) - expiresAt = &expiresAtValue - } - dbTransaction = db.Transaction{ - Type: constants.TRANSACTION_TYPE_INCOMING, - AmountMsat: uint64(lnClientTransaction.AmountMsat), - PaymentRequest: lnClientTransaction.Invoice, - PaymentHash: lnClientTransaction.PaymentHash, - Description: description, - DescriptionHash: lnClientTransaction.DescriptionHash, - ExpiresAt: expiresAt, - Metadata: datatypes.JSON(metadataBytes), - Boostagram: datatypes.JSON(boostagramBytes), - AppId: appId, - } - err := tx.Create(&dbTransaction).Error - if err != nil { - logger.Logger.WithFields(logrus.Fields{ - "payment_hash": lnClientTransaction.PaymentHash, - }).WithError(err).Error("Failed to create transaction") - return err - } - } - - _, err := svc.markTransactionSettled(tx, &dbTransaction, lnClientTransaction.Preimage, uint64(lnClientTransaction.FeesPaidMsat), false) - return err + result := svc.db.Limit(1).Find(&dbTransaction, &db.Transaction{ + Type: constants.TRANSACTION_TYPE_INCOMING, + PaymentHash: lnClientTransaction.PaymentHash, }) - if err != nil { + if result.Error != nil { logger.Logger.WithFields(logrus.Fields{ "payment_hash": lnClientTransaction.PaymentHash, - }).WithError(err).Error("Failed to execute DB transaction") + }).WithError(result.Error).Error("Failed to find transaction") return } + if result.RowsAffected == 0 { + var appId *uint + description := lnClientTransaction.Description + var metadataBytes []byte + var boostagramBytes []byte + if lnClientTransaction.Metadata != nil { + var err error + metadataBytes, err = json.Marshal(lnClientTransaction.Metadata) + if err != nil { + logger.Logger.WithError(err).Error("Failed to serialize transaction metadata") + return + } + + var customRecords []lnclient.TLVRecord + customRecords, _ = lnClientTransaction.Metadata["tlv_records"].([]lnclient.TLVRecord) + boostagramBytes = svc.getBoostagramBytesFromCustomRecords(customRecords) + extractedDescription := svc.getDescriptionFromCustomRecords(customRecords) + if extractedDescription != "" { + description = extractedDescription + } + // find app by custom key/value records + appId = svc.getAppIdFromCustomRecords(customRecords, svc.db) + } + var expiresAt *time.Time + if lnClientTransaction.ExpiresAt != nil { + expiresAtValue := time.Unix(*lnClientTransaction.ExpiresAt, 0) + expiresAt = &expiresAtValue + } + dbTransaction = db.Transaction{ + Type: constants.TRANSACTION_TYPE_INCOMING, + AmountMsat: uint64(lnClientTransaction.AmountMsat), + PaymentRequest: lnClientTransaction.Invoice, + PaymentHash: lnClientTransaction.PaymentHash, + Description: description, + DescriptionHash: lnClientTransaction.DescriptionHash, + ExpiresAt: expiresAt, + Metadata: datatypes.JSON(metadataBytes), + Boostagram: datatypes.JSON(boostagramBytes), + AppId: appId, + } + if _, err := svc.createSettledTransactionFromNotification(&dbTransaction, lnClientTransaction.Preimage, uint64(lnClientTransaction.FeesPaidMsat), false); err != nil { + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": lnClientTransaction.PaymentHash, + }).WithError(err).Error("Failed to create settled transaction") + } + return + } + + if _, err := svc.markTransactionSettled(&dbTransaction, lnClientTransaction.Preimage, uint64(lnClientTransaction.FeesPaidMsat), false); err != nil { + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": lnClientTransaction.PaymentHash, + }).WithError(err).Error("Failed to mark transaction as settled") + } + case "nwc_lnclient_hold_invoice_accepted": lnClientTransaction, ok := event.Properties.(*lnclient.Transaction) if !ok { @@ -903,78 +887,79 @@ func (svc *transactionsService) ConsumeEvent(ctx context.Context, event *events. } var dbTransaction db.Transaction - err := svc.db.Transaction(func(tx *gorm.DB) error { - // first lookup by pending - result := tx.Limit(1).Find(&dbTransaction, &db.Transaction{ + // first lookup by pending + result := svc.db.Limit(1).Find(&dbTransaction, &db.Transaction{ + Type: constants.TRANSACTION_TYPE_OUTGOING, + State: constants.TRANSACTION_STATE_PENDING, + PaymentHash: lnClientTransaction.PaymentHash, + }) + + if result.Error != nil { + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": lnClientTransaction.PaymentHash, + }).WithError(result.Error).Error("Failed to find transaction") + return + } + + if result.RowsAffected == 0 { + // if no pending payment was found, lookup by failed, latest updated first + result := svc.db.Limit(1).Order("updated_at DESC").Find(&dbTransaction, &db.Transaction{ Type: constants.TRANSACTION_TYPE_OUTGOING, - State: constants.TRANSACTION_STATE_PENDING, + State: constants.TRANSACTION_STATE_FAILED, PaymentHash: lnClientTransaction.PaymentHash, }) if result.Error != nil { - return result.Error + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": lnClientTransaction.PaymentHash, + }).WithError(result.Error).Error("Failed to find transaction") + return } if result.RowsAffected == 0 { - // if no pending payment was found, lookup by failed, latest updated first - result := tx.Limit(1).Order("updated_at DESC").Find(&dbTransaction, &db.Transaction{ + result := svc.db.Limit(1).Find(&dbTransaction, &db.Transaction{ Type: constants.TRANSACTION_TYPE_OUTGOING, - State: constants.TRANSACTION_STATE_FAILED, PaymentHash: lnClientTransaction.PaymentHash, }) if result.Error != nil { - return result.Error + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": lnClientTransaction.PaymentHash, + }).WithError(result.Error).Error("Failed to find transaction") + return } if result.RowsAffected == 0 { - result := tx.Limit(1).Find(&dbTransaction, &db.Transaction{ - Type: constants.TRANSACTION_TYPE_OUTGOING, - PaymentHash: lnClientTransaction.PaymentHash, - }) - - if result.Error != nil { - return result.Error + dbTransaction = db.Transaction{ + Type: constants.TRANSACTION_TYPE_OUTGOING, + AmountMsat: uint64(lnClientTransaction.AmountMsat), + FeeReserveMsat: 0, + PaymentRequest: lnClientTransaction.Invoice, + PaymentHash: lnClientTransaction.PaymentHash, + Description: lnClientTransaction.Description, + DescriptionHash: lnClientTransaction.DescriptionHash, } - if result.RowsAffected == 0 { - dbTransaction = db.Transaction{ - Type: constants.TRANSACTION_TYPE_OUTGOING, - State: constants.TRANSACTION_STATE_PENDING, - AmountMsat: uint64(lnClientTransaction.AmountMsat), - FeeReserveMsat: 0, - PaymentRequest: lnClientTransaction.Invoice, - PaymentHash: lnClientTransaction.PaymentHash, - Description: lnClientTransaction.Description, - DescriptionHash: lnClientTransaction.DescriptionHash, - } - - if lnClientTransaction.ExpiresAt != nil { - expiresAtValue := time.Unix(*lnClientTransaction.ExpiresAt, 0) - dbTransaction.ExpiresAt = &expiresAtValue - } - - err := tx.Create(&dbTransaction).Error - if err != nil { - logger.Logger.WithFields(logrus.Fields{ - "payment_hash": lnClientTransaction.PaymentHash, - }).WithError(err).Error("Failed to create outgoing transaction") - return err - } + if lnClientTransaction.ExpiresAt != nil { + expiresAtValue := time.Unix(*lnClientTransaction.ExpiresAt, 0) + dbTransaction.ExpiresAt = &expiresAtValue } + + if _, err := svc.createSettledTransactionFromNotification(&dbTransaction, lnClientTransaction.Preimage, uint64(lnClientTransaction.FeesPaidMsat), false); err != nil { + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": lnClientTransaction.PaymentHash, + }).WithError(err).Error("Failed to create settled transaction") + } + return } } + } - _, err := svc.markTransactionSettled(tx, &dbTransaction, lnClientTransaction.Preimage, uint64(lnClientTransaction.FeesPaidMsat), false) - return err - }) - - if err != nil { + if _, err := svc.markTransactionSettled(&dbTransaction, lnClientTransaction.Preimage, uint64(lnClientTransaction.FeesPaidMsat), false); err != nil { logger.Logger.WithFields(logrus.Fields{ "payment_hash": lnClientTransaction.PaymentHash, }).WithError(err).Error("Failed to update transaction") - return } case "nwc_lnclient_payment_failed": paymentFailedAsyncProperties, ok := event.Properties.(*lnclient.PaymentFailedEventProperties) @@ -997,9 +982,11 @@ func (svc *transactionsService) ConsumeEvent(ctx context.Context, event *events. return } - svc.db.Transaction(func(tx *gorm.DB) error { - return svc.markPaymentFailed(tx, &dbTransaction, paymentFailedAsyncProperties.Reason) - }) + if _, err := svc.markPaymentFailed(&dbTransaction, paymentFailedAsyncProperties.Reason); err != nil { + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": lnClientTransaction.PaymentHash, + }).WithError(err).Error("Failed to mark payment as failed") + } } } @@ -1094,11 +1081,7 @@ func (svc *transactionsService) interceptSelfPayment(paymentRequest string, paym return nil, errors.New("preimage is not set on transaction. Self payments not supported") } - err := svc.db.Transaction(func(tx *gorm.DB) error { - _, err := svc.markTransactionSettled(tx, &incomingTransaction, *incomingTransaction.Preimage, uint64(0), true) - return err - }) - + _, err := svc.markTransactionSettled(&incomingTransaction, *incomingTransaction.Preimage, uint64(0), true) if err != nil { return nil, err } @@ -1356,18 +1339,12 @@ func (svc *transactionsService) SettleHoldInvoice(ctx context.Context, preimage return nil, err } - var settledTransaction *db.Transaction - err = svc.db.Transaction(func(tx *gorm.DB) error { - var err error - settledTransaction, err = svc.markTransactionSettled(tx, &dbTransaction, preimage, 0, dbTransaction.SelfPayment) - return err - }) - + settledTransaction, err := svc.markTransactionSettled(&dbTransaction, preimage, 0, dbTransaction.SelfPayment) if err != nil { logger.Logger.WithFields(logrus.Fields{ "payment_hash": paymentHash, "preimage": preimage, - }).WithError(err).Error("Failed DB transaction while settling hold invoice") + }).WithError(err).Error("Failed to mark hold invoice as settled") return nil, err } @@ -1399,37 +1376,23 @@ func (svc *transactionsService) CancelHoldInvoice(ctx context.Context, paymentHa } } - err := svc.db.Transaction(func(tx *gorm.DB) error { - var dbTransaction db.Transaction - result := tx.Limit(1).Find(&dbTransaction, &db.Transaction{ - Type: constants.TRANSACTION_TYPE_INCOMING, - State: constants.TRANSACTION_STATE_ACCEPTED, - PaymentHash: paymentHash, - }) - - if result.Error != nil { - logger.Logger.WithFields(logrus.Fields{ - "payment_hash": paymentHash, - }).WithError(result.Error).Error("Failed to find accepted hold invoice in DB for cancellation") - return result.Error - } - if result.RowsAffected == 0 { - logger.Logger.WithFields(logrus.Fields{ - "payment_hash": paymentHash, - }).Warn("No accepted hold invoice found in DB to mark as failed due to cancellation") - return NewNotFoundError() - } - - return svc.markPaymentFailed(tx, &dbTransaction, "Hold invoice was cancelled") - }) - + markedFailed, err := svc.markPaymentFailed(&dbTransaction, "Hold invoice was cancelled") if err != nil { logger.Logger.WithFields(logrus.Fields{ "payment_hash": paymentHash, - }).WithError(err).Error("Failed DB transaction while canceling hold invoice") + }).WithError(err).Error("Failed to mark hold invoice as failed due to cancellation") return err } + if !markedFailed { + // a concurrent cancellation already marked the invoice as failed and + // published the canceled event + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": paymentHash, + }).Info("Hold invoice was already marked as failed") + return nil + } + logger.Logger.WithFields(logrus.Fields{ "payment_hash": paymentHash, }).Info("Marked hold invoice as failed in DB due to cancellation") @@ -1501,61 +1464,171 @@ func (svc *transactionsService) SetTransactionUserLabels(ctx context.Context, id return svc.SetTransactionMetadata(ctx, id, metadata) } -func (svc *transactionsService) markTransactionSettled(tx *gorm.DB, dbTransaction *db.Transaction, preimage string, feeMsat uint64, selfPayment bool) (*db.Transaction, error) { +// markTransactionSettled marks an existing transaction as settled in its own +// database transaction and publishes the corresponding events after it +// commits, so subscribers never observe uncommitted state. +func (svc *transactionsService) markTransactionSettled(dbTransaction *db.Transaction, preimage string, feeMsat uint64, selfPayment bool) (*db.Transaction, error) { if preimage == "" { return nil, errors.New("no preimage in payment") } - if tx.Dialector.Name() == "postgres" { - // lock based on payment hash to ensure we only mark one transaction as settled - // (in sqlite transactions are serializable by default) - transactionsWithPaymentHash := []db.Transaction{} - tx.Where(&db.Transaction{ - PaymentHash: dbTransaction.PaymentHash, - }).Clauses(clause.Locking{Strength: "UPDATE"}).Find(&transactionsWithPaymentHash) + var settledTransaction *db.Transaction + var eventsToPublish []*events.Event + err := svc.db.Transaction(func(tx *gorm.DB) error { + existingSettledTransaction, err := svc.findSettledTransaction(tx, dbTransaction) + if err != nil { + return err + } + if existingSettledTransaction != nil { + logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).Debug("payment already marked as sent") + settledTransaction = existingSettledTransaction + return nil + } + + settledAt := time.Now() + err = tx.Model(dbTransaction).Updates(map[string]interface{}{ + "State": constants.TRANSACTION_STATE_SETTLED, + "Preimage": &preimage, + "FeeMsat": feeMsat, + "FeeReserveMsat": 0, + "SettledAt": &settledAt, + "SelfPayment": selfPayment, + }).Error + if err != nil { + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": dbTransaction.PaymentHash, + }).WithError(err).Error("Failed to update DB transaction") + return err + } + + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": dbTransaction.PaymentHash, + "type": dbTransaction.Type, + }).Info("Marked transaction as settled") + + settledTransaction = dbTransaction + eventsToPublish = svc.afterTransactionSettled(tx, dbTransaction, &settledAt) + return nil + }) + if err != nil { + return nil, err + } + svc.publishEvents(eventsToPublish) + + return settledTransaction, nil +} + +// createSettledTransactionFromNotification inserts a transaction directly in +// its settled state, in its own database transaction, and publishes the +// corresponding events after it commits. It is for the case where the +// LNClient notifies us of a sent or received payment we didn't already know +// about (e.g. if the LNClient is an external node, and the payment was made +// or received outside of Alby Hub, or a received keysend, which has no +// invoice created upfront). +func (svc *transactionsService) createSettledTransactionFromNotification(dbTransaction *db.Transaction, preimage string, feeMsat uint64, selfPayment bool) (*db.Transaction, error) { + if preimage == "" { + return nil, errors.New("no preimage in payment") + } + + var settledTransaction *db.Transaction + var eventsToPublish []*events.Event + err := svc.db.Transaction(func(tx *gorm.DB) error { + existingSettledTransaction, err := svc.findSettledTransaction(tx, dbTransaction) + if err != nil { + return err + } + if existingSettledTransaction != nil { + logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).Debug("payment already marked as settled") + settledTransaction = existingSettledTransaction + return nil + } + + settledAt := time.Now() + dbTransaction.State = constants.TRANSACTION_STATE_SETTLED + dbTransaction.Preimage = &preimage + dbTransaction.FeeMsat = feeMsat + dbTransaction.FeeReserveMsat = 0 + dbTransaction.SettledAt = &settledAt + dbTransaction.SelfPayment = selfPayment + if err := tx.Create(dbTransaction).Error; err != nil { + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": dbTransaction.PaymentHash, + }).WithError(err).Error("Failed to create settled DB transaction") + return err + } + + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": dbTransaction.PaymentHash, + "type": dbTransaction.Type, + }).Info("Created settled transaction") + + settledTransaction = dbTransaction + eventsToPublish = svc.afterTransactionSettled(tx, dbTransaction, &settledAt) + return nil + }) + if err != nil { + return nil, err + } + svc.publishEvents(eventsToPublish) + + return settledTransaction, nil +} + +// lockTransactionsByPaymentHash takes a row lock on all transactions with the +// given payment hash on postgres, so that concurrent state changes for the +// same payment serialize (in sqlite transactions are serializable by default). +func (svc *transactionsService) lockTransactionsByPaymentHash(tx *gorm.DB, paymentHash string) error { + if tx.Dialector.Name() != "postgres" { + return nil + } + transactionsWithPaymentHash := []db.Transaction{} + err := tx.Where(&db.Transaction{ + PaymentHash: paymentHash, + }).Clauses(clause.Locking{Strength: "UPDATE"}).Find(&transactionsWithPaymentHash).Error + if err != nil { + logger.Logger.WithField("payment_hash", paymentHash).WithError(err).Error("Failed to lock transactions by payment hash") + } + return err +} + +// findSettledTransaction returns the already-settled transaction matching +// dbTransaction if one exists, locking all transactions with the same payment +// hash to ensure only one transaction is settled per payment. +func (svc *transactionsService) findSettledTransaction(tx *gorm.DB, dbTransaction *db.Transaction) (*db.Transaction, error) { + if err := svc.lockTransactionsByPaymentHash(tx, dbTransaction.PaymentHash); err != nil { + return nil, err } var existingSettledTransaction db.Transaction - if tx.Limit(1).Find(&existingSettledTransaction, &db.Transaction{ + result := tx.Limit(1).Find(&existingSettledTransaction, &db.Transaction{ Type: dbTransaction.Type, PaymentRequest: dbTransaction.PaymentRequest, PaymentHash: dbTransaction.PaymentHash, State: constants.TRANSACTION_STATE_SETTLED, - }).RowsAffected > 0 { - logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).Debug("payment already marked as sent") + }) + if result.Error != nil { + logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).WithError(result.Error).Error("Failed to check for existing settled transaction") + return nil, result.Error + } + if result.RowsAffected > 0 { return &existingSettledTransaction, nil } + return nil, nil +} - settledAt := time.Now() - err := tx.Model(dbTransaction).Updates(map[string]interface{}{ - "State": constants.TRANSACTION_STATE_SETTLED, - "Preimage": &preimage, - "FeeMsat": feeMsat, - "FeeReserveMsat": 0, - "SettledAt": &settledAt, - "SelfPayment": selfPayment, - }).Error - if err != nil { - logger.Logger.WithFields(logrus.Fields{ - "payment_hash": dbTransaction.PaymentHash, - }).WithError(err).Error("Failed to update DB transaction") - return nil, err - } - - logger.Logger.WithFields(logrus.Fields{ - "payment_hash": dbTransaction.PaymentHash, - "type": dbTransaction.Type, - }).Info("Marked transaction as settled") - +// afterTransactionSettled runs the post-settlement side effects within the +// caller's database transaction and returns the events to publish after it +// commits. +func (svc *transactionsService) afterTransactionSettled(tx *gorm.DB, dbTransaction *db.Transaction, settledAt *time.Time) []*events.Event { event := "nwc_payment_sent" if dbTransaction.Type == constants.TRANSACTION_TYPE_INCOMING { event = "nwc_payment_received" } - svc.eventPublisher.Publish(&events.Event{ + eventsToPublish := []*events.Event{{ Event: event, Properties: dbTransaction, - }) + }} if dbTransaction.AppId != nil { var app db.App @@ -1564,17 +1637,25 @@ func (svc *transactionsService) markTransactionSettled(tx *gorm.DB, dbTransactio }) if result.RowsAffected == 0 { logger.Logger.WithField("app_id", dbTransaction.AppId).Error("failed to find app by id") - return dbTransaction, nil + return eventsToPublish } - svc.updateAppLastSettledTransactionAt(&app, tx, &settledAt) + svc.updateAppLastSettledTransactionAt(&app, tx, settledAt) if dbTransaction.Type == constants.TRANSACTION_TYPE_OUTGOING { - svc.checkBudgetUsage(&app, dbTransaction, tx) + if budgetWarningEvent := svc.checkBudgetUsage(&app, dbTransaction, tx); budgetWarningEvent != nil { + eventsToPublish = append(eventsToPublish, budgetWarningEvent) + } } } - return dbTransaction, nil + return eventsToPublish +} + +func (svc *transactionsService) publishEvents(eventsToPublish []*events.Event) { + for _, event := range eventsToPublish { + svc.eventPublisher.Publish(event) + } } func (svc *transactionsService) updateAppLastSettledTransactionAt(app *db.App, gormTransaction *gorm.DB, settledAt *time.Time) { @@ -1584,9 +1665,11 @@ func (svc *transactionsService) updateAppLastSettledTransactionAt(app *db.App, g } } -func (svc *transactionsService) checkBudgetUsage(app *db.App, dbTransaction *db.Transaction, gormTransaction *gorm.DB) { +// checkBudgetUsage returns a budget warning event to publish after the +// caller's database transaction commits, or nil if no warning is due. +func (svc *transactionsService) checkBudgetUsage(app *db.App, dbTransaction *db.Transaction, gormTransaction *gorm.DB) *events.Event { if app.Isolated { - return + return nil } var appPermission db.AppPermission @@ -1596,59 +1679,91 @@ func (svc *transactionsService) checkBudgetUsage(app *db.App, dbTransaction *db. }) if result.RowsAffected == 0 { logger.Logger.WithField("app_id", dbTransaction.AppId).Error("failed to find pay_invoice scope") - return + return nil } budgetUsageMsat, err := queries.GetBudgetUsageMsat(gormTransaction, &appPermission) if err != nil { logger.Logger.WithField("app_id", dbTransaction.AppId).WithError(err).Error("failed to get budget usage") - return + return nil } budgetUsageSat := budgetUsageMsat / 1000 warningUsage := uint64(math.Floor(float64(appPermission.MaxAmountSat) * 0.8)) if budgetUsageSat >= warningUsage && budgetUsageSat-dbTransaction.AmountMsat/1000 < warningUsage { - svc.eventPublisher.Publish(&events.Event{ + return &events.Event{ Event: "nwc_budget_warning", Properties: map[string]interface{}{ "name": app.Name, "id": app.ID, }, - }) + } } -} - -func (svc *transactionsService) markPaymentFailed(tx *gorm.DB, dbTransaction *db.Transaction, reason string) error { - var existingTransaction db.Transaction - result := tx.Limit(1).Find(&existingTransaction, &db.Transaction{ - ID: dbTransaction.ID, - }) - - if result.Error != nil { - logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).WithError(result.Error).Error("could not find transaction to mark as failed") - return result.Error - } - - if existingTransaction.State == constants.TRANSACTION_STATE_FAILED { - logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).Info("payment already marked as failed") - return nil - } - - err := tx.Model(dbTransaction).Updates(map[string]interface{}{ - "State": constants.TRANSACTION_STATE_FAILED, - "FeeReserveMsat": 0, - "FailureReason": reason, - }).Error - if err != nil { - logger.Logger.WithFields(logrus.Fields{ - "payment_hash": dbTransaction.PaymentHash, - }).WithError(err).Error("Failed to mark transaction as failed") - return err - } - logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).Info("Marked transaction as failed") - - svc.eventPublisher.Publish(&events.Event{ - Event: "nwc_payment_failed", - Properties: dbTransaction, - }) return nil } + +// markPaymentFailed marks the transaction as failed in its own database +// transaction and publishes the failed event after it commits, so subscribers +// never observe uncommitted state. It returns whether this call transitioned +// the transaction to failed (false if it was already failed), and refuses to +// mark a settled transaction as failed. +func (svc *transactionsService) markPaymentFailed(dbTransaction *db.Transaction, reason string) (bool, error) { + markedFailed := false + var eventsToPublish []*events.Event + err := svc.db.Transaction(func(tx *gorm.DB) error { + // lock all transactions with the same payment hash so a concurrent + // settlement cannot slip in between the state check and the update + if err := svc.lockTransactionsByPaymentHash(tx, dbTransaction.PaymentHash); err != nil { + return err + } + + var existingTransaction db.Transaction + result := tx.Limit(1).Find(&existingTransaction, &db.Transaction{ + ID: dbTransaction.ID, + }) + + if result.Error != nil { + logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).WithError(result.Error).Error("could not find transaction to mark as failed") + return result.Error + } + + if result.RowsAffected == 0 { + logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).Error("could not find transaction to mark as failed") + return NewNotFoundError() + } + + if existingTransaction.State == constants.TRANSACTION_STATE_FAILED { + logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).Info("payment already marked as failed") + return nil + } + + if existingTransaction.State == constants.TRANSACTION_STATE_SETTLED { + logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).Error("cannot mark settled payment as failed") + return errors.New("cannot mark settled payment as failed") + } + + err := tx.Model(dbTransaction).Updates(map[string]interface{}{ + "State": constants.TRANSACTION_STATE_FAILED, + "FeeReserveMsat": 0, + "FailureReason": reason, + }).Error + if err != nil { + logger.Logger.WithFields(logrus.Fields{ + "payment_hash": dbTransaction.PaymentHash, + }).WithError(err).Error("Failed to mark transaction as failed") + return err + } + logger.Logger.WithField("payment_hash", dbTransaction.PaymentHash).Info("Marked transaction as failed") + + markedFailed = true + eventsToPublish = append(eventsToPublish, &events.Event{ + Event: "nwc_payment_failed", + Properties: dbTransaction, + }) + return nil + }) + if err != nil { + return false, err + } + svc.publishEvents(eventsToPublish) + return markedFailed, nil +}