diff --git a/internal/store/postgres/policy_repository.go b/internal/store/postgres/policy_repository.go index 47171750d..bf80a229c 100644 --- a/internal/store/postgres/policy_repository.go +++ b/internal/store/postgres/policy_repository.go @@ -11,10 +11,12 @@ import ( "time" "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" "github.com/jmoiron/sqlx" "github.com/lib/pq" "github.com/raystack/frontier/core/namespace" "github.com/raystack/frontier/core/policy" + "github.com/raystack/frontier/core/role" "github.com/raystack/frontier/internal/bootstrap/schema" "github.com/raystack/frontier/pkg/auditrecord" "github.com/raystack/frontier/pkg/db" @@ -210,6 +212,15 @@ func (r PolicyRepository) Upsert(ctx context.Context, pol policy.Policy) (policy return policy.Policy{}, fmt.Errorf("%w: %w", errParse, err) } + lockQuery, lockParams, err := fromLive(TABLE_ROLES). + Select("id"). + Where(goqu.Ex{"id": pol.RoleID}). + ForKeyShare(exp.Wait). + ToSQL() + if err != nil { + return policy.Policy{}, fmt.Errorf("%w: %w", errQuery, err) + } + query, params, err := dialect.Insert(TABLE_POLICIES).Rows( goqu.Record{ "role_id": pol.RoleID, @@ -234,6 +245,16 @@ func (r PolicyRepository) Upsert(ctx context.Context, pol policy.Policy) (policy var policyDB Policy if err = r.dbc.WithTxn(ctx, sql.TxOptions{}, func(tx *sqlx.Tx) error { return r.dbc.WithTimeout(ctx, TABLE_POLICIES, "Upsert", func(ctx context.Context) error { + // A soft-deleted role keeps its row, so the foreign key no longer proves the + // role is live. This lock waits for a running role delete but not for other policy creates. + var roleID string + if err := tx.QueryRowContext(ctx, lockQuery, lockParams...).Scan(&roleID); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return role.ErrNotExist + } + return err + } + if err := tx.QueryRowxContext(ctx, query, params...).StructScan(&policyDB); err != nil { return err } @@ -258,6 +279,11 @@ func (r PolicyRepository) Upsert(ctx context.Context, pol policy.Policy) (policy return InsertAuditRecordInTx(ctx, tx, auditRecord) }) }); err != nil { + // WithTxn wraps every error it rolled back on as "rollback: ...". Here + // the rollback is the normal path, so return the plain sentinel. + if errors.Is(err, role.ErrNotExist) { + return policy.Policy{}, role.ErrNotExist + } err = checkPostgresError(err) switch { case errors.Is(err, ErrForeignKeyViolation): diff --git a/internal/store/postgres/policy_repository_test.go b/internal/store/postgres/policy_repository_test.go index 8d36aa835..124b5ca55 100644 --- a/internal/store/postgres/policy_repository_test.go +++ b/internal/store/postgres/policy_repository_test.go @@ -2,9 +2,9 @@ package postgres_test import ( "context" - "errors" "fmt" "testing" + "time" "github.com/raystack/frontier/core/role" @@ -172,18 +172,24 @@ func (s *PolicyRepositoryTestSuite) TestCreate() { }, }, { - Description: "should return error if role id does not exist", + Description: "should return not found if the role does not exist", PolicyToCreate: policy.Policy{ - RoleID: "role2-random", - ResourceType: "ns1", + RoleID: uuid.NewString(), + ResourceID: uuid.NewString(), + ResourceType: "ns1", + PrincipalID: s.userID, + PrincipalType: schema.UserPrincipal, }, - Err: policy.ErrInvalidDetail, + Err: role.ErrNotExist, }, { Description: "should return error if namespace id does not exist", PolicyToCreate: policy.Policy{ - RoleID: s.roles[0].ID, - ResourceType: "ns1-random", + RoleID: s.roles[0].ID, + ResourceID: uuid.NewString(), + ResourceType: "ns1-random", + PrincipalID: s.userID, + PrincipalType: schema.UserPrincipal, }, Err: policy.ErrInvalidDetail, }, @@ -193,9 +199,7 @@ func (s *PolicyRepositoryTestSuite) TestCreate() { s.Run(tc.Description, func() { got, err := s.repository.Upsert(s.ctx, tc.PolicyToCreate) if tc.Err != nil { - if errors.Is(tc.Err, err) { - s.T().Fatalf("got error %s, expected was %s", err.Error(), tc.Err.Error()) - } + s.Assert().ErrorIs(err, tc.Err) } else { s.Assert().NoError(err) if got.ID == "" { @@ -251,6 +255,88 @@ func (s *PolicyRepositoryTestSuite) TestCreate() { s.Assert().Equal(metadata.Metadata{"team": "maps"}, got.Metadata) s.Assert().True(got.UpdatedAt.After(before.UpdatedAt)) }) + + newRole := func(name string) role.Role { + created, err := postgres.NewRoleRepository(s.client).Upsert(s.ctx, role.Role{ + Name: name, + OrgID: s.orgID, + Metadata: metadata.Metadata{}, + }) + s.Require().NoError(err) + return created + } + + s.Run("should return not found when the role is soft-deleted", func() { + deleted := newRole("soft-deleted role for a policy create") + _, err := s.client.ExecContext(s.ctx, "UPDATE roles SET deleted_at = now() WHERE id = $1", deleted.ID) + s.Require().NoError(err) + + _, err = s.repository.Upsert(s.ctx, policy.Policy{ + RoleID: deleted.ID, + ResourceID: uuid.NewString(), + ResourceType: "ns1", + PrincipalID: uuid.NewString(), + PrincipalType: schema.UserPrincipal, + }) + s.Assert().ErrorIs(err, role.ErrNotExist) + + var created int + err = s.client.QueryRowxContext(s.ctx, "SELECT count(*) FROM policies WHERE role_id = $1", deleted.ID).Scan(&created) + s.Require().NoError(err) + s.Assert().Equal(0, created) + }) + + s.Run("should wait for a role delete that has not committed and then return not found", func() { + target := newRole("role with a running delete") + tx, err := s.client.BeginTxx(s.ctx, nil) + s.Require().NoError(err) + defer tx.Rollback() // nolint + var lockedID string + err = tx.QueryRowContext(s.ctx, "SELECT id FROM roles WHERE id = $1 AND deleted_at IS NULL FOR UPDATE", target.ID).Scan(&lockedID) + s.Require().NoError(err) + _, err = tx.ExecContext(s.ctx, "UPDATE roles SET deleted_at = now() WHERE id = $1", target.ID) + s.Require().NoError(err) + + done := make(chan error, 1) + go func() { + _, err := s.repository.Upsert(s.ctx, policy.Policy{ + RoleID: target.ID, + ResourceID: uuid.NewString(), + ResourceType: "ns1", + PrincipalID: uuid.NewString(), + PrincipalType: schema.UserPrincipal, + }) + done <- err + }() + + s.Require().Eventually(func() bool { + var waiting int + err := s.client.QueryRowxContext(s.ctx, "SELECT count(*) FROM pg_stat_activity WHERE datname = current_database() AND wait_event_type = 'Lock'").Scan(&waiting) + return err == nil && waiting == 1 + }, time.Second, 10*time.Millisecond, "the create did not wait for the role delete") + + s.Require().NoError(tx.Commit()) + s.Assert().ErrorIs(<-done, role.ErrNotExist) + }) + + s.Run("should not wait for another policy create on the same role", func() { + target := newRole("role with two policy creates") + tx, err := s.client.BeginTxx(s.ctx, nil) + s.Require().NoError(err) + defer tx.Rollback() // nolint + _, err = tx.ExecContext(s.ctx, "INSERT INTO policies (role_id, resource_id, resource_type, principal_id, principal_type) VALUES ($1, $2, 'ns1', $3, 'app/user')", + target.ID, s.orgID, uuid.NewString()) + s.Require().NoError(err) + + _, err = s.repository.Upsert(s.ctx, policy.Policy{ + RoleID: target.ID, + ResourceID: s.orgID, + ResourceType: "ns1", + PrincipalID: uuid.NewString(), + PrincipalType: schema.UserPrincipal, + }) + s.Assert().NoError(err) + }) } func (s *PolicyRepositoryTestSuite) TestList() {