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
11 changes: 11 additions & 0 deletions pkg/stream/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -224,6 +224,17 @@ func (c *Config) SnapshotCreateTargetDB() bool {
return c.Listener.Postgres.Snapshot.Schema.DumpRestore.CreateTargetDB
}

func (c *Config) SnapshotRestoresRoles() bool {
if c.Listener.Postgres == nil ||
c.Listener.Postgres.Snapshot == nil ||
c.Listener.Postgres.Snapshot.Schema == nil ||
c.Listener.Postgres.Snapshot.Schema.DumpRestore == nil {
return false
}
mode := c.Listener.Postgres.Snapshot.Schema.DumpRestore.RolesSnapshotMode
return mode == "enabled" || mode == "no_passwords"
}

// restoreConflictTargetsBeforeData reports whether the schema snapshot must
// restore primary keys, unique constraints and unique indexes before the data
// snapshot runs. This is required when the postgres batch writer emits
Expand Down
50 changes: 50 additions & 0 deletions pkg/stream/config_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,12 +7,62 @@ import (

"github.com/stretchr/testify/require"
pgsnapshotgenerator "github.com/xataio/pgstream/pkg/snapshot/generator/postgres/data"
pgdumprestore "github.com/xataio/pgstream/pkg/snapshot/generator/postgres/schema/pgdumprestore"
"github.com/xataio/pgstream/pkg/wal/listener/snapshot/adapter"
snapshotbuilder "github.com/xataio/pgstream/pkg/wal/listener/snapshot/builder"
"github.com/xataio/pgstream/pkg/wal/processor/filter"
pgreplication "github.com/xataio/pgstream/pkg/wal/replication/postgres"
)

func TestConfig_SnapshotRestoresRoles(t *testing.T) {
t.Parallel()

tests := []struct {
name string
cfg *Config
want bool
}{
{name: "no postgres listener", cfg: &Config{}},
{
name: "no snapshot configuration",
cfg: &Config{Listener: ListenerConfig{Postgres: &PostgresListenerConfig{}}},
},
{
name: "enabled",
cfg: configWithRolesSnapshotMode("enabled"),
want: true,
},
{
name: "no passwords",
cfg: configWithRolesSnapshotMode("no_passwords"),
want: true,
},
{name: "disabled", cfg: configWithRolesSnapshotMode("disabled")},
{name: "empty mode", cfg: configWithRolesSnapshotMode("")},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
require.Equal(t, tc.want, tc.cfg.SnapshotRestoresRoles())
})
}
}

func configWithRolesSnapshotMode(mode string) *Config {
return &Config{
Listener: ListenerConfig{
Postgres: &PostgresListenerConfig{
Snapshot: &snapshotbuilder.SnapshotListenerConfig{
Schema: &snapshotbuilder.SchemaSnapshotConfig{
DumpRestore: &pgdumprestore.Config{RolesSnapshotMode: mode},
},
},
},
},
}
}

func TestConfig_ReplicationTableSelection(t *testing.T) {
t.Parallel()

Expand Down
46 changes: 46 additions & 0 deletions pkg/stream/preflight/access.go
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,14 @@ func (c *TargetCreateDBPrivilegeCheck) Name() string {
return "target_createdb_privilege"
}

type TargetCreateRolePrivilegeCheck struct {
Target postgres.AcquireFunc
}

func (c *TargetCreateRolePrivilegeCheck) Name() string {
return "target_createrole_privilege"
}

const sourceTableSelectPrivilegesQuery = `
SELECT
current_user,
Expand Down Expand Up @@ -110,6 +118,14 @@ FROM pg_roles
WHERE rolname = current_user
`

const targetCreateRolePrivilegeQuery = `
SELECT
current_user,
rolcreaterole OR rolsuper
FROM pg_roles
WHERE rolname = current_user
`

func (c *SourceTableSelectPrivilegesCheck) Run(ctx context.Context) ([]Finding, error) {
conn, err := c.Source(ctx)
if err != nil {
Expand Down Expand Up @@ -197,6 +213,29 @@ func (c *TargetCreateDBPrivilegeCheck) Run(ctx context.Context) ([]Finding, erro
}, nil
}

func (c *TargetCreateRolePrivilegeCheck) Run(ctx context.Context) ([]Finding, error) {
conn, err := c.Target(ctx)
if err != nil {
return nil, fmt.Errorf("connecting to target: %w", err)
}
var role string
var hasCreateRole bool

if err := conn.QueryRow(
ctx,
[]any{&role, &hasCreateRole},
targetCreateRolePrivilegeQuery,
); err != nil {
return nil, fmt.Errorf("querying target CREATEROLE privilege: %w", err)
}
if hasCreateRole {
return nil, nil
}
return []Finding{{
Message: targetCreateRolePrivilegeMessage(role),
}}, nil
}

type sourceTableSelectPrivilegeRow struct {
Role string
Schema string
Expand Down Expand Up @@ -284,3 +323,10 @@ func targetCreateDBPrivilegeMessage(role string) string {
role, quotedRole,
)
}

func targetCreateRolePrivilegeMessage(role string) string {
quotedRole := postgres.QuoteIdentifier(role)
return fmt.Sprintf(
"target role %q lacks CREATEROLE; run ALTER ROLE %s CREATEROLE", role, quotedRole,
)
}
178 changes: 178 additions & 0 deletions pkg/stream/preflight/access_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -503,6 +503,106 @@ func TestTargetCreateDBPrivilegeCheck_Name(t *testing.T) {
require.Equal(t, "target_createdb_privilege", (&TargetCreateDBPrivilegeCheck{}).Name())
}

func TestTargetCreateRolePrivilegeCheck_Run_HasCreateRole(t *testing.T) {
t.Parallel()

check := &TargetCreateRolePrivilegeCheck{
Target: targetWithCreateRolePrivilege(t, "pgstreamtarget", true),
}

findings, err := check.Run(context.Background())

require.NoError(t, err)
require.Empty(t, findings)
}

func TestTargetCreateRolePrivilegeCheck_Run_ChecksSuperuserCapability(t *testing.T) {
t.Parallel()

check := &TargetCreateRolePrivilegeCheck{
Target: func(context.Context) (postgres.Querier, error) {
return &mocks.Querier{
QueryRowFn: func(_ context.Context, dest []any, query string, _ ...any) error {
require.Contains(t, query, "rolcreaterole OR rolsuper")
require.Len(t, dest, 2)
role, ok := dest[0].(*string)
require.True(t, ok)
hasCreateRole, ok := dest[1].(*bool)
require.True(t, ok)
*role = "postgres"
*hasCreateRole = true
return nil
},
}, nil
},
}

findings, err := check.Run(context.Background())

require.NoError(t, err)
require.Empty(t, findings)
}

func TestTargetCreateRolePrivilegeCheck_Run_MissingCreateRoleReturnsFinding(t *testing.T) {
t.Parallel()

check := &TargetCreateRolePrivilegeCheck{
Target: targetWithCreateRolePrivilege(t, "pgstreamtarget", false),
}

findings, err := check.Run(context.Background())

require.NoError(t, err)
require.Len(t, findings, 1)
require.Contains(t, findings[0].Message, `target role "pgstreamtarget"`)
require.Contains(t, findings[0].Message, "lacks CREATEROLE")
require.Contains(t, findings[0].Message, `ALTER ROLE "pgstreamtarget" CREATEROLE`)
}

func TestTargetCreateRolePrivilegeCheck_Run_TargetAcquireFails(t *testing.T) {
t.Parallel()

checkErr := errors.New("boom")
check := &TargetCreateRolePrivilegeCheck{
Target: func(context.Context) (postgres.Querier, error) {
return nil, checkErr
},
}

findings, err := check.Run(context.Background())

require.Nil(t, findings)
require.ErrorIs(t, err, checkErr)
require.ErrorContains(t, err, "connecting to target")
}

func TestTargetCreateRolePrivilegeCheck_Run_QueryFails(t *testing.T) {
t.Parallel()

queryErr := errors.New("query failed")
check := &TargetCreateRolePrivilegeCheck{
Target: func(context.Context) (postgres.Querier, error) {
return &mocks.Querier{
QueryRowFn: func(context.Context, []any, string, ...any) error {
return queryErr
},
}, nil
},
}

findings, err := check.Run(context.Background())

require.Nil(t, findings)
require.ErrorIs(t, err, queryErr)
require.ErrorContains(t, err, "querying target CREATEROLE privilege")
}

func TestTargetCreateRolePrivilegeCheck_Name(t *testing.T) {
t.Parallel()

require.Equal(t, "target_createrole_privilege", (&TargetCreateRolePrivilegeCheck{}).Name())
}

func TestSourceTableSelectPrivilegeMessage(t *testing.T) {
t.Parallel()

Expand Down Expand Up @@ -557,6 +657,16 @@ func TestTargetCreateDBPrivilegeMessage(t *testing.T) {
require.Contains(t, msg, `ALTER ROLE "pgstreamtarget" CREATEDB`)
}

func TestTargetCreateRolePrivilegeMessage(t *testing.T) {
t.Parallel()

msg := targetCreateRolePrivilegeMessage("pgstreamtarget")

require.Contains(t, msg, `target role "pgstreamtarget"`)
require.Contains(t, msg, "lacks CREATEROLE")
require.Contains(t, msg, `ALTER ROLE "pgstreamtarget" CREATEROLE`)
}

func TestBuildAccessChecks(t *testing.T) {
t.Parallel()

Expand Down Expand Up @@ -670,6 +780,55 @@ func TestBuildChecks_SelectedAccessOnly(t *testing.T) {
require.Equal(t, "source_sequence_select_privileges", checks[1].Name())
}

func TestBuildAccessChecks_RolesSnapshotMode(t *testing.T) {
t.Parallel()

tests := []struct {
name string
mode string
wantCheck bool
}{
{name: "enabled adds target createrole check", mode: "enabled", wantCheck: true},
{name: "no passwords adds target createrole check", mode: "no_passwords", wantCheck: true},
{name: "disabled omits target createrole check", mode: "disabled"},
{name: "empty mode omits target createrole check"},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()

cfg := &stream.Config{
Listener: stream.ListenerConfig{
Postgres: &stream.PostgresListenerConfig{
URL: "postgres://source",
Snapshot: &snapshotbuilder.SnapshotListenerConfig{
Schema: &snapshotbuilder.SchemaSnapshotConfig{
DumpRestore: &pgdumprestore.Config{
TargetPGURL: "postgres://target",
RolesSnapshotMode: tc.mode,
},
},
},
},
},
}

checks, cleanup := BuildAccessChecks(cfg)
require.NotNil(t, cleanup)
defer func() { require.NoError(t, cleanup(context.Background())) }()

var found bool
for _, check := range checks {
if _, ok := check.(*TargetCreateRolePrivilegeCheck); ok {
found = true
}
}
require.Equal(t, tc.wantCheck, found)
})
}
}

func sourceWithRows(t *testing.T, rows []sourceTableSelectPrivilegeRow) postgres.AcquireFunc {
t.Helper()
return sourceWithMockRows(privilegeRows(t, rows))
Expand Down Expand Up @@ -735,6 +894,25 @@ func targetWithCreateDBPrivilege(t *testing.T, role string, hasCreateDB bool) po
}
}

func targetWithCreateRolePrivilege(t *testing.T, role string, hasCreateRole bool) postgres.AcquireFunc {
t.Helper()
return func(context.Context) (postgres.Querier, error) {
return &mocks.Querier{
QueryRowFn: func(_ context.Context, dest []any, query string, _ ...any) error {
require.Contains(t, query, "rolcreaterole OR rolsuper")
require.Len(t, dest, 2)
roleDest, ok := dest[0].(*string)
require.True(t, ok)
createRoleDest, ok := dest[1].(*bool)
require.True(t, ok)
*roleDest = role
*createRoleDest = hasCreateRole
return nil
},
}, nil
}
}

func sequencePrivilegeRows(t *testing.T, rows []sourceSequenceSelectPrivilegeRow) postgres.Rows {
t.Helper()
return &mocks.Rows{
Expand Down
18 changes: 18 additions & 0 deletions pkg/stream/preflight/builder.go
Original file line number Diff line number Diff line change
Expand Up @@ -154,6 +154,24 @@ func BuildAccessChecks(cfg *stream.Config) ([]Check, CleanupFunc) {
}
}

if cfg.SnapshotRestoresRoles() {
if targetURL := cfg.SnapshotTargetPostgresURL(); targetURL != "" {
var err error
targetURL, err = postgres.RemoveDatabaseFromConnectionString(targetURL)
if err != nil {
checks = append(checks, &TargetCreateRolePrivilegeCheck{
Target: func(context.Context) (postgres.Querier, error) {
return nil, err
},
})
return checks, joinCleanups(cleanups)
}
target := postgres.NewLazyConn(targetURL)
checks = append(checks, &TargetCreateRolePrivilegeCheck{Target: target.Acquire})
cleanups = append(cleanups, target.Close)
}
}

return checks, joinCleanups(cleanups)
}

Expand Down