diff --git a/pkg/stream/config.go b/pkg/stream/config.go index 7fb77aa5..cfbe55e4 100644 --- a/pkg/stream/config.go +++ b/pkg/stream/config.go @@ -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 diff --git a/pkg/stream/config_test.go b/pkg/stream/config_test.go index 2b08a4d9..dab20b10 100644 --- a/pkg/stream/config_test.go +++ b/pkg/stream/config_test.go @@ -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() diff --git a/pkg/stream/preflight/access.go b/pkg/stream/preflight/access.go index 3810d50b..22ee0464 100644 --- a/pkg/stream/preflight/access.go +++ b/pkg/stream/preflight/access.go @@ -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, @@ -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 { @@ -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 @@ -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, + ) +} diff --git a/pkg/stream/preflight/access_test.go b/pkg/stream/preflight/access_test.go index 65706618..5df89f90 100644 --- a/pkg/stream/preflight/access_test.go +++ b/pkg/stream/preflight/access_test.go @@ -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() @@ -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() @@ -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)) @@ -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{ diff --git a/pkg/stream/preflight/builder.go b/pkg/stream/preflight/builder.go index 789ce668..4e9f9b57 100644 --- a/pkg/stream/preflight/builder.go +++ b/pkg/stream/preflight/builder.go @@ -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) }