restructure project, add claudemd
This commit is contained in:
380
go/cmd/migrate/database.go
Normal file
380
go/cmd/migrate/database.go
Normal file
@@ -0,0 +1,380 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user