228 lines
5.8 KiB
Go
228 lines
5.8 KiB
Go
|
|
package database
|
||
|
|
|
||
|
|
import (
|
||
|
|
"context"
|
||
|
|
"database/sql"
|
||
|
|
"embed"
|
||
|
|
"errors"
|
||
|
|
"fmt"
|
||
|
|
"io/fs"
|
||
|
|
"path"
|
||
|
|
"sort"
|
||
|
|
"strconv"
|
||
|
|
"strings"
|
||
|
|
"time"
|
||
|
|
|
||
|
|
"gorm.io/gorm"
|
||
|
|
)
|
||
|
|
|
||
|
|
const (
|
||
|
|
migrationLockID int64 = 67549013341001
|
||
|
|
migrationUnlockTimeout = 5 * time.Second
|
||
|
|
)
|
||
|
|
|
||
|
|
const createSchemaMigrationsSQL = `
|
||
|
|
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||
|
|
version BIGINT PRIMARY KEY,
|
||
|
|
name VARCHAR(255) NOT NULL,
|
||
|
|
applied_at TIMESTAMPTZ NOT NULL
|
||
|
|
)`
|
||
|
|
|
||
|
|
//go:embed migrations/*.sql
|
||
|
|
var migrationFiles embed.FS
|
||
|
|
|
||
|
|
type migration struct {
|
||
|
|
version int64
|
||
|
|
name string
|
||
|
|
sql string
|
||
|
|
}
|
||
|
|
|
||
|
|
type appliedMigration struct {
|
||
|
|
version int64
|
||
|
|
name string
|
||
|
|
}
|
||
|
|
|
||
|
|
func Migrate(ctx context.Context, db *gorm.DB) (resultErr error) {
|
||
|
|
if db == nil {
|
||
|
|
return errors.New("数据库实例不能为空")
|
||
|
|
}
|
||
|
|
|
||
|
|
migrations, err := loadMigrations()
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
|
||
|
|
sqlDB, err := db.DB()
|
||
|
|
if err != nil {
|
||
|
|
return fmt.Errorf("获取数据库连接池: %w", err)
|
||
|
|
}
|
||
|
|
conn, err := sqlDB.Conn(ctx)
|
||
|
|
if err != nil {
|
||
|
|
return fmt.Errorf("获取迁移专用连接: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
lockAcquired := false
|
||
|
|
defer func() {
|
||
|
|
cleanupContext, cancel := context.WithTimeout(context.Background(), migrationUnlockTimeout)
|
||
|
|
defer cancel()
|
||
|
|
|
||
|
|
var unlocked bool
|
||
|
|
unlockErr := conn.QueryRowContext(
|
||
|
|
cleanupContext,
|
||
|
|
"SELECT pg_advisory_unlock($1)",
|
||
|
|
migrationLockID,
|
||
|
|
).Scan(&unlocked)
|
||
|
|
if unlockErr != nil {
|
||
|
|
unlockErr = fmt.Errorf("释放数据库迁移锁: %w", unlockErr)
|
||
|
|
} else if lockAcquired && !unlocked {
|
||
|
|
unlockErr = errors.New("数据库迁移锁未由当前连接持有")
|
||
|
|
}
|
||
|
|
|
||
|
|
closeErr := conn.Close()
|
||
|
|
if closeErr != nil {
|
||
|
|
closeErr = fmt.Errorf("关闭迁移专用连接: %w", closeErr)
|
||
|
|
}
|
||
|
|
resultErr = errors.Join(resultErr, unlockErr, closeErr)
|
||
|
|
}()
|
||
|
|
|
||
|
|
if _, err := conn.ExecContext(ctx, "SELECT pg_advisory_lock($1)", migrationLockID); err != nil {
|
||
|
|
return fmt.Errorf("获取数据库迁移锁: %w", err)
|
||
|
|
}
|
||
|
|
lockAcquired = true
|
||
|
|
|
||
|
|
if _, err := conn.ExecContext(ctx, createSchemaMigrationsSQL); err != nil {
|
||
|
|
return fmt.Errorf("创建迁移版本表: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
applied, err := readAppliedMigrations(ctx, conn)
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
if err := validateAppliedMigrations(migrations, applied); err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
|
||
|
|
for _, item := range migrations[len(applied):] {
|
||
|
|
if err := applyMigration(ctx, conn, item); err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func loadMigrations() ([]migration, error) {
|
||
|
|
fileNames, err := fs.Glob(migrationFiles, "migrations/*.sql")
|
||
|
|
if err != nil {
|
||
|
|
return nil, fmt.Errorf("读取内嵌迁移文件列表: %w", err)
|
||
|
|
}
|
||
|
|
if len(fileNames) == 0 {
|
||
|
|
return nil, errors.New("未找到内嵌迁移文件")
|
||
|
|
}
|
||
|
|
|
||
|
|
migrations := make([]migration, 0, len(fileNames))
|
||
|
|
for _, fileName := range fileNames {
|
||
|
|
item, err := loadMigration(fileName)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
migrations = append(migrations, item)
|
||
|
|
}
|
||
|
|
sort.Slice(migrations, func(i, j int) bool {
|
||
|
|
return migrations[i].version < migrations[j].version
|
||
|
|
})
|
||
|
|
for index := 1; index < len(migrations); index++ {
|
||
|
|
if migrations[index-1].version == migrations[index].version {
|
||
|
|
return nil, fmt.Errorf("迁移版本 %d 重复", migrations[index].version)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return migrations, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func loadMigration(fileName string) (migration, error) {
|
||
|
|
baseName := strings.TrimSuffix(path.Base(fileName), ".sql")
|
||
|
|
parts := strings.SplitN(baseName, "_", 2)
|
||
|
|
if len(parts) != 2 || parts[1] == "" {
|
||
|
|
return migration{}, fmt.Errorf("迁移文件名 %q 无效", fileName)
|
||
|
|
}
|
||
|
|
|
||
|
|
version, err := strconv.ParseInt(parts[0], 10, 64)
|
||
|
|
if err != nil || version <= 0 {
|
||
|
|
return migration{}, fmt.Errorf("迁移文件名 %q 的版本无效", fileName)
|
||
|
|
}
|
||
|
|
source, err := migrationFiles.ReadFile(fileName)
|
||
|
|
if err != nil {
|
||
|
|
return migration{}, fmt.Errorf("读取迁移文件 %q: %w", fileName, err)
|
||
|
|
}
|
||
|
|
if strings.TrimSpace(string(source)) == "" {
|
||
|
|
return migration{}, fmt.Errorf("迁移文件 %q 不能为空", fileName)
|
||
|
|
}
|
||
|
|
|
||
|
|
return migration{
|
||
|
|
version: version,
|
||
|
|
name: parts[1],
|
||
|
|
sql: string(source),
|
||
|
|
}, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func readAppliedMigrations(ctx context.Context, conn *sql.Conn) ([]appliedMigration, error) {
|
||
|
|
rows, err := conn.QueryContext(ctx, "SELECT version, name FROM schema_migrations ORDER BY version")
|
||
|
|
if err != nil {
|
||
|
|
return nil, fmt.Errorf("读取已应用迁移: %w", err)
|
||
|
|
}
|
||
|
|
defer rows.Close()
|
||
|
|
|
||
|
|
var applied []appliedMigration
|
||
|
|
for rows.Next() {
|
||
|
|
var item appliedMigration
|
||
|
|
if err := rows.Scan(&item.version, &item.name); err != nil {
|
||
|
|
return nil, fmt.Errorf("解析已应用迁移: %w", err)
|
||
|
|
}
|
||
|
|
applied = append(applied, item)
|
||
|
|
}
|
||
|
|
if err := rows.Err(); err != nil {
|
||
|
|
return nil, fmt.Errorf("遍历已应用迁移: %w", err)
|
||
|
|
}
|
||
|
|
return applied, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func validateAppliedMigrations(embedded []migration, applied []appliedMigration) error {
|
||
|
|
if len(applied) > len(embedded) {
|
||
|
|
return errors.New("数据库包含未知迁移版本")
|
||
|
|
}
|
||
|
|
for index, item := range applied {
|
||
|
|
expected := embedded[index]
|
||
|
|
if item.version != expected.version {
|
||
|
|
return fmt.Errorf("数据库迁移历史不是内嵌迁移的精确前缀,位置 %d 版本不匹配", index+1)
|
||
|
|
}
|
||
|
|
if item.name != expected.name {
|
||
|
|
return fmt.Errorf("数据库迁移版本 %d 的名称不匹配", item.version)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func applyMigration(ctx context.Context, conn *sql.Conn, item migration) error {
|
||
|
|
tx, err := conn.BeginTx(ctx, nil)
|
||
|
|
if err != nil {
|
||
|
|
return fmt.Errorf("开始迁移 %d 事务: %w", item.version, err)
|
||
|
|
}
|
||
|
|
defer tx.Rollback()
|
||
|
|
|
||
|
|
if _, err := tx.ExecContext(ctx, item.sql); err != nil {
|
||
|
|
return fmt.Errorf("执行迁移 %d: %w", item.version, err)
|
||
|
|
}
|
||
|
|
if _, err := tx.ExecContext(
|
||
|
|
ctx,
|
||
|
|
"INSERT INTO schema_migrations (version, name, applied_at) VALUES ($1, $2, NOW())",
|
||
|
|
item.version,
|
||
|
|
item.name,
|
||
|
|
); err != nil {
|
||
|
|
return fmt.Errorf("记录迁移 %d: %w", item.version, err)
|
||
|
|
}
|
||
|
|
if err := tx.Commit(); err != nil {
|
||
|
|
return fmt.Errorf("提交迁移 %d: %w", item.version, err)
|
||
|
|
}
|
||
|
|
return nil
|
||
|
|
}
|