diff --git a/core/deleter/service.go b/core/deleter/service.go index 78de9e797..5cd2f885a 100644 --- a/core/deleter/service.go +++ b/core/deleter/service.go @@ -403,6 +403,11 @@ func (d Service) DeleteCustomers(ctx context.Context, id string) error { // here. func (d Service) deleteCustomers(ctx context.Context, id string, customers []customer.Customer, amounts map[string]accountTokens) error { for _, c := range customers { + // TODO(fix): the subscription and checkout reads here skip rows with + // deleted_at. Once these deletes turn soft, a row left behind with + // deleted_at set is invisible to this loop. A subscription then blocks the + // customer delete on its foreign key, and a checkout is removed without + // its audit record below. Make these deletes soft in the same change. // cancels active subscriptions on the billing provider and removes local records if err := d.subService.DeleteByCustomer(ctx, c); err != nil { return fmt.Errorf("failed to delete org while deleting a billing account subscriptions[%s]: %w", c.ID, err) diff --git a/internal/store/postgres/billing_checkout_repository.go b/internal/store/postgres/billing_checkout_repository.go index 881a96a9d..5d5cf481d 100644 --- a/internal/store/postgres/billing_checkout_repository.go +++ b/internal/store/postgres/billing_checkout_repository.go @@ -212,7 +212,7 @@ func (r BillingCheckoutRepository) Create(ctx context.Context, toCreate checkout } func (r BillingCheckoutRepository) GetByID(ctx context.Context, id string) (checkout.Checkout, error) { - stmt := dialect.Select().From(TABLE_BILLING_CHECKOUTS).Where(goqu.Ex{ + stmt := fromLive(TABLE_BILLING_CHECKOUTS).Where(goqu.Ex{ "id": id, }) query, params, err := stmt.ToSQL() @@ -235,30 +235,6 @@ func (r BillingCheckoutRepository) GetByID(ctx context.Context, id string) (chec return checkoutModel.transform() } -func (r BillingCheckoutRepository) GetByName(ctx context.Context, name string) (checkout.Checkout, error) { - stmt := dialect.Select().From(TABLE_BILLING_CHECKOUTS).Where(goqu.Ex{ - "name": name, - }) - query, params, err := stmt.ToSQL() - if err != nil { - return checkout.Checkout{}, fmt.Errorf("%w: %s", errParse, err) - } - - var checkoutModel Checkout - if err = r.dbc.WithTimeout(ctx, TABLE_BILLING_CHECKOUTS, "GetByName", func(ctx context.Context) error { - return r.dbc.QueryRowxContext(ctx, query, params...).StructScan(&checkoutModel) - }); err != nil { - err = checkPostgresError(err) - switch { - case errors.Is(err, sql.ErrNoRows): - return checkout.Checkout{}, checkout.ErrNotFound - } - return checkout.Checkout{}, fmt.Errorf("%w: %s", errDB, err) - } - - return checkoutModel.transform() -} - func (r BillingCheckoutRepository) UpdateByID(ctx context.Context, toUpdate checkout.Checkout) (checkout.Checkout, error) { if strings.TrimSpace(toUpdate.ID) == "" { return checkout.Checkout{}, checkout.ErrInvalidID @@ -323,7 +299,7 @@ func (r BillingCheckoutRepository) DeleteByCustomerID(ctx context.Context, custo } func (r BillingCheckoutRepository) List(ctx context.Context, flt checkout.Filter) ([]checkout.Checkout, error) { - stmt := dialect.Select().From(TABLE_BILLING_CHECKOUTS).Order(goqu.I("created_at").Desc()) + stmt := fromLive(TABLE_BILLING_CHECKOUTS).Order(goqu.I("created_at").Desc()) if flt.CustomerID != "" { stmt = stmt.Where(goqu.Ex{ "customer_id": flt.CustomerID, diff --git a/internal/store/postgres/billing_checkout_repository_pg_test.go b/internal/store/postgres/billing_checkout_repository_pg_test.go new file mode 100644 index 000000000..9d54f3d2d --- /dev/null +++ b/internal/store/postgres/billing_checkout_repository_pg_test.go @@ -0,0 +1,101 @@ +package postgres_test + +import ( + "context" + "fmt" + "testing" + + "github.com/raystack/frontier/billing/checkout" + "github.com/raystack/frontier/internal/store/postgres" + "github.com/raystack/frontier/pkg/db" + "github.com/stretchr/testify/suite" +) + +// Runs the billing checkout reads against a real postgres to check that a +// soft-deleted checkout stays out of every read. +type BillingCheckoutRepositoryPGTestSuite struct { + suite.Suite + ctx context.Context + client *db.Client + repository *postgres.BillingCheckoutRepository +} + +func (s *BillingCheckoutRepositoryPGTestSuite) SetupSuite() { + var err error + s.client, err = newTestClient() + if err != nil { + s.T().Fatal(err) + } + s.ctx = context.TODO() + s.repository = postgres.NewBillingCheckoutRepository(s.client) +} + +func (s *BillingCheckoutRepositoryPGTestSuite) TearDownSuite() { + if err := closeTestClient(s.client); err != nil { + s.T().Fatal(err) + } +} + +func (s *BillingCheckoutRepositoryPGTestSuite) SetupTest() { + s.exec(`INSERT INTO organizations (name, title) VALUES ('bch-live', 'Live Org')`) + s.exec(`INSERT INTO billing_customers (org_id, provider_id, name, email) + VALUES ((SELECT id FROM organizations WHERE name = 'bch-live'), 'bch-cust', 'bch-cust', 'bch-cust')`) + + s.checkout("bch-live-session") + s.checkout("bch-gone-session") + + s.exec(`UPDATE billing_checkouts SET deleted_at = now() WHERE provider_id = 'bch-gone-session'`) +} + +func (s *BillingCheckoutRepositoryPGTestSuite) TearDownTest() { + queries := []string{} + for _, table := range []string{postgres.TABLE_BILLING_CHECKOUTS, postgres.TABLE_BILLING_CUSTOMERS, + postgres.TABLE_ORGANIZATIONS} { + queries = append(queries, fmt.Sprintf("TRUNCATE TABLE %s RESTART IDENTITY CASCADE", table)) + } + if err := execQueries(s.ctx, s.client, queries); err != nil { + s.T().Fatal(err) + } +} + +func (s *BillingCheckoutRepositoryPGTestSuite) exec(query string, args ...any) { + s.T().Helper() + execSQL(s.T(), s.ctx, s.client, query, args...) +} + +func (s *BillingCheckoutRepositoryPGTestSuite) customerID() string { + s.T().Helper() + return scalarSQL(s.T(), s.ctx, s.client, `SELECT id FROM billing_customers WHERE name = 'bch-cust'`) +} + +func (s *BillingCheckoutRepositoryPGTestSuite) checkoutID(providerID string) string { + s.T().Helper() + return scalarSQL(s.T(), s.ctx, s.client, + `SELECT id FROM billing_checkouts WHERE provider_id = $1`, providerID) +} + +func (s *BillingCheckoutRepositoryPGTestSuite) checkout(providerID string) { + s.T().Helper() + s.exec(`INSERT INTO billing_checkouts (customer_id, provider_id, checkout_url, state) + VALUES ((SELECT id FROM billing_customers WHERE name = 'bch-cust'), $1, $1, 'pending')`, providerID) +} + +func (s *BillingCheckoutRepositoryPGTestSuite) TestGetByIDSkipsDeleted() { + got, err := s.repository.GetByID(s.ctx, s.checkoutID("bch-live-session")) + s.Require().NoError(err) + s.Equal("bch-live-session", got.ProviderID) + + _, err = s.repository.GetByID(s.ctx, s.checkoutID("bch-gone-session")) + s.ErrorIs(err, checkout.ErrNotFound) +} + +func (s *BillingCheckoutRepositoryPGTestSuite) TestListSkipsDeleted() { + got, err := s.repository.List(s.ctx, checkout.Filter{CustomerID: s.customerID()}) + s.Require().NoError(err) + s.Require().Len(got, 1) + s.Equal("bch-live-session", got[0].ProviderID) +} + +func TestBillingCheckoutRepositoryPG(t *testing.T) { + suite.Run(t, new(BillingCheckoutRepositoryPGTestSuite)) +} diff --git a/internal/store/postgres/billing_subscription_repository.go b/internal/store/postgres/billing_subscription_repository.go index e2f32ac0f..4396d7ed5 100644 --- a/internal/store/postgres/billing_subscription_repository.go +++ b/internal/store/postgres/billing_subscription_repository.go @@ -241,7 +241,7 @@ func (r BillingSubscriptionRepository) Create(ctx context.Context, toCreate subs } func (r BillingSubscriptionRepository) GetByID(ctx context.Context, id string) (subscription.Subscription, error) { - stmt := dialect.Select().From(TABLE_BILLING_SUBSCRIPTIONS).Where(goqu.Ex{ + stmt := fromLive(TABLE_BILLING_SUBSCRIPTIONS).Where(goqu.Ex{ "id": id, }) query, params, err := stmt.ToSQL() @@ -264,32 +264,8 @@ func (r BillingSubscriptionRepository) GetByID(ctx context.Context, id string) ( return subscriptionModel.transform() } -func (r BillingSubscriptionRepository) GetByName(ctx context.Context, name string) (subscription.Subscription, error) { - stmt := dialect.Select().From(TABLE_BILLING_SUBSCRIPTIONS).Where(goqu.Ex{ - "name": name, - }) - query, params, err := stmt.ToSQL() - if err != nil { - return subscription.Subscription{}, fmt.Errorf("%w: %s", errParse, err) - } - - var subscriptionModel Subscription - if err = r.dbc.WithTimeout(ctx, TABLE_BILLING_SUBSCRIPTIONS, "GetByName", func(ctx context.Context) error { - return r.dbc.QueryRowxContext(ctx, query, params...).StructScan(&subscriptionModel) - }); err != nil { - err = checkPostgresError(err) - switch { - case errors.Is(err, sql.ErrNoRows): - return subscription.Subscription{}, subscription.ErrNotFound - } - return subscription.Subscription{}, fmt.Errorf("%w: %s", errDB, err) - } - - return subscriptionModel.transform() -} - func (r BillingSubscriptionRepository) GetByProviderID(ctx context.Context, id string) (subscription.Subscription, error) { - stmt := dialect.Select().From(TABLE_BILLING_SUBSCRIPTIONS).Where(goqu.Ex{ + stmt := fromLive(TABLE_BILLING_SUBSCRIPTIONS).Where(goqu.Ex{ "provider_id": id, }) query, params, err := stmt.ToSQL() @@ -429,7 +405,7 @@ func (r BillingSubscriptionRepository) toSubscriptionChanges(toUpdate subscripti } func (r BillingSubscriptionRepository) List(ctx context.Context, filter subscription.Filter) ([]subscription.Subscription, error) { - stmt := dialect.Select().From(TABLE_BILLING_SUBSCRIPTIONS).Order(goqu.I("created_at").Desc()) + stmt := fromLive(TABLE_BILLING_SUBSCRIPTIONS).Order(goqu.I("created_at").Desc()) if filter.CustomerID != "" { stmt = stmt.Where(goqu.Ex{ "customer_id": filter.CustomerID, diff --git a/internal/store/postgres/billing_subscription_repository_pg_test.go b/internal/store/postgres/billing_subscription_repository_pg_test.go new file mode 100644 index 000000000..bd57707d1 --- /dev/null +++ b/internal/store/postgres/billing_subscription_repository_pg_test.go @@ -0,0 +1,112 @@ +package postgres_test + +import ( + "context" + "fmt" + "testing" + + "github.com/raystack/frontier/billing/subscription" + "github.com/raystack/frontier/internal/store/postgres" + "github.com/raystack/frontier/pkg/db" + "github.com/stretchr/testify/suite" +) + +// Runs the billing subscription reads against a real postgres to check that a +// soft-deleted subscription stays out of every read. +type BillingSubscriptionRepositoryPGTestSuite struct { + suite.Suite + ctx context.Context + client *db.Client + repository *postgres.BillingSubscriptionRepository +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) SetupSuite() { + var err error + s.client, err = newTestClient() + if err != nil { + s.T().Fatal(err) + } + s.ctx = context.TODO() + s.repository = postgres.NewBillingSubscriptionRepository(s.client) +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) TearDownSuite() { + if err := closeTestClient(s.client); err != nil { + s.T().Fatal(err) + } +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) SetupTest() { + s.exec(`INSERT INTO organizations (name, title) VALUES ('bs-live', 'Live Org')`) + s.exec(`INSERT INTO billing_customers (org_id, provider_id, name, email) + VALUES ((SELECT id FROM organizations WHERE name = 'bs-live'), 'bs-cust', 'bs-cust', 'bs-cust')`) + s.exec(`INSERT INTO billing_plans (name, description) VALUES ('bs-plan', 'plan under test')`) + + s.subscription("bs-sub-live") + s.subscription("bs-sub-gone") + + s.exec(`UPDATE billing_subscriptions SET deleted_at = now() WHERE provider_id = 'bs-sub-gone'`) +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) TearDownTest() { + queries := []string{} + for _, table := range []string{postgres.TABLE_BILLING_SUBSCRIPTIONS, postgres.TABLE_BILLING_CUSTOMERS, + postgres.TABLE_BILLING_PLANS, postgres.TABLE_ORGANIZATIONS} { + queries = append(queries, fmt.Sprintf("TRUNCATE TABLE %s RESTART IDENTITY CASCADE", table)) + } + if err := execQueries(s.ctx, s.client, queries); err != nil { + s.T().Fatal(err) + } +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) exec(query string, args ...any) { + s.T().Helper() + execSQL(s.T(), s.ctx, s.client, query, args...) +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) customerID() string { + s.T().Helper() + return scalarSQL(s.T(), s.ctx, s.client, `SELECT id FROM billing_customers WHERE name = 'bs-cust'`) +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) subscriptionID(providerID string) string { + s.T().Helper() + return scalarSQL(s.T(), s.ctx, s.client, + `SELECT id FROM billing_subscriptions WHERE provider_id = $1`, providerID) +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) subscription(providerID string) { + s.T().Helper() + s.exec(`INSERT INTO billing_subscriptions (customer_id, provider_id, plan_id, state) + VALUES ((SELECT id FROM billing_customers WHERE name = 'bs-cust'), $1, + (SELECT id FROM billing_plans WHERE name = 'bs-plan'), 'active')`, providerID) +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) TestGetByIDSkipsDeleted() { + got, err := s.repository.GetByID(s.ctx, s.subscriptionID("bs-sub-live")) + s.Require().NoError(err) + s.Equal("bs-sub-live", got.ProviderID) + + _, err = s.repository.GetByID(s.ctx, s.subscriptionID("bs-sub-gone")) + s.ErrorIs(err, subscription.ErrNotFound) +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) TestGetByProviderIDSkipsDeleted() { + got, err := s.repository.GetByProviderID(s.ctx, "bs-sub-live") + s.Require().NoError(err) + s.Equal("bs-sub-live", got.ProviderID) + + _, err = s.repository.GetByProviderID(s.ctx, "bs-sub-gone") + s.ErrorIs(err, subscription.ErrNotFound) +} + +func (s *BillingSubscriptionRepositoryPGTestSuite) TestListSkipsDeleted() { + got, err := s.repository.List(s.ctx, subscription.Filter{CustomerID: s.customerID()}) + s.Require().NoError(err) + s.Require().Len(got, 1) + s.Equal("bs-sub-live", got[0].ProviderID) +} + +func TestBillingSubscriptionRepositoryPG(t *testing.T) { + suite.Run(t, new(BillingSubscriptionRepositoryPGTestSuite)) +}