381 lines
10 KiB
Go
381 lines
10 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"fmt"
|
|
"hash/crc32"
|
|
"os"
|
|
"strconv"
|
|
"strings"
|
|
|
|
"github.com/lib/pq"
|
|
)
|
|
|
|
const (
|
|
nilVersion int = -1
|
|
migrationsTable = "schema_migrations"
|
|
advisoryLockIDSalt uint = 1486364155
|
|
)
|
|
|
|
type migrateDB struct {
|
|
conn *sql.Conn
|
|
db *sql.DB
|
|
schemaName string
|
|
dbName string
|
|
lockID string
|
|
isLocked bool
|
|
}
|
|
|
|
func openDatabase(connStr string, schemaName string) (*migrateDB, error) {
|
|
db, err := sql.Open("postgres", connStr)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to open database: %w", err)
|
|
}
|
|
|
|
if err := db.Ping(); err != nil {
|
|
db.Close()
|
|
return nil, fmt.Errorf("failed to connect to database: %w", err)
|
|
}
|
|
|
|
conn, err := db.Conn(context.Background())
|
|
if err != nil {
|
|
db.Close()
|
|
return nil, fmt.Errorf("failed to acquire connection: %w", err)
|
|
}
|
|
|
|
mdb := &migrateDB{
|
|
conn: conn,
|
|
db: db,
|
|
}
|
|
|
|
// get actual database name
|
|
if err := conn.QueryRowContext(context.Background(), "SELECT current_database()").Scan(&mdb.dbName); err != nil {
|
|
mdb.close()
|
|
return nil, fmt.Errorf("failed to get database name: %w", err)
|
|
}
|
|
|
|
// determine schema name
|
|
if schemaName != "" {
|
|
mdb.schemaName = schemaName
|
|
} else {
|
|
if err := conn.QueryRowContext(context.Background(), "SELECT current_schema()").Scan(&mdb.schemaName); err != nil {
|
|
mdb.close()
|
|
return nil, fmt.Errorf("failed to get schema name: %w", err)
|
|
}
|
|
}
|
|
|
|
mdb.lockID = generateAdvisoryLockID(mdb.dbName, mdb.schemaName)
|
|
|
|
// ensure the schema_migrations table exists
|
|
if err := mdb.ensureVersionTable(); err != nil {
|
|
mdb.close()
|
|
return nil, err
|
|
}
|
|
|
|
return mdb, nil
|
|
}
|
|
|
|
func (d *migrateDB) close() error {
|
|
var connErr, dbErr error
|
|
if d.conn != nil {
|
|
connErr = d.conn.Close()
|
|
}
|
|
if d.db != nil {
|
|
dbErr = d.db.Close()
|
|
}
|
|
if connErr != nil {
|
|
return connErr
|
|
}
|
|
return dbErr
|
|
}
|
|
|
|
func (d *migrateDB) ensureVersionTable() error {
|
|
if err := d.lock(); err != nil {
|
|
return err
|
|
}
|
|
defer d.unlock()
|
|
|
|
// check if table already exists
|
|
var count int
|
|
query := `SELECT COUNT(1) FROM information_schema.tables WHERE table_schema = $1 AND table_name = $2 LIMIT 1`
|
|
if err := d.conn.QueryRowContext(context.Background(), query, d.schemaName, migrationsTable).Scan(&count); err != nil {
|
|
return fmt.Errorf("failed to check for migrations table: %w", err)
|
|
}
|
|
|
|
if count > 0 {
|
|
return nil
|
|
}
|
|
|
|
// create the table
|
|
createQuery := `CREATE TABLE IF NOT EXISTS ` +
|
|
pq.QuoteIdentifier(d.schemaName) + `.` + pq.QuoteIdentifier(migrationsTable) +
|
|
` (version bigint NOT NULL PRIMARY KEY, dirty boolean NOT NULL)`
|
|
|
|
if _, err := d.conn.ExecContext(context.Background(), createQuery); err != nil {
|
|
return fmt.Errorf("failed to create migrations table: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// generateAdvisoryLockID replicates golang-migrate's lock ID generation.
|
|
// Internally: strings.Join(append([]string{schemaName, tableName}, databaseName), "\x00")
|
|
// Result: "schemaName\x00tableName\x00databaseName"
|
|
// Then: CRC32(result) * 1486364155
|
|
func generateAdvisoryLockID(dbName, schemaName string) string {
|
|
combined := strings.Join([]string{schemaName, migrationsTable, dbName}, "\x00")
|
|
sum := crc32.ChecksumIEEE([]byte(combined))
|
|
sum = sum * uint32(advisoryLockIDSalt)
|
|
return fmt.Sprint(sum)
|
|
}
|
|
|
|
func (d *migrateDB) lock() error {
|
|
if d.isLocked {
|
|
return fmt.Errorf("database already locked")
|
|
}
|
|
|
|
query := `SELECT pg_advisory_lock($1)`
|
|
if _, err := d.conn.ExecContext(context.Background(), query, d.lockID); err != nil {
|
|
return fmt.Errorf("failed to acquire advisory lock: %w", err)
|
|
}
|
|
|
|
d.isLocked = true
|
|
return nil
|
|
}
|
|
|
|
func (d *migrateDB) unlock() error {
|
|
if !d.isLocked {
|
|
return nil
|
|
}
|
|
|
|
query := `SELECT pg_advisory_unlock($1)`
|
|
if _, err := d.conn.ExecContext(context.Background(), query, d.lockID); err != nil {
|
|
return fmt.Errorf("failed to release advisory lock: %w", err)
|
|
}
|
|
|
|
d.isLocked = false
|
|
return nil
|
|
}
|
|
|
|
func (d *migrateDB) version() (int, bool, error) {
|
|
query := `SELECT version, dirty FROM ` +
|
|
pq.QuoteIdentifier(d.schemaName) + `.` + pq.QuoteIdentifier(migrationsTable) +
|
|
` LIMIT 1`
|
|
|
|
var version int
|
|
var dirty bool
|
|
err := d.conn.QueryRowContext(context.Background(), query).Scan(&version, &dirty)
|
|
|
|
switch {
|
|
case err == sql.ErrNoRows:
|
|
return nilVersion, false, nil
|
|
case err != nil:
|
|
if e, ok := err.(*pq.Error); ok {
|
|
if e.Code.Name() == "undefined_table" {
|
|
return nilVersion, false, nil
|
|
}
|
|
}
|
|
return 0, false, fmt.Errorf("failed to get migration version: %w", err)
|
|
default:
|
|
return version, dirty, nil
|
|
}
|
|
}
|
|
|
|
func (d *migrateDB) setVersion(version int, dirty bool) error {
|
|
tx, err := d.conn.BeginTx(context.Background(), &sql.TxOptions{})
|
|
if err != nil {
|
|
return fmt.Errorf("failed to begin transaction: %w", err)
|
|
}
|
|
|
|
truncateQuery := `TRUNCATE ` +
|
|
pq.QuoteIdentifier(d.schemaName) + `.` + pq.QuoteIdentifier(migrationsTable)
|
|
|
|
if _, err := tx.Exec(truncateQuery); err != nil {
|
|
tx.Rollback()
|
|
return fmt.Errorf("failed to truncate migrations table: %w", err)
|
|
}
|
|
|
|
// re-write the schema version for nil dirty versions to prevent
|
|
// empty schema version for failed down migration on the first migration
|
|
if version >= 0 || (version == nilVersion && dirty) {
|
|
insertQuery := `INSERT INTO ` +
|
|
pq.QuoteIdentifier(d.schemaName) + `.` + pq.QuoteIdentifier(migrationsTable) +
|
|
` (version, dirty) VALUES ($1, $2)`
|
|
|
|
if _, err := tx.Exec(insertQuery, version, dirty); err != nil {
|
|
tx.Rollback()
|
|
return fmt.Errorf("failed to insert migration version: %w", err)
|
|
}
|
|
}
|
|
|
|
if err := tx.Commit(); err != nil {
|
|
return fmt.Errorf("failed to commit version update: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (d *migrateDB) run(filePath string) error {
|
|
content, err := os.ReadFile(filePath)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to read migration file %s: %w", filePath, err)
|
|
}
|
|
|
|
query := string(content)
|
|
if strings.TrimSpace(query) == "" {
|
|
return nil
|
|
}
|
|
|
|
if _, err := d.conn.ExecContext(context.Background(), query); err != nil {
|
|
if pgErr, ok := err.(*pq.Error); ok {
|
|
message := fmt.Sprintf("migration failed: %s", pgErr.Message)
|
|
if pgErr.Position != "" {
|
|
if pos, parseErr := strconv.ParseUint(pgErr.Position, 10, 64); parseErr == nil {
|
|
line, col, ok := computeLineFromPos(query, int(pos))
|
|
if ok {
|
|
message = fmt.Sprintf("%s (line %d, column %d)", message, line, col)
|
|
}
|
|
}
|
|
}
|
|
if pgErr.Detail != "" {
|
|
message = fmt.Sprintf("%s, %s", message, pgErr.Detail)
|
|
}
|
|
return fmt.Errorf("%s", message)
|
|
}
|
|
return fmt.Errorf("migration failed: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (d *migrateDB) drop() error {
|
|
// Drop tables.
|
|
tableQuery := `SELECT table_name FROM information_schema.tables WHERE table_schema=$1 AND table_type='BASE TABLE'`
|
|
rows, err := d.conn.QueryContext(context.Background(), tableQuery, d.schemaName)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to query tables: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
|
|
var tableNames []string
|
|
for rows.Next() {
|
|
var tableName string
|
|
if err := rows.Scan(&tableName); err != nil {
|
|
return fmt.Errorf("failed to scan table name: %w", err)
|
|
}
|
|
if len(tableName) > 0 {
|
|
tableNames = append(tableNames, tableName)
|
|
}
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return fmt.Errorf("failed to iterate tables: %w", err)
|
|
}
|
|
|
|
for _, t := range tableNames {
|
|
dropQuery := `DROP TABLE IF EXISTS ` + pq.QuoteIdentifier(d.schemaName) + `.` + pq.QuoteIdentifier(t) + ` CASCADE`
|
|
if _, err := d.conn.ExecContext(context.Background(), dropQuery); err != nil {
|
|
return fmt.Errorf("failed to drop table %s: %w", t, err)
|
|
}
|
|
}
|
|
|
|
// Drop views.
|
|
viewQuery := `SELECT table_name FROM information_schema.views WHERE table_schema=$1`
|
|
viewRows, err := d.conn.QueryContext(context.Background(), viewQuery, d.schemaName)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to query views: %w", err)
|
|
}
|
|
defer viewRows.Close()
|
|
|
|
var viewNames []string
|
|
for viewRows.Next() {
|
|
var viewName string
|
|
if err := viewRows.Scan(&viewName); err != nil {
|
|
return fmt.Errorf("failed to scan view name: %w", err)
|
|
}
|
|
viewNames = append(viewNames, viewName)
|
|
}
|
|
if err := viewRows.Err(); err != nil {
|
|
return fmt.Errorf("failed to iterate views: %w", err)
|
|
}
|
|
|
|
for _, v := range viewNames {
|
|
dropQuery := `DROP VIEW IF EXISTS ` + pq.QuoteIdentifier(d.schemaName) + `.` + pq.QuoteIdentifier(v) + ` CASCADE`
|
|
if _, err := d.conn.ExecContext(context.Background(), dropQuery); err != nil {
|
|
return fmt.Errorf("failed to drop view %s: %w", v, err)
|
|
}
|
|
}
|
|
|
|
// Drop custom types (enums, composites, domains, ranges). typtype 'b' (base)
|
|
// already excludes Postgres' auto-generated array types, so no further
|
|
// filtering by name is needed.
|
|
typeQuery := `
|
|
SELECT t.typname
|
|
FROM pg_type t
|
|
JOIN pg_namespace n ON n.oid = t.typnamespace
|
|
WHERE n.nspname = $1
|
|
AND t.typtype IN ('e', 'c', 'd', 'r')
|
|
AND NOT EXISTS (
|
|
SELECT 1 FROM pg_class c
|
|
WHERE c.reltype = t.oid AND c.relkind IN ('r', 'v', 'm', 'f', 'p')
|
|
)
|
|
`
|
|
typeRows, err := d.conn.QueryContext(context.Background(), typeQuery, d.schemaName)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to query types: %w", err)
|
|
}
|
|
defer typeRows.Close()
|
|
|
|
var typeNames []string
|
|
for typeRows.Next() {
|
|
var typeName string
|
|
if err := typeRows.Scan(&typeName); err != nil {
|
|
return fmt.Errorf("failed to scan type name: %w", err)
|
|
}
|
|
typeNames = append(typeNames, typeName)
|
|
}
|
|
if err := typeRows.Err(); err != nil {
|
|
return fmt.Errorf("failed to iterate types: %w", err)
|
|
}
|
|
|
|
for _, tn := range typeNames {
|
|
dropQuery := `DROP TYPE IF EXISTS ` + pq.QuoteIdentifier(d.schemaName) + `.` + pq.QuoteIdentifier(tn) + ` CASCADE`
|
|
if _, err := d.conn.ExecContext(context.Background(), dropQuery); err != nil {
|
|
return fmt.Errorf("failed to drop type %s: %w", tn, err)
|
|
}
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func computeLineFromPos(s string, pos int) (line uint, col uint, ok bool) {
|
|
s = strings.ReplaceAll(s, "\r\n", "\n")
|
|
runes := []rune(s)
|
|
if pos > len(runes) {
|
|
return 0, 0, false
|
|
}
|
|
sel := runes[:pos]
|
|
line = uint(runesCount(sel, '\n') + 1)
|
|
col = uint(pos - 1 - runesLastIndex(sel, '\n'))
|
|
return line, col, true
|
|
}
|
|
|
|
func runesCount(input []rune, target rune) int {
|
|
var count int
|
|
for _, r := range input {
|
|
if r == target {
|
|
count++
|
|
}
|
|
}
|
|
return count
|
|
}
|
|
|
|
func runesLastIndex(input []rune, target rune) int {
|
|
for i := len(input) - 1; i >= 0; i-- {
|
|
if input[i] == target {
|
|
return i
|
|
}
|
|
}
|
|
return -1
|
|
}
|