diff --git a/database/database.go b/database/database.go index 67543936..9a6f32ea 100644 --- a/database/database.go +++ b/database/database.go @@ -36,11 +36,15 @@ func New(ctx context.Context, logger log.Logger, config DatabaseConfig) (*sql.DB } return db, nil } else if config.Postgres != nil { + // Pool settings are applied to pgxpool inside postgresConnection. + // Do not call ApplyConnectionsConfig on the returned *sql.DB: + // OpenDBFromPool requires MaxIdleConns=0, and sql.DB setters do not + // configure the underlying pgxpool. db, err := postgresConnection(ctx, logger, *config.Postgres, config.DatabaseName) if err != nil { return nil, fmt.Errorf("connecting to postgres: %w", err) } - return ApplyConnectionsConfig(db, &config.Postgres.Connections, logger), nil + return db, nil } return nil, ErrMissingConfig @@ -112,3 +116,17 @@ func ApplyConnectionsConfig(db *sql.DB, connections *ConnectionsConfig, logger l return db } + +// ApplyPostgresConnectionsConfig applies connection pool settings with safe defaults +// onto a *sql.DB. +// +// Deprecated: Postgres connections from New use pgxpool under the hood. Pool +// settings are applied via ApplyPostgresPoolConfig inside postgresConnection. +// Calling this on a Postgres *sql.DB from New is incorrect: SetMaxIdleConns +// with a non-zero value breaks OpenDBFromPool, and the other setters do not +// configure the underlying pgxpool. Prefer ConnectionsConfig on PostgresConfig +// (applied automatically) or ApplyPostgresPoolConfig when building a pool. +func ApplyPostgresConnectionsConfig(db *sql.DB, connections *ConnectionsConfig, logger log.Logger) *sql.DB { + applied := ResolvePostgresConnectionsConfig(*connections) + return ApplyConnectionsConfig(db, &applied, logger) +} diff --git a/database/database_test.go b/database/database_test.go index a648aa02..92baeee5 100644 --- a/database/database_test.go +++ b/database/database_test.go @@ -8,6 +8,7 @@ import ( "errors" "os" "testing" + "time" gomysql "github.com/go-sql-driver/mysql" "github.com/jackc/pgx/v5/pgconn" @@ -121,6 +122,33 @@ func TestDataTooLong(t *testing.T) { } } +func TestResolvePostgresConnectionsConfig_Defaults(t *testing.T) { + defaults := database.DefaultPostgresConnectionsConfig() + got := database.ResolvePostgresConnectionsConfig(database.ConnectionsConfig{}) + require.Equal(t, defaults, got) +} + +func TestResolvePostgresConnectionsConfig_Overrides(t *testing.T) { + in := database.ConnectionsConfig{ + MaxOpen: 10, + MaxIdle: 3, + MaxLifetime: time.Minute, + MaxIdleTime: time.Second * 15, + } + got := database.ResolvePostgresConnectionsConfig(in) + require.Equal(t, in, got) +} + +func TestDefaultPostgresConnectionsConfig(t *testing.T) { + defaults := database.DefaultPostgresConnectionsConfig() + require.Greater(t, defaults.MaxOpen, 0) + require.Greater(t, defaults.MaxIdle, 0) + require.Greater(t, defaults.MaxLifetime, time.Duration(0)) + require.Greater(t, defaults.MaxIdleTime, time.Duration(0)) + // MaxIdleTime should be shorter than MaxLifetime + require.Less(t, defaults.MaxIdleTime, defaults.MaxLifetime) +} + func TestConnectionsConfigOrder(t *testing.T) { bs, err := os.ReadFile("database.go") require.NoError(t, err) diff --git a/database/model_config.go b/database/model_config.go index ea8b9026..374cd7aa 100644 --- a/database/model_config.go +++ b/database/model_config.go @@ -99,6 +99,28 @@ type ConnectionsConfig struct { MaxIdleTime time.Duration } +// DefaultPostgresConnectionsConfig returns connection pool defaults tuned for +// database failover recovery (e.g. AlloyDB maintenance switchovers). +// +// These are applied by ResolvePostgresConnectionsConfig / ApplyPostgresPoolConfig +// whenever a field on ConnectionsConfig is zero. pgxpool always uses a finite +// MaxConns (unlike database/sql, where MaxOpen=0 means unlimited), so leaving +// MaxOpen unset would otherwise silently fall back to max(4, NumCPU()). +// Explicit defaults keep pool size and eviction policy predictable across +// services that never set Connections. +// +// Short MaxLifetime / MaxIdleTime help the background reaper drop stale +// connections after a primary change; acquire-time liveness is separate +// (pgxpool ShouldPing / PingTimeout). +func DefaultPostgresConnectionsConfig() ConnectionsConfig { + return ConnectionsConfig{ + MaxOpen: 25, + MaxIdle: 5, + MaxLifetime: 5 * time.Minute, + MaxIdleTime: 30 * time.Second, + } +} + type RetryConfig struct { MaxAttempts int MinDuration time.Duration diff --git a/database/postgres.go b/database/postgres.go index 7c98cdd2..09b86eac 100644 --- a/database/postgres.go +++ b/database/postgres.go @@ -3,14 +3,18 @@ package database import ( "context" "database/sql" + "database/sql/driver" "errors" "fmt" + "io" "net" "strings" + "sync" + "time" "cloud.google.com/go/alloydbconn" - "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgconn" + "github.com/jackc/pgx/v5/pgxpool" "github.com/jackc/pgx/v5/stdlib" "github.com/moov-io/base/log" ) @@ -20,70 +24,174 @@ const ( // https://www.postgresql.org/docs/current/errcodes-appendix.html postgresErrUniqueViolation = "23505" postgresErrDeadlockFound = "40P01" + + // Bound ShouldPing wait so acquire retries during failover stay inside + // typical request budgets. AlloyDB disconnects usually fail fast (TCP RST); + // this caps hung/TIME_WAIT peers. + defaultPostgresPingTimeout = time.Second ) func postgresConnection(ctx context.Context, logger log.Logger, config PostgresConfig, databaseName string) (*sql.DB, error) { - var connStr string - if config.Alloy != nil { - c, err := getAlloyDBConnectorConnStr(ctx, config, databaseName) - if err != nil { - return nil, logger.LogErrorf("creating alloydb connection: %w", err).Err() - } - connStr = c - } else { - c, err := getPostgresConnStr(config, databaseName) - if err != nil { - return nil, logger.LogErrorf("creating postgres connection: %w", err).Err() - } - connStr = c + poolConfig, dialer, err := buildPgxPoolConfig(ctx, config, databaseName) + if err != nil { + return nil, logger.LogErrorf("building pgx pool config: %w", err).Err() } - db, err := sql.Open("pgx", connStr) + // Apply connection limits to pgxpool (not database/sql). OpenDBFromPool + // requires sql.DB MaxIdleConns=0; sql.DB setters do not configure the + // underlying pool and SetMaxIdleConns(n>0) actively breaks it. + ApplyPostgresPoolConfig(logger, poolConfig, config.Connections) + + // HealthCheckPeriod is the background reaper cadence. It only evicts + // connections that exceed MaxConnLifetime / MaxConnIdleTime — it does not + // ping. Liveness at acquire time is handled by pgxpool's default + // ShouldPing (idle > 1s) unless callers customize it later. + poolConfig.HealthCheckPeriod = 1 * time.Second + + if poolConfig.PingTimeout <= 0 { + poolConfig.PingTimeout = defaultPostgresPingTimeout + } + + pool, err := pgxpool.NewWithConfig(ctx, poolConfig) if err != nil { - return nil, logger.LogErrorf("opening database: %w", err).Err() + _ = closeAlloyDialer(dialer) + return nil, logger.LogErrorf("creating pgx pool: %w", err).Err() } - err = db.Ping() + err = pool.Ping(ctx) if err != nil { - _ = db.Close() + pool.Close() + _ = closeAlloyDialer(dialer) return nil, logger.LogErrorf("connecting to database: %w", err).Err() } + // OpenDBFromPool does not close the pool when *sql.DB is closed. Wrap the + // connector so db.Close() shuts down the pool (and AlloyDB dialer). + db := openDBFromPool(pool, dialer) + return db, nil } -func getPostgresConnStr(config PostgresConfig, databaseName string) (string, error) { - url := fmt.Sprintf("postgres://%s:%s@%s/%s", config.User, config.Password, config.Address, databaseName) +// ApplyPostgresPoolConfig fills zero-valued fields in connections with +// DefaultPostgresConnectionsConfig, then maps them onto poolConfig. +// +// Unlike database/sql (where MaxOpen=0 means unlimited), pgxpool always has a +// finite MaxConns. Leaving MaxOpen unset previously fell through to pgxpool's +// default of max(4, NumCPU()), which silently shrinks pools for services that +// never configured Connections. We instead apply explicit library defaults so +// behavior is predictable and logged. +// +// MaxIdle has no pgxpool "max idle" equivalent. When set (or defaulted), it is +// applied as MinIdleConns (warm floor), capped by MaxConns, so the field still +// influences pool shape rather than being dropped on the floor. +func ApplyPostgresPoolConfig(logger log.Logger, poolConfig *pgxpool.Config, connections ConnectionsConfig) { + if poolConfig == nil { + return + } - params := "" + applied := ResolvePostgresConnectionsConfig(connections) - if config.TLS != nil { - if len(config.TLS.Mode) < 1 { - config.TLS.Mode = "verify-full" - } + logger.Logf("setting pgx pool MaxConns to %d", applied.MaxOpen) + poolConfig.MaxConns = int32(applied.MaxOpen) - params += "sslmode=" + config.TLS.Mode + minIdle := applied.MaxIdle + if minIdle > applied.MaxOpen { + minIdle = applied.MaxOpen + } + if minIdle < 0 { + minIdle = 0 + } + logger.Logf("setting pgx pool MinIdleConns to %d (from ConnectionsConfig.MaxIdle)", minIdle) + poolConfig.MinIdleConns = int32(minIdle) - if len(config.TLS.CACertFile) > 0 { - params += "&sslrootcert=" + config.TLS.CACertFile - } + logger.Logf("setting pgx pool MaxConnIdleTime to %v", applied.MaxIdleTime) + poolConfig.MaxConnIdleTime = applied.MaxIdleTime - if len(config.TLS.ClientCertFile) > 0 { - params += "&sslcert=" + config.TLS.ClientCertFile - } + logger.Logf("setting pgx pool MaxConnLifetime to %v", applied.MaxLifetime) + poolConfig.MaxConnLifetime = applied.MaxLifetime +} - if len(config.TLS.ClientKeyFile) > 0 { - params += "&sslkey=" + config.TLS.ClientKeyFile +// ResolvePostgresConnectionsConfig returns connections with zero-valued fields +// replaced by DefaultPostgresConnectionsConfig. +func ResolvePostgresConnectionsConfig(connections ConnectionsConfig) ConnectionsConfig { + defaults := DefaultPostgresConnectionsConfig() + if connections.MaxOpen <= 0 { + connections.MaxOpen = defaults.MaxOpen + } + if connections.MaxIdle <= 0 { + connections.MaxIdle = defaults.MaxIdle + } + if connections.MaxLifetime <= 0 { + connections.MaxLifetime = defaults.MaxLifetime + } + if connections.MaxIdleTime <= 0 { + connections.MaxIdleTime = defaults.MaxIdleTime + } + return connections +} + +// openDBFromPool wraps pgxpool in a *sql.DB whose Close also closes the pool +// and optional AlloyDB dialer. stdlib.OpenDBFromPool alone leaks both. +func openDBFromPool(pool *pgxpool.Pool, dialer *alloydbconn.Dialer) *sql.DB { + c := &poolConnector{ + Connector: stdlib.GetPoolConnector(pool), + pool: pool, + dialer: dialer, + } + db := sql.OpenDB(c) + // Required when using a pgxpool-backed connector: non-zero idle conns on + // sql.DB prevent connections from being released back to the pool. + db.SetMaxIdleConns(0) + return db +} + +// poolConnector delegates to pgx stdlib's pool connector and implements +// io.Closer so database/sql.DB.Close shuts down the underlying pgxpool. +type poolConnector struct { + driver.Connector + pool *pgxpool.Pool + dialer *alloydbconn.Dialer + + closeOnce sync.Once + closeErr error +} + +func (c *poolConnector) Close() error { + c.closeOnce.Do(func() { + if c.pool != nil { + c.pool.Close() } + c.closeErr = closeAlloyDialer(c.dialer) + }) + return c.closeErr +} + +func closeAlloyDialer(dialer *alloydbconn.Dialer) error { + if dialer == nil { + return nil } + return dialer.Close() +} - connStr := fmt.Sprintf("%s?%s", url, params) - return connStr, nil +func buildPgxPoolConfig(ctx context.Context, config PostgresConfig, databaseName string) (*pgxpool.Config, *alloydbconn.Dialer, error) { + if config.Alloy != nil { + return buildAlloyDBPoolConfig(ctx, config, databaseName) + } + + connStr, err := getPostgresConnStr(config, databaseName) + if err != nil { + return nil, nil, err + } + poolConfig, err := pgxpool.ParseConfig(connStr) + if err != nil { + return nil, nil, err + } + return poolConfig, nil, nil } -func getAlloyDBConnectorConnStr(ctx context.Context, config PostgresConfig, databaseName string) (string, error) { +func buildAlloyDBPoolConfig(ctx context.Context, config PostgresConfig, databaseName string) (*pgxpool.Config, *alloydbconn.Dialer, error) { if config.Alloy == nil { - return "", fmt.Errorf("missing alloy config") + return nil, nil, fmt.Errorf("missing alloy config") } var dialer *alloydbconn.Dialer @@ -92,7 +200,7 @@ func getAlloyDBConnectorConnStr(ctx context.Context, config PostgresConfig, data if config.Alloy.UseIAM { d, err := alloydbconn.NewDialer(ctx, alloydbconn.WithIAMAuthN()) if err != nil { - return "", fmt.Errorf("creating alloydb dialer: %v", err) + return nil, nil, fmt.Errorf("creating alloydb dialer: %w", err) } dialer = d dsn = fmt.Sprintf( @@ -104,7 +212,7 @@ func getAlloyDBConnectorConnStr(ctx context.Context, config PostgresConfig, data } else { d, err := alloydbconn.NewDialer(ctx) if err != nil { - return "", fmt.Errorf("creating alloydb dialer: %v", err) + return nil, nil, fmt.Errorf("creating alloydb dialer: %w", err) } dialer = d dsn = fmt.Sprintf( @@ -114,12 +222,10 @@ func getAlloyDBConnectorConnStr(ctx context.Context, config PostgresConfig, data ) } - // TODO - //cleanup := func() error { return d.Close() } - - connConfig, err := pgx.ParseConfig(dsn) + poolConfig, err := pgxpool.ParseConfig(dsn) if err != nil { - return "", fmt.Errorf("failed to parse pgx config: %v", err) + _ = closeAlloyDialer(dialer) + return nil, nil, fmt.Errorf("failed to parse pgx pool config: %w", err) } var connOptions []alloydbconn.DialOption @@ -127,11 +233,39 @@ func getAlloyDBConnectorConnStr(ctx context.Context, config PostgresConfig, data connOptions = append(connOptions, alloydbconn.WithPSC()) } - connConfig.DialFunc = func(ctx context.Context, _ string, _ string) (net.Conn, error) { + poolConfig.ConnConfig.DialFunc = func(ctx context.Context, _ string, _ string) (net.Conn, error) { return dialer.Dial(ctx, config.Alloy.InstanceURI, connOptions...) } - connStr := stdlib.RegisterConnConfig(connConfig) + return poolConfig, dialer, nil +} + +func getPostgresConnStr(config PostgresConfig, databaseName string) (string, error) { + url := fmt.Sprintf("postgres://%s:%s@%s/%s", config.User, config.Password, config.Address, databaseName) + + params := "" + + if config.TLS != nil { + if len(config.TLS.Mode) < 1 { + config.TLS.Mode = "verify-full" + } + + params += "sslmode=" + config.TLS.Mode + + if len(config.TLS.CACertFile) > 0 { + params += "&sslrootcert=" + config.TLS.CACertFile + } + + if len(config.TLS.ClientCertFile) > 0 { + params += "&sslcert=" + config.TLS.ClientCertFile + } + + if len(config.TLS.ClientKeyFile) > 0 { + params += "&sslkey=" + config.TLS.ClientKeyFile + } + } + + connStr := fmt.Sprintf("%s?%s", url, params) return connStr, nil } @@ -164,3 +298,77 @@ func PostgresDeadlockFound(err error) bool { return strings.Contains(err.Error(), postgresErrDeadlockFound) } + +// IsRetryablePostgresError returns true if the error is a transient connection-level +// error that is safe to retry. This covers the errors seen during AlloyDB maintenance +// switchovers and other transient network failures. +func IsRetryablePostgresError(err error) bool { + if err == nil { + return false + } + + // PostgreSQL error codes indicating the server is shutting down or unavailable + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) { + switch pgErr.Code { + case "57P01", "57P02", "57P03": // admin_shutdown, crash_shutdown, cannot_connect_now + return true + case "08000", "08001", "08003", "08004", "08006": // connection_exception class + return true + } + return false + } + + // Network-level errors: connection reset, broken pipe, EOF, etc. + // These occur when the TCP connection is severed during a switchover. + var netErr *net.OpError + if errors.As(err, &netErr) { + return true + } + if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) { + return true + } + if errors.Is(err, context.DeadlineExceeded) { + return false // don't retry if the caller's context timed out + } + + // pgx wraps connection errors with these messages + msg := err.Error() + if strings.Contains(msg, "connection reset by peer") || + strings.Contains(msg, "broken pipe") || + strings.Contains(msg, "connection refused") || + strings.Contains(msg, "unexpected EOF") || + strings.Contains(msg, "conn closed") { + return true + } + + return false +} + +// RetryPostgres executes fn up to maxAttempts times, retrying on transient +// connection errors. This is intended for use around individual database +// operations to survive brief outages like AlloyDB maintenance switchovers. +func RetryPostgres(ctx context.Context, maxAttempts int, fn func() error) error { + if maxAttempts <= 0 { + maxAttempts = 3 + } + var err error + for attempt := 0; attempt < maxAttempts; attempt++ { + err = fn() + if err == nil { + return nil + } + if !IsRetryablePostgresError(err) { + return err + } + if attempt < maxAttempts-1 { + backoff := time.Duration(attempt+1) * 200 * time.Millisecond + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(backoff): + } + } + } + return err +} diff --git a/database/postgres_pool_test.go b/database/postgres_pool_test.go new file mode 100644 index 00000000..d6839d8c --- /dev/null +++ b/database/postgres_pool_test.go @@ -0,0 +1,44 @@ +package database + +import ( + "context" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + "github.com/stretchr/testify/require" +) + +func TestOpenDBFromPool_CloseClosesPool(t *testing.T) { + // NewWithConfig does not dial until a connection is needed; use an + // unreachable address so this stays a pure unit test. + cfg, err := pgxpool.ParseConfig("postgres://user:pass@127.0.0.1:1/db?sslmode=disable") + require.NoError(t, err) + cfg.MaxConns = 2 + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + pool, err := pgxpool.NewWithConfig(ctx, cfg) + require.NoError(t, err) + + db := openDBFromPool(pool, nil) + require.NoError(t, db.Close()) + + // Closing the sql.DB must close the underlying pgxpool. + _, err = pool.Acquire(context.Background()) + require.Error(t, err) + require.Contains(t, err.Error(), "closed") +} + +func TestOpenDBFromPool_CloseIdempotent(t *testing.T) { + cfg, err := pgxpool.ParseConfig("postgres://user:pass@127.0.0.1:1/db?sslmode=disable") + require.NoError(t, err) + + pool, err := pgxpool.NewWithConfig(context.Background(), cfg) + require.NoError(t, err) + + db := openDBFromPool(pool, nil) + require.NoError(t, db.Close()) + require.NoError(t, db.Close()) +} diff --git a/database/postgres_test.go b/database/postgres_test.go index 68dcb14d..6a65014c 100644 --- a/database/postgres_test.go +++ b/database/postgres_test.go @@ -2,11 +2,16 @@ package database_test import ( "context" + "errors" + "io" + "net" "os" "path/filepath" "testing" "time" + "github.com/jackc/pgx/v5/pgconn" + "github.com/jackc/pgx/v5/pgxpool" "github.com/moov-io/base" "github.com/moov-io/base/database" "github.com/moov-io/base/database/testdb" @@ -180,6 +185,110 @@ func Test_Postgres_Alloy_Migrations(t *testing.T) { defer db.Close() } +func TestIsRetryablePostgresError(t *testing.T) { + // nil error is not retryable + require.False(t, database.IsRetryablePostgresError(nil)) + + // admin_shutdown is retryable (seen during AlloyDB maintenance) + require.True(t, database.IsRetryablePostgresError(&pgconn.PgError{Code: "57P01"})) + + // crash_shutdown is retryable + require.True(t, database.IsRetryablePostgresError(&pgconn.PgError{Code: "57P02"})) + + // cannot_connect_now is retryable + require.True(t, database.IsRetryablePostgresError(&pgconn.PgError{Code: "57P03"})) + + // connection_exception class is retryable + require.True(t, database.IsRetryablePostgresError(&pgconn.PgError{Code: "08006"})) + + // unique_violation is NOT retryable (application-level error) + require.False(t, database.IsRetryablePostgresError(&pgconn.PgError{Code: "23505"})) + + // syntax_error is NOT retryable + require.False(t, database.IsRetryablePostgresError(&pgconn.PgError{Code: "42601"})) + + // EOF is retryable (connection severed) + require.True(t, database.IsRetryablePostgresError(io.EOF)) + require.True(t, database.IsRetryablePostgresError(io.ErrUnexpectedEOF)) + + // net.OpError is retryable + require.True(t, database.IsRetryablePostgresError(&net.OpError{ + Op: "read", + Err: errors.New("connection reset by peer"), + })) + + // context.DeadlineExceeded is NOT retryable + require.False(t, database.IsRetryablePostgresError(context.DeadlineExceeded)) + + // String-matched connection errors + require.True(t, database.IsRetryablePostgresError(errors.New("connection reset by peer"))) + require.True(t, database.IsRetryablePostgresError(errors.New("broken pipe"))) + require.True(t, database.IsRetryablePostgresError(errors.New("conn closed"))) + + // Random application error is NOT retryable + require.False(t, database.IsRetryablePostgresError(errors.New("invalid input"))) +} + +func TestRetryPostgres(t *testing.T) { + t.Run("succeeds on first attempt", func(t *testing.T) { + calls := 0 + err := database.RetryPostgres(context.Background(), 3, func() error { + calls++ + return nil + }) + require.NoError(t, err) + require.Equal(t, 1, calls) + }) + + t.Run("retries on transient error then succeeds", func(t *testing.T) { + calls := 0 + err := database.RetryPostgres(context.Background(), 3, func() error { + calls++ + if calls < 3 { + return &pgconn.PgError{Code: "57P01"} // admin_shutdown + } + return nil + }) + require.NoError(t, err) + require.Equal(t, 3, calls) + }) + + t.Run("does not retry non-retryable errors", func(t *testing.T) { + calls := 0 + err := database.RetryPostgres(context.Background(), 3, func() error { + calls++ + return &pgconn.PgError{Code: "23505"} // unique_violation + }) + require.Error(t, err) + require.Equal(t, 1, calls) + }) + + t.Run("respects context cancellation", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() // cancel immediately + + calls := 0 + err := database.RetryPostgres(ctx, 3, func() error { + calls++ + return io.EOF // retryable, but context is done + }) + // First call happens, then context cancellation is detected + require.ErrorIs(t, err, context.Canceled) + require.Equal(t, 1, calls) + }) + + t.Run("exhausts all attempts", func(t *testing.T) { + calls := 0 + err := database.RetryPostgres(context.Background(), 3, func() error { + calls++ + return io.EOF + }) + require.Error(t, err) + require.Equal(t, 3, calls) + require.ErrorIs(t, err, io.EOF) + }) +} + func Test_Postgres_UniqueViolation(t *testing.T) { if testing.Short() { t.Skip("-short flag enabled") @@ -213,3 +322,47 @@ func Test_Postgres_UniqueViolation(t *testing.T) { require.Error(t, err) require.True(t, database.UniqueViolation(err)) } + +func TestApplyPostgresPoolConfig_Defaults(t *testing.T) { + poolConfig, err := pgxpool.ParseConfig("postgres://user:pass@localhost:5432/db") + require.NoError(t, err) + + database.ApplyPostgresPoolConfig(log.NewTestLogger(), poolConfig, database.ConnectionsConfig{}) + + defaults := database.DefaultPostgresConnectionsConfig() + require.Equal(t, int32(defaults.MaxOpen), poolConfig.MaxConns) + require.Equal(t, int32(defaults.MaxIdle), poolConfig.MinIdleConns) + require.Equal(t, defaults.MaxLifetime, poolConfig.MaxConnLifetime) + require.Equal(t, defaults.MaxIdleTime, poolConfig.MaxConnIdleTime) +} + +func TestApplyPostgresPoolConfig_Overrides(t *testing.T) { + poolConfig, err := pgxpool.ParseConfig("postgres://user:pass@localhost:5432/db") + require.NoError(t, err) + + in := database.ConnectionsConfig{ + MaxOpen: 10, + MaxIdle: 3, + MaxLifetime: time.Minute, + MaxIdleTime: 15 * time.Second, + } + database.ApplyPostgresPoolConfig(log.NewTestLogger(), poolConfig, in) + + require.Equal(t, int32(10), poolConfig.MaxConns) + require.Equal(t, int32(3), poolConfig.MinIdleConns) + require.Equal(t, time.Minute, poolConfig.MaxConnLifetime) + require.Equal(t, 15*time.Second, poolConfig.MaxConnIdleTime) +} + +func TestApplyPostgresPoolConfig_MaxIdleCappedByMaxOpen(t *testing.T) { + poolConfig, err := pgxpool.ParseConfig("postgres://user:pass@localhost:5432/db") + require.NoError(t, err) + + database.ApplyPostgresPoolConfig(log.NewTestLogger(), poolConfig, database.ConnectionsConfig{ + MaxOpen: 4, + MaxIdle: 10, // larger than MaxOpen + }) + + require.Equal(t, int32(4), poolConfig.MaxConns) + require.Equal(t, int32(4), poolConfig.MinIdleConns) +} diff --git a/database/testdb/migrations.go b/database/testdb/migrations.go new file mode 100644 index 00000000..8d5d76a0 --- /dev/null +++ b/database/testdb/migrations.go @@ -0,0 +1,68 @@ +package testdb + +import ( + "crypto/sha256" + "encoding/hex" + "fmt" + "io" + "io/fs" + "sort" + "strconv" + "strings" +) + +// migrationFiles walks the embedded migrations FS, filters by suffix, sorts by +// filename (which starts with a zero-padded version number), and returns the +// sorted filenames and their contents. +func migrationFiles(migrations fs.FS, suffix string) ([]string, []string, error) { + entries, err := fs.ReadDir(migrations, "migrations") + if err != nil { + return nil, nil, fmt.Errorf("reading migrations directory: %w", err) + } + + var names []string + for _, e := range entries { + if e.IsDir() || !strings.HasSuffix(e.Name(), suffix) { + continue + } + names = append(names, e.Name()) + } + sort.Strings(names) + + contents := make([]string, len(names)) + for i, name := range names { + data, err := fs.ReadFile(migrations, "migrations/"+name) + if err != nil { + return nil, nil, fmt.Errorf("reading migration %s: %w", name, err) + } + contents[i] = string(data) + } + + return names, contents, nil +} + +// migrationHash computes a SHA256 hash of migration filenames and contents. +// The hash is content-addressed: it changes when any migration is added or +// modified, which automatically invalidates caches keyed by the hash. +func migrationHash(names, contents []string) string { + h := sha256.New() + for i := range names { + io.WriteString(h, names[i]) + io.WriteString(h, contents[i]) + } + return hex.EncodeToString(h.Sum(nil)) +} + +// migrationVersion extracts the version number from the last migration filename. +// Filenames follow the pattern NNN_description.up.{suffix}.sql +func migrationVersion(names []string) (int, error) { + if len(names) == 0 { + return 0, nil + } + parts := strings.SplitN(names[len(names)-1], "_", 2) + v, err := strconv.Atoi(parts[0]) + if err != nil { + return 0, fmt.Errorf("parsing migration version from %s: %w", names[len(names)-1], err) + } + return v, nil +} diff --git a/database/testdb/migrations_test.go b/database/testdb/migrations_test.go new file mode 100644 index 00000000..5d79c5ad --- /dev/null +++ b/database/testdb/migrations_test.go @@ -0,0 +1,179 @@ +package testdb + +import ( + "embed" + "io/fs" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/moov-io/base/database" + "github.com/moov-io/base/log" +) + +//go:embed all:testdata +var testEmbedFS embed.FS + +// testMigrationsFS wraps the embedded files so that migrationFiles can call +// fs.ReadDir(migrations, "migrations"). We embed testdata/ which contains +// a migrations/ subdirectory, matching the convention used by +// database.WithEmbeddedMigrations (iofs.New(f, "migrations")). +var testMigrationsFS fs.FS + +func init() { + sub, err := fs.Sub(testEmbedFS, "testdata") + if err != nil { + panic(err) + } + testMigrationsFS = sub +} + +func TestMigrationFiles_postgres(t *testing.T) { + names, contents, err := migrationFiles(testMigrationsFS, ".up.postgres.sql") + require.NoError(t, err) + assert.Len(t, names, 2) + assert.Equal(t, "001_create_users.up.postgres.sql", names[0]) + assert.Equal(t, "002_add_email.up.postgres.sql", names[1]) + assert.Contains(t, contents[0], "CREATE TABLE users (id TEXT PRIMARY KEY)") + assert.Contains(t, contents[1], "ALTER TABLE users ADD COLUMN email TEXT") +} + +func TestMigrationFiles_spanner(t *testing.T) { + names, contents, err := migrationFiles(testMigrationsFS, ".up.spanner.sql") + require.NoError(t, err) + assert.Len(t, names, 2) + assert.Equal(t, "001_create_items.up.spanner.sql", names[0]) + assert.Equal(t, "002_add_price.up.spanner.sql", names[1]) + assert.Contains(t, contents[0], "CREATE TABLE items") + assert.Contains(t, contents[1], "ADD COLUMN Price") +} + +func TestMigrationFiles_wrongSuffix(t *testing.T) { + names, _, err := migrationFiles(testMigrationsFS, ".up.mysql.sql") + require.NoError(t, err) + assert.Empty(t, names) +} + +func TestMigrationHash_stable(t *testing.T) { + names := []string{"001_a.up.sql", "002_b.up.sql"} + contents := []string{"-- a\n", "-- b\n"} + h1 := migrationHash(names, contents) + h2 := migrationHash(names, contents) + assert.Equal(t, h1, h2) + assert.Len(t, h1, 64) +} + +func TestMigrationHash_changesOnContent(t *testing.T) { + names := []string{"001_a.up.sql"} + h1 := migrationHash(names, []string{"-- a\n"}) + h2 := migrationHash(names, []string{"-- b\n"}) + assert.NotEqual(t, h1, h2) +} + +func TestMigrationHash_changesOnName(t *testing.T) { + contents := []string{"-- a\n"} + h1 := migrationHash([]string{"001_a.up.sql"}, contents) + h2 := migrationHash([]string{"001_c.up.sql"}, contents) + assert.NotEqual(t, h1, h2) +} + +func TestMigrationVersion(t *testing.T) { + tests := []struct { + name string + names []string + want int + wantErr bool + }{ + {"empty", nil, 0, false}, + {"single", []string{"001_a.up.sql"}, 1, false}, + {"multi", []string{"001_a.up.sql", "010_b.up.sql", "003_c.up.sql"}, 3, false}, + {"large", []string{"001_a.up.sql", "999_final.up.sql"}, 999, false}, + {"bad", []string{"abc.up.sql"}, 0, true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := migrationVersion(tt.names) + if tt.wantErr { + assert.Error(t, err) + } else { + require.NoError(t, err) + assert.Equal(t, tt.want, got) + } + }) + } +} + +func TestHashToLockKey(t *testing.T) { + assert.NotZero(t, hashToLockKey("0123456789abcdef")) + assert.Equal(t, int64('a'), hashToLockKey("a")) +} + +func TestEnsureServiceDatabase(t *testing.T) { + t.Run("empty name", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + ensureServiceDatabase(t, db, "") + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("already exists", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectQuery("SELECT EXISTS"). + WithArgs("svc"). + WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + + ensureServiceDatabase(t, db, "svc") + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("create succeeds", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectQuery("SELECT EXISTS"). + WithArgs("svc"). + WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(false)) + mock.ExpectExec("CREATE DATABASE"). + WillReturnResult(sqlmock.NewResult(0, 1)) + + ensureServiceDatabase(t, db, "svc") + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("create loses race", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectQuery("SELECT EXISTS"). + WithArgs("svc"). + WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(false)) + mock.ExpectExec("CREATE DATABASE"). + WillReturnError(assert.AnError) + mock.ExpectQuery("SELECT EXISTS"). + WithArgs("svc"). + WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + + ensureServiceDatabase(t, db, "svc") + + require.NoError(t, mock.ExpectationsWereMet()) + }) +} + +func TestCreateSpannerDatabaseFromMigrations_nilMigrations(t *testing.T) { + cfg, dropFn, err := CreateSpannerDatabaseFromMigrations(log.NewNopLogger(), database.DatabaseConfig{}, nil) + require.Error(t, err) + assert.Nil(t, dropFn) + assert.Empty(t, cfg.DatabaseName) +} diff --git a/database/testdb/postgres_template.go b/database/testdb/postgres_template.go new file mode 100644 index 00000000..1af2511c --- /dev/null +++ b/database/testdb/postgres_template.go @@ -0,0 +1,221 @@ +package testdb + +import ( + "context" + "database/sql" + "fmt" + "io/fs" + "testing" + + "github.com/jackc/pgx/v5" + "github.com/moov-io/base" + "github.com/moov-io/base/database" + "github.com/moov-io/base/log" +) + +// NewPostgresDatabaseFromTemplate creates an isolated test database by cloning +// a template database that has already been migrated. The template is created +// once per unique migration set (content-addressed by SHA256) and reused across +// all tests in the process. This avoids running migrations for every test. +// +// The template database is marked IS_TEMPLATE=true and ALLOW_CONNECTIONS=false +// so no test can accidentally connect to it. Each test gets its own clone with +// full isolation. +func NewPostgresDatabaseFromTemplate( + tb testing.TB, + logger log.Logger, + cfg database.DatabaseConfig, + migrations fs.FS, +) (database.DatabaseConfig, func()) { + tb.Helper() + + if cfg.Postgres == nil { + tb.Fatal("NewPostgresDatabaseFromTemplate: postgres config not defined") + } + if migrations == nil { + tb.Fatal("NewPostgresDatabaseFromTemplate: migrations FS is nil") + } + + names, contents, err := migrationFiles(migrations, ".up.postgres.sql") + if err != nil { + tb.Fatal(fmt.Errorf("reading postgres migrations: %w", err)) + } + + hash := migrationHash(names, contents) + templateName := fmt.Sprintf("tmpl_%s", hash[:16]) + + // Connect to an admin database to manage templates. + // Try "moov" first (the standard Moov docker-compose DB), then "postgres". + adminDb := openAdminDb(tb, cfg.Postgres) + + // Ensure the service database exists. Some tests (e.g. TestEnvironment) + // call service.NewEnvironment directly without going through + // CreateTestDatabase, relying on the service database existing as a + // side effect of prior test setup. The old CreateTestDatabase created + // it via openOrCreateDatabase; we preserve that behavior here. + ensureServiceDatabase(tb, adminDb, cfg.DatabaseName) + + ensurePostgresTemplate(tb, logger, adminDb, cfg, migrations, hash, templateName) + + // Create the per-test database from the template. The template creation lock + // is intentionally released before this point: the lock protects the template + // definition, not every clone created from it. + testDbName := "test" + base.ID()[0:26] + _, err = adminDb.Exec(fmt.Sprintf("CREATE DATABASE %s TEMPLATE %s", + pgx.Identifier{testDbName}.Sanitize(), + pgx.Identifier{templateName}.Sanitize())) + if err != nil { + tb.Fatal(fmt.Errorf("creating test database from template: %w", err)) + } + + dropFn := func() { + _, err := adminDb.Exec(fmt.Sprintf("DROP DATABASE %s", pgx.Identifier{testDbName}.Sanitize())) + if err != nil { + tb.Logf("cleanup: drop database %s: %v", testDbName, err) + } + if err := adminDb.Close(); err != nil { + tb.Logf("cleanup: close admin database: %v", err) + } + } + + cfg.DatabaseName = testDbName + return cfg, dropFn +} + +func ensurePostgresTemplate( + tb testing.TB, + logger log.Logger, + adminDb *sql.DB, + cfg database.DatabaseConfig, + migrations fs.FS, + hash string, + templateName string, +) { + tb.Helper() + + ctx := context.Background() + conn, err := adminDb.Conn(ctx) + if err != nil { + tb.Fatal(fmt.Errorf("acquiring admin connection: %w", err)) + } + defer func() { + if err := conn.Close(); err != nil { + tb.Logf("cleanup: close admin connection: %v", err) + } + }() + + // Serialize template creation across processes using a session-level advisory + // lock keyed by the migration hash. The lock and unlock must use the same + // physical connection. + lockKey := hashToLockKey(hash) + if _, err := conn.ExecContext(ctx, "SELECT pg_advisory_lock($1)", lockKey); err != nil { + tb.Fatal(fmt.Errorf("acquiring advisory lock: %w", err)) + } + defer func() { + if _, err := conn.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", lockKey); err != nil { + tb.Logf("cleanup: release advisory lock: %v", err) + } + }() + + var exists bool + err = conn.QueryRowContext( + ctx, + "SELECT EXISTS(SELECT 1 FROM pg_database WHERE datname = $1 AND datistemplate = true)", + templateName, + ).Scan(&exists) + if err != nil { + tb.Fatal(fmt.Errorf("checking template existence: %w", err)) + } + + if exists { + return + } + + _, err = conn.ExecContext(ctx, fmt.Sprintf("CREATE DATABASE %s", pgx.Identifier{templateName}.Sanitize())) + if err != nil { + tb.Fatal(fmt.Errorf("creating template database: %w", err)) + } + + templateCfg := database.DatabaseConfig{ + DatabaseName: templateName, + Postgres: cfg.Postgres, + } + if err := database.RunMigrations(logger, templateCfg, database.WithEmbeddedMigrations(migrations)); err != nil { + if _, dropErr := conn.ExecContext(ctx, fmt.Sprintf("DROP DATABASE IF EXISTS %s", pgx.Identifier{templateName}.Sanitize())); dropErr != nil { + tb.Logf("cleanup: drop template database %s: %v", templateName, dropErr) + } + tb.Fatal(fmt.Errorf("running migrations on template: %w", err)) + } + + _, err = conn.ExecContext(ctx, fmt.Sprintf("ALTER DATABASE %s IS_TEMPLATE true", pgx.Identifier{templateName}.Sanitize())) + if err != nil { + tb.Fatal(fmt.Errorf("marking template: %w", err)) + } + _, err = conn.ExecContext(ctx, fmt.Sprintf("ALTER DATABASE %s ALLOW_CONNECTIONS false", pgx.Identifier{templateName}.Sanitize())) + if err != nil { + tb.Fatal(fmt.Errorf("disallowing connections to template: %w", err)) + } +} + +// openAdminDb connects to an admin database for managing templates. +// Tries "moov" first, then falls back to "postgres". +func openAdminDb(tb testing.TB, pgCfg *database.PostgresConfig) *sql.DB { + tb.Helper() + + for _, adminDb := range []string{"moov", "postgres"} { + db, err := sql.Open("pgx", fmt.Sprintf("postgres://%s:%s@%s/%s", + pgCfg.User, pgCfg.Password, pgCfg.Address, adminDb)) + if err != nil { + continue + } + if err := db.Ping(); err != nil { + db.Close() + continue + } + return db + } + + tb.Fatal("could not connect to admin database (tried 'moov' and 'postgres')") + return nil +} + +// hashToLockKey converts the first 8 bytes of a hex hash string into an int64 +// for use with pg_advisory_lock. +func hashToLockKey(hash string) int64 { + var key int64 + for i := 0; i < 8 && i < len(hash); i++ { + key = key<<8 | int64(hash[i]) + } + return key +} + +// ensureServiceDatabase creates the service database if it doesn't exist. +// This preserves the side-effect behavior of the old CreateTestDatabase, +// which created the service database via openOrCreateDatabase. Some tests +// call service.NewEnvironment directly without CreateTestDatabase and rely +// on the service database existing. +func ensureServiceDatabase(tb testing.TB, adminDb *sql.DB, dbName string) { + tb.Helper() + if dbName == "" { + return + } + var exists bool + err := adminDb.QueryRow( + "SELECT EXISTS(SELECT 1 FROM pg_database WHERE datname = $1)", + dbName, + ).Scan(&exists) + if err != nil { + tb.Fatal(fmt.Errorf("checking service database existence: %w", err)) + } + if !exists { + _, err = adminDb.Exec(fmt.Sprintf("CREATE DATABASE %s", pgx.Identifier{dbName}.Sanitize())) + if err != nil { + if checkErr := adminDb.QueryRow( + "SELECT EXISTS(SELECT 1 FROM pg_database WHERE datname = $1)", + dbName, + ).Scan(&exists); checkErr != nil || !exists { + tb.Fatal(fmt.Errorf("creating service database %s: %w", dbName, err)) + } + } + } +} diff --git a/database/testdb/spanner_fast.go b/database/testdb/spanner_fast.go new file mode 100644 index 00000000..00bbeabc --- /dev/null +++ b/database/testdb/spanner_fast.go @@ -0,0 +1,166 @@ +package testdb + +import ( + "context" + "fmt" + "io/fs" + "testing" + + "cloud.google.com/go/spanner" + spannerdb "cloud.google.com/go/spanner/admin/database/apiv1" + "cloud.google.com/go/spanner/admin/database/apiv1/databasepb" + "cloud.google.com/go/spanner/spansql" + "github.com/moov-io/base/database" + "github.com/moov-io/base/log" +) + +// NewSpannerDatabaseFromMigrations creates a Spanner test database with all +// migrations applied in a single batched DDL update, instead of running each +// migration file individually through golang-migrate. The +// spanner_schema_migrations table is pre-populated with the latest version so +// subsequent RunMigrations calls are no-ops (ErrNoChange). +// +// If the batched DDL fails or cannot be parsed, the function falls back to +// creating a fresh empty database and lets the caller run migrations normally. +func NewSpannerDatabaseFromMigrations( + tb testing.TB, + logger log.Logger, + cfg database.DatabaseConfig, + migrations fs.FS, +) (database.DatabaseConfig, func()) { + tb.Helper() + cfg2, dropFn, err := CreateSpannerDatabaseFromMigrations(logger, cfg, migrations) + if err != nil { + tb.Fatal(err) + } + return cfg2, dropFn +} + +// CreateSpannerDatabaseFromMigrations is the error-returning variant of +// NewSpannerDatabaseFromMigrations. It is safe to call from contexts that +// cannot use tb.Fatal (e.g. inside sync.Once.Do). +func CreateSpannerDatabaseFromMigrations( + logger log.Logger, + cfg database.DatabaseConfig, + migrations fs.FS, +) (database.DatabaseConfig, func(), error) { + if migrations == nil { + return cfg, nil, fmt.Errorf("migrations FS is nil") + } + + cfg, err := NewSpannerDatabase(cfg.DatabaseName, cfg.Spanner) + if err != nil { + return cfg, nil, fmt.Errorf("creating spanner database: %w", err) + } + + dropFn := func() { dropSpannerDBByCfg(cfg) } + + names, contents, err := migrationFiles(migrations, ".up.spanner.sql") + if err != nil { + dropSpannerDBByCfg(cfg) + return cfg, nil, fmt.Errorf("reading spanner migrations: %w", err) + } + + if len(names) == 0 { + return cfg, dropFn, nil + } + + version, err := migrationVersion(names) + if err != nil { + dropSpannerDBByCfg(cfg) + return cfg, nil, err + } + + var allStmts []string + for i, content := range contents { + ddl, err := spansql.ParseDDL(names[i], content) + if err != nil { + logger.Info().Logf("spanner fast create: DDL parse failed for %s, falling back: %v", names[i], err) + return fallbackSpannerE(cfg) + } + for _, stmt := range ddl.List { + allStmts = append(allStmts, stmt.SQL()) + } + } + + schemaMigrationsDDL := `CREATE TABLE spanner_schema_migrations ( + Version INT64 NOT NULL, + Dirty BOOL NOT NULL +) PRIMARY KEY(Version)` + allStmts = append([]string{schemaMigrationsDDL}, allStmts...) + + ctx := context.Background() + adminClient, err := spannerdb.NewDatabaseAdminClient(ctx) + if err != nil { + dropSpannerDBByCfg(cfg) + return cfg, nil, fmt.Errorf("creating spanner admin client: %w", err) + } + defer adminClient.Close() + + dbPath := fmt.Sprintf("projects/%s/instances/%s/databases/%s", + cfg.Spanner.Project, cfg.Spanner.Instance, cfg.DatabaseName) + + op, err := adminClient.UpdateDatabaseDdl(ctx, &databasepb.UpdateDatabaseDdlRequest{ + Database: dbPath, + Statements: allStmts, + }) + if err != nil { + logger.Info().Logf("spanner fast create: batched DDL request failed, falling back: %v", err) + return fallbackSpannerE(cfg) + } + if err := op.Wait(ctx); err != nil { + logger.Info().Logf("spanner fast create: batched DDL operation failed, falling back: %v", err) + return fallbackSpannerE(cfg) + } + + dataClient, err := spanner.NewClient(ctx, dbPath) + if err != nil { + dropSpannerDBByCfg(cfg) + return cfg, nil, fmt.Errorf("creating spanner data client for version insert: %w", err) + } + _, err = dataClient.ReadWriteTransaction(ctx, func(ctx context.Context, txn *spanner.ReadWriteTransaction) error { + return txn.BufferWrite([]*spanner.Mutation{ + spanner.Delete("spanner_schema_migrations", spanner.AllKeys()), + spanner.Insert("spanner_schema_migrations", + []string{"Version", "Dirty"}, + []interface{}{version, false}, + ), + }) + }) + dataClient.Close() + if err != nil { + dropSpannerDBByCfg(cfg) + return cfg, nil, fmt.Errorf("inserting migration version: %w", err) + } + + logger.Info().Logf("spanner fast create: applied %d migration files as %d DDL statements (version %d)", len(names), len(allStmts)-1, version) + return cfg, dropFn, nil +} + +// fallbackSpannerE drops the current database, creates a fresh empty one, and +// returns it so the caller can run migrations normally via LoadDatabase. +func fallbackSpannerE(cfg database.DatabaseConfig) (database.DatabaseConfig, func(), error) { + dropSpannerDBByCfg(cfg) + cfg2, err := NewSpannerDatabase(cfg.DatabaseName, cfg.Spanner) + if err != nil { + return cfg, nil, fmt.Errorf("fallback: creating fresh spanner database: %w", err) + } + return cfg2, func() { dropSpannerDBByCfg(cfg2) }, nil +} + +// dropSpannerDBByCfg drops a Spanner database described by a DatabaseConfig. +func dropSpannerDBByCfg(cfg database.DatabaseConfig) { + ctx := context.Background() + adminClient, err := spannerdb.NewDatabaseAdminClient(ctx) + if err != nil { + return + } + defer adminClient.Close() + + dbPath := fmt.Sprintf("projects/%s/instances/%s/databases/%s", + cfg.Spanner.Project, cfg.Spanner.Instance, cfg.DatabaseName) + + _ = adminClient.DropDatabase(ctx, &databasepb.DropDatabaseRequest{ + Database: dbPath, + }) +} diff --git a/database/testdb/testdata/migrations/001_create_items.up.spanner.sql b/database/testdb/testdata/migrations/001_create_items.up.spanner.sql new file mode 100644 index 00000000..27d39e2d --- /dev/null +++ b/database/testdb/testdata/migrations/001_create_items.up.spanner.sql @@ -0,0 +1,3 @@ +CREATE TABLE items ( + ItemID STRING(36) NOT NULL, +) PRIMARY KEY(ItemID); diff --git a/database/testdb/testdata/migrations/001_create_users.up.postgres.sql b/database/testdb/testdata/migrations/001_create_users.up.postgres.sql new file mode 100644 index 00000000..d8ed438a --- /dev/null +++ b/database/testdb/testdata/migrations/001_create_users.up.postgres.sql @@ -0,0 +1 @@ +CREATE TABLE users (id TEXT PRIMARY KEY); diff --git a/database/testdb/testdata/migrations/002_add_email.up.postgres.sql b/database/testdb/testdata/migrations/002_add_email.up.postgres.sql new file mode 100644 index 00000000..79fc1d86 --- /dev/null +++ b/database/testdb/testdata/migrations/002_add_email.up.postgres.sql @@ -0,0 +1 @@ +ALTER TABLE users ADD COLUMN email TEXT; diff --git a/database/testdb/testdata/migrations/002_add_price.up.spanner.sql b/database/testdb/testdata/migrations/002_add_price.up.spanner.sql new file mode 100644 index 00000000..c1bc3fbb --- /dev/null +++ b/database/testdb/testdata/migrations/002_add_price.up.spanner.sql @@ -0,0 +1 @@ +ALTER TABLE items ADD COLUMN Price INT64; diff --git a/go.mod b/go.mod index 9ead87e7..6b48a913 100644 --- a/go.mod +++ b/go.mod @@ -42,6 +42,7 @@ require ( cloud.google.com/go/longrunning v0.13.0 // indirect cloud.google.com/go/monitoring v1.29.0 // indirect filippo.io/edwards25519 v1.2.0 // indirect + github.com/DATA-DOG/go-sqlmock v1.5.2 // indirect github.com/GoogleCloudPlatform/grpc-gcp-go/grpcgcp v1.6.0 // indirect github.com/GoogleCloudPlatform/opentelemetry-operations-go/detectors/gcp v1.32.0 // indirect github.com/GoogleCloudPlatform/opentelemetry-operations-go/exporter/metric v0.56.0 // indirect @@ -93,10 +94,10 @@ require ( go.yaml.in/yaml/v2 v2.4.4 // indirect go.yaml.in/yaml/v3 v3.0.4 // indirect golang.org/x/crypto v0.51.0 // indirect - golang.org/x/net v0.54.0 // indirect + golang.org/x/net v0.55.0 // indirect golang.org/x/oauth2 v0.36.0 // indirect golang.org/x/sync v0.20.0 // indirect - golang.org/x/sys v0.44.0 // indirect + golang.org/x/sys v0.45.0 // indirect golang.org/x/text v0.37.0 // indirect google.golang.org/api v0.278.0 // indirect google.golang.org/genproto v0.0.0-20260504160031-60b97b32f348 // indirect diff --git a/go.sum b/go.sum index a0194c29..ad2a3b99 100644 --- a/go.sum +++ b/go.sum @@ -5,8 +5,6 @@ cloud.google.com/go v0.123.0 h1:2NAUJwPR47q+E35uaJeYoNhuNEM9kM8SjgRgdeOJUSE= cloud.google.com/go v0.123.0/go.mod h1:xBoMV08QcqUGuPW65Qfm1o9Y4zKZBpGS+7bImXLTAZU= cloud.google.com/go/alloydb v1.26.0 h1:UTzyumJ8tEo0CqwzLkV4WMGnCxvvhw3BDy1nXfCt9KE= cloud.google.com/go/alloydb v1.26.0/go.mod h1:oqHGc/Xb5fWtH+wIDpu2wcPJX9oML/fGJuH/sp8ysyo= -cloud.google.com/go/alloydbconn v1.18.2 h1:zxxUnSU50d1sS/nswGBldpZsaxF3emirbO+xF+Rtgig= -cloud.google.com/go/alloydbconn v1.18.2/go.mod h1:0vfrUSwleLoK13ycnoaZ1A4yCfg3r7rig7Xm5vKlEtM= cloud.google.com/go/alloydbconn v1.18.3 h1:7A8QN5DtPkyGToqekWSm4+Ryrx2+xYWkvRYh0A0fAVg= cloud.google.com/go/alloydbconn v1.18.3/go.mod h1:cbwUQHP9eLp3Y66xZyQZrKmlcvqLnlb91qKc6kXV46Q= cloud.google.com/go/auth v0.20.0 h1:kXTssoVb4azsVDoUiF8KvxAqrsQcQtB53DcSgta74CA= @@ -32,6 +30,8 @@ filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q github.com/Azure/go-ansiterm v0.0.0-20230124172434-306776ec8161 h1:L/gRVlceqvL25UVaW/CKtUDjefjrs0SPonmDGUVOYP0= github.com/Azure/go-ansiterm v0.0.0-20230124172434-306776ec8161/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E= github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= +github.com/DATA-DOG/go-sqlmock v1.5.2 h1:OcvFkGmslmlZibjAjaHm3L//6LiuBgolP7OputlJIzU= +github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU= github.com/GoogleCloudPlatform/grpc-gcp-go/grpcgcp v1.6.0 h1:BzsL0qE7LvtTEtXG7Dt5NS1EP0CQwI21HZfj9aGghhw= github.com/GoogleCloudPlatform/grpc-gcp-go/grpcgcp v1.6.0/go.mod h1:I7kE2kM3qCr9QPT4cU4cCFYkEpVyVr16YOGUHzy+nR0= github.com/GoogleCloudPlatform/opentelemetry-operations-go/detectors/gcp v1.32.0 h1:rIkQfkCOVKc1OiRCNcSDD8ml5RJlZbH/Xsq7lbpynwc= @@ -168,6 +168,7 @@ github.com/jackc/pgx/v5 v5.9.2 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw= github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= +github.com/kisielk/sqlstruct v0.0.0-20201105191214-5f3e10d3ab46/go.mod h1:yyMNCyc/Ib3bDTKd379tNMpB/7/H5TjM2Y9QJ5THLbE= github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo= github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ= github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORNo= @@ -280,8 +281,6 @@ go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= -golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI= -golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q= golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI= golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8= golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= @@ -294,10 +293,8 @@ golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73r golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20201110031124-69a78807bb2b/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= -golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA= -golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs= -golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w= -golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ= +golang.org/x/net v0.55.0 h1:bcvxaJn3e1U6InsFWt1JUq1aSjnRxLzT2rtD2KfkDF8= +golang.org/x/net v0.55.0/go.mod h1:L5U2KuzuOe1lY7Z+aWVIKK6qEeJXnXV9yzGA+WCHJww= golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U= golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs= golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q= @@ -310,12 +307,10 @@ golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5h golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ= -golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY= +golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= -golang.org/x/text v0.36.0 h1:JfKh3XmcRPqZPKevfXVpI1wXPTqbkE5f7JA92a55Yxg= -golang.org/x/text v0.36.0/go.mod h1:NIdBknypM8iqVmPiuco0Dh6P5Jcdk8lJL0CUebqK164= golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc= golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38= golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= @@ -339,8 +334,6 @@ google.golang.org/genproto v0.0.0-20260504160031-60b97b32f348 h1:JjVGDZYWkJWZcxv google.golang.org/genproto v0.0.0-20260504160031-60b97b32f348/go.mod h1:95PqD4xM+AdOcBGsmgfaofXsiA37uXDtDufVbntT3TU= google.golang.org/genproto/googleapis/api v0.0.0-20260504160031-60b97b32f348 h1:U8orV30l6KpDsi9dxU0CoJZGbjS8EEpw+6ba+XwGPQA= google.golang.org/genproto/googleapis/api v0.0.0-20260504160031-60b97b32f348/go.mod h1:Yzdzr5OOZFgSsEV2D/Xi9NL3bszpXFAg0hFJiRohcD8= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260504160031-60b97b32f348 h1:pfIbyB44sWzHiCpRqIen67ZQnVXSfIxWrqUMk1qwODE= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260504160031-60b97b32f348/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= google.golang.org/genproto/googleapis/rpc v0.0.0-20260511170946-3700d4141b60 h1:seT2EwLWM78plQ7wcDfuWBc/4FAEAXDDiaSol4ku4qo= google.golang.org/genproto/googleapis/rpc v0.0.0-20260511170946-3700d4141b60/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c= @@ -348,8 +341,6 @@ google.golang.org/grpc v1.23.0/go.mod h1:Y5yQAOtifL1yxbo5wqy6BxZv8vAUGQwXBOALyac google.golang.org/grpc v1.25.1/go.mod h1:c3i+UQWmh7LiEpx4sFZnkU36qjEYZ0imhYfXVyQciAY= google.golang.org/grpc v1.27.0/go.mod h1:qbnxyOmOxrQa7FizSgH+ReBfzJrCY1pSN7KXBS8abTk= google.golang.org/grpc v1.33.2/go.mod h1:JMHMWHQWaTccqQQlmk3MJZS+GWXOdAesneDmEnv2fbc= -google.golang.org/grpc v1.81.0 h1:W3G9N3KQf3BU+YuCtGKJk0CmxQNbAISICD/9AORxLIw= -google.golang.org/grpc v1.81.0/go.mod h1:xGH9GfzOyMTGIOXBJmXt+BX/V0kcdQbdcuwQ/zNw42I= google.golang.org/grpc v1.81.1 h1:VnnIIZ88UzOOKLukQi+ImGz8O1Wdp8nAGGnvOfEIWQQ= google.golang.org/grpc v1.81.1/go.mod h1:xGH9GfzOyMTGIOXBJmXt+BX/V0kcdQbdcuwQ/zNw42I= google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=