67 lines
1.8 KiB
Go
67 lines
1.8 KiB
Go
package database
|
|
|
|
import (
|
|
"errors"
|
|
"fmt"
|
|
"strconv"
|
|
|
|
"github.com/golang-migrate/migrate/v4"
|
|
"github.com/golang-migrate/migrate/v4/database/postgres"
|
|
"github.com/golang-migrate/migrate/v4/source/iofs"
|
|
"silk-server-go/migrations"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
// CurrentSchemaVersion 是当前后端代码期望的迁移版本。
|
|
const CurrentSchemaVersion = "1"
|
|
|
|
// RunMigrations 使用嵌入式 SQL 迁移文件将数据库升级到最新版本。
|
|
func RunMigrations(db *gorm.DB) error {
|
|
sqlDB, err := db.DB()
|
|
if err != nil {
|
|
return fmt.Errorf("获取数据库连接: %w", err)
|
|
}
|
|
|
|
sourceDriver, err := iofs.New(migrations.FS, ".")
|
|
if err != nil {
|
|
return fmt.Errorf("加载嵌入式迁移文件: %w", err)
|
|
}
|
|
|
|
databaseDriver, err := postgres.WithInstance(sqlDB, &postgres.Config{})
|
|
if err != nil {
|
|
return fmt.Errorf("初始化 PostgreSQL 迁移驱动: %w", err)
|
|
}
|
|
|
|
m, err := migrate.NewWithInstance("iofs", sourceDriver, "postgres", databaseDriver)
|
|
if err != nil {
|
|
return fmt.Errorf("创建迁移实例: %w", err)
|
|
}
|
|
defer func() {
|
|
_, _ = m.Close()
|
|
}()
|
|
|
|
if err := m.Up(); err != nil && !errors.Is(err, migrate.ErrNoChange) {
|
|
return fmt.Errorf("执行数据库迁移: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// CheckSchemaVersion 校验 schema_migrations 当前版本与 expected 一致,且不是 dirty。
|
|
func CheckSchemaVersion(db *gorm.DB, expected string) error {
|
|
var version int64
|
|
var dirty bool
|
|
|
|
row := db.Raw("SELECT version, dirty FROM schema_migrations ORDER BY version DESC LIMIT 1").Row()
|
|
if err := row.Scan(&version, &dirty); err != nil {
|
|
return fmt.Errorf("读取 schema_migrations: %w", err)
|
|
}
|
|
if dirty {
|
|
return fmt.Errorf("schema 迁移处于 dirty 状态,版本 %d", version)
|
|
}
|
|
if strconv.FormatInt(version, 10) != expected {
|
|
return fmt.Errorf("schema 版本不匹配:当前 %d,期望 %s", version, expected)
|
|
}
|
|
return nil
|
|
}
|