Files
license/internal/database/migrate.go
2026-07-31 16:59:55 +08:00

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
}