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
17 changes: 9 additions & 8 deletions internal/store/postgres/organization_repository.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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 {
Expand All @@ -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)
}
Expand Down Expand Up @@ -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 {
Expand All @@ -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)
}
Expand Down Expand Up @@ -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)
Expand Down
36 changes: 36 additions & 0 deletions internal/store/postgres/organization_repository_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
12 changes: 12 additions & 0 deletions internal/store/postgres/postgres.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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"
Expand Down
22 changes: 10 additions & 12 deletions internal/store/postgres/project_repository.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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
}
Expand Down
32 changes: 32 additions & 0 deletions internal/store/postgres/project_repository_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
18 changes: 11 additions & 7 deletions internal/store/postgres/user_repository.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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),
Expand Down Expand Up @@ -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"),
Expand Down
Loading
Loading