Skip to content
Merged
Show file tree
Hide file tree
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
5 changes: 5 additions & 0 deletions core/deleter/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
28 changes: 2 additions & 26 deletions internal/store/postgres/billing_checkout_repository.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand Down
101 changes: 101 additions & 0 deletions internal/store/postgres/billing_checkout_repository_pg_test.go
Original file line number Diff line number Diff line change
@@ -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))
}
30 changes: 3 additions & 27 deletions internal/store/postgres/billing_subscription_repository.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand All @@ -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()
Expand Down Expand Up @@ -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,
Expand Down
112 changes: 112 additions & 0 deletions internal/store/postgres/billing_subscription_repository_pg_test.go
Original file line number Diff line number Diff line change
@@ -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))
}
Loading