This commit is contained in:
@@ -0,0 +1,129 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"context"
|
||||
"embed"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"path"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
//go:embed migrations/*.sql
|
||||
var migrationFiles embed.FS
|
||||
|
||||
const (
|
||||
defaultMaxConns = 10
|
||||
defaultMinConns = 1
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
URL string
|
||||
MaxConns int32
|
||||
MinConns int32
|
||||
MaxConnIdle time.Duration
|
||||
}
|
||||
|
||||
func NewPool(ctx context.Context, config Config) (*pgxpool.Pool, error) {
|
||||
if strings.TrimSpace(config.URL) == "" {
|
||||
return nil, errors.New("database URL is required")
|
||||
}
|
||||
poolConfig, err := pgxpool.ParseConfig(config.URL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse database URL: %w", err)
|
||||
}
|
||||
if config.MaxConns == 0 {
|
||||
config.MaxConns = defaultMaxConns
|
||||
}
|
||||
if config.MinConns == 0 {
|
||||
config.MinConns = defaultMinConns
|
||||
}
|
||||
if config.MaxConns < config.MinConns || config.MinConns < 0 {
|
||||
return nil, errors.New("database pool limits are invalid")
|
||||
}
|
||||
poolConfig.MaxConns = config.MaxConns
|
||||
poolConfig.MinConns = config.MinConns
|
||||
if config.MaxConnIdle > 0 {
|
||||
poolConfig.MaxConnIdleTime = config.MaxConnIdle
|
||||
}
|
||||
pool, err := pgxpool.NewWithConfig(ctx, poolConfig)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create database pool: %w", err)
|
||||
}
|
||||
return pool, nil
|
||||
}
|
||||
|
||||
func Ping(ctx context.Context, pool *pgxpool.Pool) error {
|
||||
if pool == nil {
|
||||
return errors.New("database pool is nil")
|
||||
}
|
||||
if err := pool.Ping(ctx); err != nil {
|
||||
return fmt.Errorf("database ping: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func Migrate(ctx context.Context, pool *pgxpool.Pool) error {
|
||||
if pool == nil {
|
||||
return errors.New("database pool is nil")
|
||||
}
|
||||
entries, err := fs.Glob(migrationFiles, "migrations/*.sql")
|
||||
if err != nil {
|
||||
return fmt.Errorf("list migrations: %w", err)
|
||||
}
|
||||
sort.Strings(entries)
|
||||
for _, entry := range entries {
|
||||
if err := applyMigration(ctx, pool, entry); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func applyMigration(ctx context.Context, pool *pgxpool.Pool, entry string) error {
|
||||
migrationID := strings.TrimSuffix(path.Base(entry), path.Ext(entry))
|
||||
sqlBytes, err := migrationFiles.ReadFile(entry)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read migration %s: %w", migrationID, err)
|
||||
}
|
||||
tx, err := pool.BeginTx(ctx, pgx.TxOptions{})
|
||||
if err != nil {
|
||||
return fmt.Errorf("begin migration %s: %w", migrationID, err)
|
||||
}
|
||||
defer func() { _ = tx.Rollback(ctx) }()
|
||||
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtext('itworx-pulse:schema-migrations'))`); err != nil {
|
||||
return fmt.Errorf("lock migrations: %w", err)
|
||||
}
|
||||
if _, err := tx.Exec(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
id text PRIMARY KEY,
|
||||
applied_at timestamptz NOT NULL DEFAULT now()
|
||||
)`); err != nil {
|
||||
return fmt.Errorf("create migration table: %w", err)
|
||||
}
|
||||
var exists bool
|
||||
if err := tx.QueryRow(ctx, `SELECT EXISTS (SELECT 1 FROM schema_migrations WHERE id = $1)`, migrationID).Scan(&exists); err != nil {
|
||||
return fmt.Errorf("check migration %s: %w", migrationID, err)
|
||||
}
|
||||
if exists {
|
||||
if err := tx.Commit(ctx); err != nil {
|
||||
return fmt.Errorf("commit migration check %s: %w", migrationID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if _, err := tx.Exec(ctx, string(sqlBytes)); err != nil {
|
||||
return fmt.Errorf("apply migration %s: %w", migrationID, err)
|
||||
}
|
||||
if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations (id) VALUES ($1)`, migrationID); err != nil {
|
||||
return fmt.Errorf("record migration %s: %w", migrationID, err)
|
||||
}
|
||||
if err := tx.Commit(ctx); err != nil {
|
||||
return fmt.Errorf("commit migration %s: %w", migrationID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
Reference in New Issue
Block a user