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 }