Skip to content
Open
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
26 changes: 26 additions & 0 deletions internal/store/postgres/policy_repository.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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,
Expand All @@ -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
}
Expand All @@ -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):
Expand Down
106 changes: 96 additions & 10 deletions internal/store/postgres/policy_repository_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,9 @@ package postgres_test

import (
"context"
"errors"
"fmt"
"testing"
"time"

"github.com/raystack/frontier/core/role"

Expand Down Expand Up @@ -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,
},
Expand All @@ -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 == "" {
Expand Down Expand Up @@ -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
Comment on lines +314 to +315

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🎯 Functional Correctness | 🟡 Minor | ⚡ Quick win

Count waits on this transaction, not all database waits.

If another session waits for a lock in the same database, waiting == 1 can pass before this policy upsert blocks or fail when both sessions block. Record the deletion transaction’s backend PID. Then check whether the policy upsert is blocked by that PID, for example with pg_blocking_pids. (postgresql.org)

}, 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() {
Expand Down
Loading