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 }