diff --git a/internal/store/postgres/organization_repository.go b/internal/store/postgres/organization_repository.go index e7a849e3f..ba15574c9 100644 --- a/internal/store/postgres/organization_repository.go +++ b/internal/store/postgres/organization_repository.go @@ -41,7 +41,7 @@ func (r OrganizationRepository) GetByID(ctx context.Context, id string) (organiz return organization.Organization{}, organization.ErrInvalidID } - query, params, err := dialect.From(TABLE_ORGANIZATIONS).Where(goqu.Ex{ + query, params, err := fromLive(TABLE_ORGANIZATIONS).Where(goqu.Ex{ "id": id, }).ToSQL() if err != nil { @@ -76,7 +76,7 @@ func (r OrganizationRepository) GetByIDs(ctx context.Context, ids []string) ([]o return nil, organization.ErrInvalidID } - query, params, err := dialect.From(TABLE_ORGANIZATIONS).Where(goqu.Ex{ + query, params, err := fromLive(TABLE_ORGANIZATIONS).Where(goqu.Ex{ "id": goqu.Op{"in": ids}, }).Where(notDisabledOrgExp).ToSQL() if err != nil { @@ -114,7 +114,7 @@ func (r OrganizationRepository) GetByName(ctx context.Context, name string) (org return organization.Organization{}, organization.ErrInvalidID } - query, params, err := dialect.From(TABLE_ORGANIZATIONS).Where(goqu.Ex{ + query, params, err := fromLive(TABLE_ORGANIZATIONS).Where(goqu.Ex{ "name": name, }).ToSQL() if err != nil { @@ -216,7 +216,7 @@ func (r OrganizationRepository) Create(ctx context.Context, org organization.Org } func (r OrganizationRepository) List(ctx context.Context, flt organization.Filter) ([]organization.Organization, error) { - stmt := dialect.From(TABLE_ORGANIZATIONS) + stmt := fromLive(TABLE_ORGANIZATIONS) if flt.State == "" { stmt = stmt.Where(notDisabledOrgExp) } else { @@ -326,7 +326,7 @@ func (r OrganizationRepository) UpdateByID(ctx context.Context, org organization } // Query to fetch org title before update - getQuery, getParams, err := dialect.From(TABLE_ORGANIZATIONS). + getQuery, getParams, err := fromLive(TABLE_ORGANIZATIONS). Select("title"). Where(goqu.Ex{"id": org.ID}).ToSQL() if err != nil { @@ -341,7 +341,7 @@ func (r OrganizationRepository) UpdateByID(ctx context.Context, org organization "updated_at": goqu.L("now()"), }).Where(goqu.Ex{ "id": org.ID, - }).Returning(&Organization{}).ToSQL() + }, live(TABLE_ORGANIZATIONS)).Returning(&Organization{}).ToSQL() if err != nil { return organization.Organization{}, fmt.Errorf("%w: %w", errQuery, err) } @@ -397,7 +397,7 @@ func (r OrganizationRepository) UpdateByName(ctx context.Context, org organizati } // Query to fetch org data before update - getQuery, getParams, err := dialect.From(TABLE_ORGANIZATIONS). + getQuery, getParams, err := fromLive(TABLE_ORGANIZATIONS). Select("title"). Where(goqu.Ex{"name": org.Name}).ToSQL() if err != nil { @@ -413,7 +413,7 @@ func (r OrganizationRepository) UpdateByName(ctx context.Context, org organizati }).Where( goqu.Ex{ "name": org.Name, - }).Returning(&Organization{}).ToSQL() + }, live(TABLE_ORGANIZATIONS)).Returning(&Organization{}).ToSQL() if err != nil { return organization.Organization{}, fmt.Errorf("%w: %w", errQuery, err) } @@ -464,6 +464,7 @@ func (r OrganizationRepository) SetState(ctx context.Context, id string, state o goqu.Ex{ "id": id, }, + live(TABLE_ORGANIZATIONS), ).Returning(&Organization{}).ToSQL() if err != nil { return fmt.Errorf("%w: %w", errQuery, err) diff --git a/internal/store/postgres/organization_repository_test.go b/internal/store/postgres/organization_repository_test.go index e145432d6..40b1027b3 100644 --- a/internal/store/postgres/organization_repository_test.go +++ b/internal/store/postgres/organization_repository_test.go @@ -566,3 +566,39 @@ func (s *OrganizationRepositoryTestSuite) TestGetByIDs() { func TestOrganizationRepository(t *testing.T) { suite.Run(t, new(OrganizationRepositoryTestSuite)) } + +func (s *OrganizationRepositoryTestSuite) TestSkipsSoftDeletedOrganizations() { + deleted := s.orgs[0] + _, err := s.client.ExecContext(s.ctx, "UPDATE organizations SET deleted_at = now() WHERE id = $1", deleted.ID) + if err != nil { + s.T().Fatal(err) + } + + _, err = s.repository.GetByID(s.ctx, deleted.ID) + s.Assert().ErrorIs(err, organization.ErrNotExist) + + _, err = s.repository.GetByName(s.ctx, deleted.Name) + s.Assert().ErrorIs(err, organization.ErrNotExist) + + byIDs, err := s.repository.GetByIDs(s.ctx, []string{deleted.ID}) + s.Assert().NoError(err) + s.Assert().Empty(byIDs) + + got, err := s.repository.List(s.ctx, organization.Filter{}) + s.Assert().NoError(err) + s.Assert().Len(got, len(s.orgs)-1) + for _, o := range got { + s.Assert().NotEqual(deleted.ID, o.ID) + } + + _, err = s.repository.UpdateByName(s.ctx, organization.Organization{Name: deleted.Name, Title: "changed"}) + s.Assert().ErrorIs(err, organization.ErrNotExist) + + byID := deleted + byID.Title = "changed" + _, err = s.repository.UpdateByID(s.ctx, byID) + s.Assert().ErrorIs(err, organization.ErrNotExist) + + err = s.repository.SetState(s.ctx, deleted.ID, organization.Disabled) + s.Assert().ErrorIs(err, organization.ErrNotExist) +} diff --git a/internal/store/postgres/postgres.go b/internal/store/postgres/postgres.go index eb3dadea6..bb85e16ea 100644 --- a/internal/store/postgres/postgres.go +++ b/internal/store/postgres/postgres.go @@ -6,6 +6,7 @@ import ( "github.com/doug-martin/goqu/v9" _ "github.com/doug-martin/goqu/v9/dialect/postgres" + "github.com/doug-martin/goqu/v9/exp" "github.com/jackc/pgerrcode" "github.com/jackc/pgx/v5/pgconn" _ "github.com/jackc/pgx/v5/stdlib" @@ -22,6 +23,17 @@ var ( dialect = goqu.Dialect("postgres") ) +// live is the filter every read of a soft-deleted table adds. The column is +// qualified with the table name or alias so it stays correct inside joins. +func live(table string) exp.BooleanExpression { + return goqu.I(table + ".deleted_at").IsNull() +} + +// fromLive reads only the rows that are not soft-deleted. +func fromLive(table string) *goqu.SelectDataset { + return dialect.From(table).Where(live(table)) +} + const ( TABLE_PERMISSIONS = "permissions" TABLE_GROUPS = "groups" diff --git a/internal/store/postgres/project_repository.go b/internal/store/postgres/project_repository.go index 5e09b4828..be4720c73 100644 --- a/internal/store/postgres/project_repository.go +++ b/internal/store/postgres/project_repository.go @@ -11,7 +11,6 @@ import ( "github.com/doug-martin/goqu/v9" "github.com/raystack/frontier/core/organization" "github.com/raystack/frontier/core/project" - "github.com/raystack/frontier/core/user" "github.com/raystack/frontier/pkg/db" ) @@ -39,7 +38,7 @@ func (r ProjectRepository) GetByID(ctx context.Context, id string) (project.Proj return project.Project{}, project.ErrInvalidID } - query, params, err := dialect.From(TABLE_PROJECTS).Where(goqu.ExOr{ + query, params, err := fromLive(TABLE_PROJECTS).Where(goqu.ExOr{ "id": id, }).Where(notDisabledProjectExp).ToSQL() if err != nil { @@ -74,7 +73,7 @@ func (r ProjectRepository) GetByName(ctx context.Context, name string) (project. return project.Project{}, project.ErrInvalidID } - query, params, err := dialect.From(TABLE_PROJECTS).Where(goqu.Ex{ + query, params, err := fromLive(TABLE_PROJECTS).Where(goqu.Ex{ "name": name, }).Where(notDisabledProjectExp).ToSQL() if err != nil { @@ -155,7 +154,7 @@ func (r ProjectRepository) Create(ctx context.Context, prj project.Project) (pro } func (r ProjectRepository) List(ctx context.Context, flt project.Filter) ([]project.Project, error) { - stmt := dialect.From(TABLE_PROJECTS) + stmt := fromLive(TABLE_PROJECTS) if flt.OrgID != "" { stmt = stmt.Where(goqu.Ex{ "org_id": flt.OrgID, @@ -245,7 +244,7 @@ func (r ProjectRepository) UpdateByID(ctx context.Context, prj project.Project) "title": prj.Title, "metadata": marshaledMetadata, "updated_at": goqu.L("now()"), - }).Where(goqu.Ex{"id": prj.ID}).Returning(&Project{}).ToSQL() + }).Where(goqu.Ex{"id": prj.ID}, live(TABLE_PROJECTS)).Returning(&Project{}).ToSQL() if err != nil { return project.Project{}, fmt.Errorf("%w: %s", errQuery, err) } @@ -290,7 +289,7 @@ func (r ProjectRepository) UpdateByName(ctx context.Context, prj project.Project "title": prj.Title, "metadata": marshaledMetadata, "updated_at": goqu.L("now()"), - }).Where(goqu.Ex{"name": prj.Name}).Returning(&Project{}).ToSQL() + }).Where(goqu.Ex{"name": prj.Name}, live(TABLE_PROJECTS)).Returning(&Project{}).ToSQL() if err != nil { return project.Project{}, fmt.Errorf("%w: %s", errQuery, err) } @@ -328,21 +327,20 @@ func (r ProjectRepository) SetState(ctx context.Context, id string, state projec goqu.Ex{ "id": id, }, - ).ToSQL() + live(TABLE_PROJECTS), + ).Returning(&Project{}).ToSQL() if err != nil { return fmt.Errorf("%w: %s", errQuery, err) } + var projectModel Project if err = r.dbc.WithTimeout(ctx, TABLE_PROJECTS, "SetState", func(ctx context.Context) error { - if _, err = r.dbc.DB.ExecContext(ctx, query, params...); err != nil { - return err - } - return nil + return r.dbc.QueryRowxContext(ctx, query, params...).StructScan(&projectModel) }); err != nil { err = checkPostgresError(err) switch { case errors.Is(err, sql.ErrNoRows): - return user.ErrNotExist + return project.ErrNotExist default: return err } diff --git a/internal/store/postgres/project_repository_test.go b/internal/store/postgres/project_repository_test.go index 9c4167cca..31151722c 100644 --- a/internal/store/postgres/project_repository_test.go +++ b/internal/store/postgres/project_repository_test.go @@ -470,3 +470,35 @@ func (s *ProjectRepositoryTestSuite) TestUpdateByName() { func TestProjectRepository(t *testing.T) { suite.Run(t, new(ProjectRepositoryTestSuite)) } + +func (s *ProjectRepositoryTestSuite) TestSkipsSoftDeletedProjects() { + deleted := s.projects[0] + _, err := s.client.ExecContext(s.ctx, "UPDATE projects SET deleted_at = now() WHERE id = $1", deleted.ID) + if err != nil { + s.T().Fatal(err) + } + + _, err = s.repository.GetByID(s.ctx, deleted.ID) + s.Assert().ErrorIs(err, project.ErrNotExist) + + _, err = s.repository.GetByName(s.ctx, deleted.Name) + s.Assert().ErrorIs(err, project.ErrNotExist) + + got, err := s.repository.List(s.ctx, project.Filter{OrgID: deleted.Organization.ID}) + s.Assert().NoError(err) + s.Assert().NotEmpty(got) + for _, p := range got { + s.Assert().NotEqual(deleted.ID, p.ID) + } + + _, err = s.repository.UpdateByName(s.ctx, project.Project{Name: deleted.Name, Title: "changed"}) + s.Assert().ErrorIs(err, project.ErrNotExist) + + byID := deleted + byID.Title = "changed" + _, err = s.repository.UpdateByID(s.ctx, byID) + s.Assert().ErrorIs(err, project.ErrNotExist) + + err = s.repository.SetState(s.ctx, deleted.ID, project.Disabled) + s.Assert().ErrorIs(err, project.ErrNotExist) +} diff --git a/internal/store/postgres/user_repository.go b/internal/store/postgres/user_repository.go index 716fbf49c..70dacb512 100644 --- a/internal/store/postgres/user_repository.go +++ b/internal/store/postgres/user_repository.go @@ -47,7 +47,7 @@ func (r UserRepository) GetByID(ctx context.Context, id string) (user.User, erro } var fetchedUser User - userQuery, params, err := dialect.From(TABLE_USERS). + userQuery, params, err := fromLive(TABLE_USERS). Where(goqu.Ex{ "id": id, }).Where(notDisabledUserExp).ToSQL() @@ -85,7 +85,7 @@ func (r UserRepository) GetByName(ctx context.Context, name string) (user.User, } var fetchedUser User - query, params, err := dialect.From(TABLE_USERS). + query, params, err := fromLive(TABLE_USERS). Where(goqu.Ex{ "name": strings.ToLower(name), }).ToSQL() @@ -247,7 +247,7 @@ func (r UserRepository) List(ctx context.Context, flt user.Filter) ([]user.User, } offset := (flt.Page - 1) * flt.Limit - sqlStmt := dialect.From(TABLE_USERS). + sqlStmt := fromLive(TABLE_USERS). Select("users.id", "name", "email", "title", "avatar", "users.created_at", "users.updated_at") if len(flt.Keyword) != 0 { @@ -297,7 +297,7 @@ func (r UserRepository) GetByIDs(ctx context.Context, userIDs []string) ([]user. } var fetchedUsers []User - query, params, err := dialect.From(TABLE_USERS).Select("id", "name", "email", "title", "avatar", "state").Where( + query, params, err := fromLive(TABLE_USERS).Select("id", "name", "email", "title", "avatar", "state").Where( goqu.Ex{ "id": goqu.Op{"in": userIDs}, }).Where(notDisabledUserExp).ToSQL() @@ -352,6 +352,7 @@ func (r UserRepository) UpdateByEmail(ctx context.Context, usr user.User) (user. goqu.Ex{ "email": strings.ToLower(usr.Email), }, + live(TABLE_USERS), ).Returning(&User{}).ToSQL() if err != nil { return fmt.Errorf("%w: %s", errQuery, err) @@ -403,6 +404,7 @@ func (r UserRepository) UpdateByID(ctx context.Context, usr user.User) (user.Use goqu.Ex{ "id": usr.ID, }, + live(TABLE_USERS), ).Returning(&User{}).ToSQL() if err != nil { return fmt.Errorf("%w: %s", errQuery, err) @@ -460,6 +462,7 @@ func (r UserRepository) UpdateByName(ctx context.Context, usr user.User) (user.U goqu.Ex{ "name": strings.ToLower(usr.Name), }, + live(TABLE_USERS), ).Returning(&User{}).ToSQL() if err != nil { return user.User{}, fmt.Errorf("%w: %s", errQuery, err) @@ -496,7 +499,7 @@ func (r UserRepository) GetByEmail(ctx context.Context, email string) (user.User } var fetchedUser User - query, params, err := dialect.From(TABLE_USERS).Where( + query, params, err := fromLive(TABLE_USERS).Where( goqu.Ex{ "email": strings.ToLower(email), }).Where(notDisabledUserExp).ToSQL() @@ -530,6 +533,7 @@ func (r UserRepository) SetState(ctx context.Context, id string, state user.Stat goqu.Ex{ "id": id, }, + live(TABLE_USERS), ).Returning(&User{}).ToSQL() if err != nil { return fmt.Errorf("%w: %s", errQuery, err) @@ -698,7 +702,7 @@ func (r UserRepository) PrepareDataQuery(input *rql.Query) (string, []any, error } func (r UserRepository) buildBaseQuery() *goqu.SelectDataset { - return dialect.From(TABLE_USERS).Prepared(true).Select( + return fromLive(TABLE_USERS).Prepared(true).Select( goqu.I(COLUMN_ID), goqu.I(COLUMN_NAME), goqu.I(COLUMN_EMAIL), @@ -772,7 +776,7 @@ func (r UserRepository) addSort(query *goqu.SelectDataset, input *rql.Query) (*g func (r UserRepository) PrepareGroupByQuery(input *rql.Query) (string, []any, error) { // Start with base query that includes COUNT and group by field - query := dialect.From(TABLE_USERS).Prepared(true). + query := fromLive(TABLE_USERS).Prepared(true). Select( goqu.COUNT("*").As("count"), goqu.I(TABLE_USERS+"."+input.GroupBy[0]).As("values"), diff --git a/internal/store/postgres/user_repository_test.go b/internal/store/postgres/user_repository_test.go index 6b23c9ddf..56938c530 100644 --- a/internal/store/postgres/user_repository_test.go +++ b/internal/store/postgres/user_repository_test.go @@ -653,7 +653,7 @@ func TestUserRepository_PrepareDataQuery(t *testing.T) { Offset: 10, Limit: 20, }, - wantSQL: `SELECT "id", "name", "email", "state", "avatar", "title", "created_at", "updated_at" FROM "users" WHERE (("CAST(users"."id AS TEXT)" = $1) AND ("users"."state" ILIKE $2) AND (("users"."email" IS NULL) OR ("users"."email" = $3)) AND ((CAST("id" AS TEXT) ILIKE $4) OR ("title" ILIKE $5) OR ("name" ILIKE $6) OR ("state" ILIKE $7))) ORDER BY "name" ASC, "created_at" DESC LIMIT $8 OFFSET $9`, + wantSQL: `SELECT "id", "name", "email", "state", "avatar", "title", "created_at", "updated_at" FROM "users" WHERE (("users"."deleted_at" IS NULL) AND ("CAST(users"."id AS TEXT)" = $1) AND ("users"."state" ILIKE $2) AND (("users"."email" IS NULL) OR ("users"."email" = $3)) AND ((CAST("id" AS TEXT) ILIKE $4) OR ("title" ILIKE $5) OR ("name" ILIKE $6) OR ("state" ILIKE $7))) ORDER BY "name" ASC, "created_at" DESC LIMIT $8 OFFSET $9`, wantParams: []any{int64(123), "%active%", "", "%john%", "%john%", "%john%", "%john%", int64(20), int64(10)}, }, { @@ -669,7 +669,7 @@ func TestUserRepository_PrepareDataQuery(t *testing.T) { Offset: 5, Limit: 15, }, - wantSQL: `SELECT "id", "name", "email", "state", "avatar", "title", "created_at", "updated_at" FROM "users" WHERE ("users"."state" = $1) ORDER BY "state" ASC, "name" ASC LIMIT $2 OFFSET $3`, + wantSQL: `SELECT "id", "name", "email", "state", "avatar", "title", "created_at", "updated_at" FROM "users" WHERE (("users"."deleted_at" IS NULL) AND ("users"."state" = $1)) ORDER BY "state" ASC, "name" ASC LIMIT $2 OFFSET $3`, wantParams: []any{ "active", int64(15), @@ -713,7 +713,7 @@ func TestUserRepository_PrepareGroupByQuery(t *testing.T) { GroupBy: []string{"state"}, Search: "test", }, - wantSQL: `SELECT COUNT(*) AS "count", "users"."state" AS "values" FROM "users" WHERE (("users"."state" = $1) AND ("CAST(users"."id AS TEXT)" = $2) AND ((CAST("id" AS TEXT) ILIKE $3) OR ("title" ILIKE $4) OR ("name" ILIKE $5) OR ("state" ILIKE $6))) GROUP BY "users"."state"`, + wantSQL: `SELECT COUNT(*) AS "count", "users"."state" AS "values" FROM "users" WHERE (("users"."deleted_at" IS NULL) AND ("users"."state" = $1) AND ("CAST(users"."id AS TEXT)" = $2) AND ((CAST("id" AS TEXT) ILIKE $3) OR ("title" ILIKE $4) OR ("name" ILIKE $5) OR ("state" ILIKE $6))) GROUP BY "users"."state"`, wantParams: []any{"active", int64(123), "%test%", "%test%", "%test%", "%test%"}, }, { @@ -722,7 +722,7 @@ func TestUserRepository_PrepareGroupByQuery(t *testing.T) { GroupBy: []string{"state"}, Search: "pending", }, - wantSQL: `SELECT COUNT(*) AS "count", "users"."state" AS "values" FROM "users" WHERE ((CAST("id" AS TEXT) ILIKE $1) OR ("title" ILIKE $2) OR ("name" ILIKE $3) OR ("state" ILIKE $4)) GROUP BY "users"."state"`, + wantSQL: `SELECT COUNT(*) AS "count", "users"."state" AS "values" FROM "users" WHERE (("users"."deleted_at" IS NULL) AND ((CAST("id" AS TEXT) ILIKE $1) OR ("title" ILIKE $2) OR ("name" ILIKE $3) OR ("state" ILIKE $4))) GROUP BY "users"."state"`, wantParams: []any{"%pending%", "%pending%", "%pending%", "%pending%"}, }, } @@ -743,3 +743,45 @@ func TestUserRepository_PrepareGroupByQuery(t *testing.T) { }) } } + +func (s *UserRepositoryTestSuite) TestSkipsSoftDeletedUsers() { + deleted := s.users[0] + _, err := s.client.ExecContext(s.ctx, "UPDATE users SET deleted_at = now() WHERE id = $1", deleted.ID) + if err != nil { + s.T().Fatal(err) + } + + _, err = s.repository.GetByID(s.ctx, deleted.ID) + s.Assert().ErrorIs(err, user.ErrNotExist) + + _, err = s.repository.GetByName(s.ctx, deleted.Name) + s.Assert().ErrorIs(err, user.ErrNotExist) + + _, err = s.repository.GetByEmail(s.ctx, deleted.Email) + s.Assert().ErrorIs(err, user.ErrNotExist) + + byIDs, err := s.repository.GetByIDs(s.ctx, []string{deleted.ID}) + s.Assert().NoError(err) + s.Assert().Empty(byIDs) + + got, err := s.repository.List(s.ctx, user.Filter{}) + s.Assert().NoError(err) + s.Assert().Len(got, len(s.users)-1) + for _, u := range got { + s.Assert().NotEqual(deleted.ID, u.ID) + } + + _, err = s.repository.UpdateByEmail(s.ctx, user.User{Email: deleted.Email, Title: "changed"}) + s.Assert().ErrorIs(err, user.ErrNotExist) + + _, err = s.repository.UpdateByName(s.ctx, user.User{Name: deleted.Name, Title: "changed"}) + s.Assert().ErrorIs(err, user.ErrNotExist) + + byID := deleted + byID.Title = "changed" + _, err = s.repository.UpdateByID(s.ctx, byID) + s.Assert().ErrorIs(err, user.ErrNotExist) + + err = s.repository.SetState(s.ctx, deleted.ID, user.Disabled) + s.Assert().ErrorIs(err, user.ErrNotExist) +}