Files
TaskPool/internal/services/migration_v3.go
T
admin e6956aa001 Initial commit: TaskPool React panel
- React frontend with route-level code splitting
- Backend rebranded from Baihu to TaskPool
- DB brand migration script and local compatibility
2026-07-26 08:43:52 +08:00

447 lines
14 KiB
Go

package services
import (
"fmt"
"os"
"path/filepath"
"strings"
"github.com/engigu/taskpool/internal/database"
"github.com/engigu/taskpool/internal/logger"
"github.com/engigu/taskpool/internal/models"
"github.com/engigu/taskpool/internal/utils"
"gorm.io/gorm"
)
// MigrationTable 定义迁移配置
type MigrationTable struct {
Model any
EntityName string
FKs map[string]string // 单个外键映射: 字段名 -> 实体名
MultiFKs map[string]string // 复合外键映射 (如逗号分隔): 字段名 -> 实体名
}
func getMigrationTables() []MigrationTable {
return []MigrationTable{
{&models.User{}, "users", nil, nil},
{&models.Agent{}, "agents", nil, nil},
{&models.AgentToken{}, "tokens", nil, nil},
{&models.EnvironmentVariable{}, "envs", map[string]string{"UserID": "users"}, nil},
{&models.Task{}, "tasks", map[string]string{"AgentID": "agents"}, map[string]string{"Envs": "envs"}},
{&models.TaskLog{}, "task_logs", map[string]string{"TaskID": "tasks", "AgentID": "agents"}, nil},
{&models.Script{}, "scripts", map[string]string{"UserID": "users"}, nil},
{&models.Setting{}, "settings", nil, nil},
{&models.SendStats{}, "send_stats", map[string]string{"TaskID": "tasks"}, nil},
{&models.Language{}, "languages", nil, nil},
{&models.Dependency{}, "deps", nil, nil},
}
}
func getTableName(db *gorm.DB, model any) string {
stmt := &gorm.Statement{DB: db}
if err := stmt.Parse(model); err != nil {
return ""
}
return stmt.Schema.Table
}
func isTableStringID(db *gorm.DB, model any) bool {
if !db.Migrator().HasTable(model) {
return true
}
columnTypes, err := db.Migrator().ColumnTypes(model)
if err != nil {
return true
}
for _, ct := range columnTypes {
if strings.ToLower(ct.Name()) == "id" {
typeName := strings.ToLower(ct.DatabaseTypeName())
return strings.Contains(typeName, "char") || strings.Contains(typeName, "text") || strings.Contains(typeName, "string")
}
}
return false
}
func getValFromMap(m map[string]interface{}, key string) (interface{}, bool) {
lowerKey := strings.ToLower(key)
for k, v := range m {
if strings.ToLower(k) == lowerKey {
return v, true
}
}
return nil, false
}
func RunMigrationV3() error {
db := database.DB
if db == nil {
return fmt.Errorf("数据库未初始化")
}
// 0. 检查迁移标记,防止重复迁移逻辑被误判触发
if db.Migrator().HasTable(&models.Setting{}) {
var migrationFlag models.Setting
res := db.Where(&models.Setting{Section: "system", Key: "migration_v3_success"}).Limit(1).Find(&migrationFlag)
if res.Error == nil && res.RowsAffected > 0 && migrationFlag.Value == "true" {
// 如果已经是字符串 ID 模式,双重确认
logger.Info("[MigrationV3] 系统已处于 V3 模式,跳过检查")
return nil
}
}
tables := getMigrationTables()
needMigration := false
for _, t := range tables {
if !isTableStringID(db, t.Model) {
needMigration = true
logger.Infof("[MigrationV3] 表 [%s] ID 为数字,需迁移", getTableName(db, t.Model))
break
}
}
if !needMigration {
// 没有数字 ID 表了,但也可能还没打标(比如之前手动修过表),此时补一个标
return markMigrationSuccess(db)
}
// 1. 备份过程
backupDir := "./data"
os.MkdirAll(backupDir, 0755)
// 如果已经有还原过的记录,为了防止反复循环,我们可以检查还原标记
backups, _ := filepath.Glob(filepath.Join(backupDir, "migration_v3_backup_*.zip"))
if len(backups) == 0 {
logger.Infof("[MigrationV3] 执行关键备份...")
backupService := NewBackupService()
zipPath, err := backupService.CreateBackup()
if err != nil {
return fmt.Errorf("自动备份失败,流程终止: %v", err)
}
baseName := strings.TrimSuffix(filepath.Base(zipPath), ".zip")
newPath := filepath.Join(backupDir, fmt.Sprintf("migration_v3_backup_%s.zip", baseName))
os.Rename(zipPath, newPath)
logger.Infof("[MigrationV3] 备份成功: %s", newPath)
}
mappings := make(map[string]map[uint]string)
// 根据数据库类型决定是否使用事务:
// - PostgreSQL: DDL 完全支持事务,使用事务保证原子性
// - MySQL: DDL 会隐式提交事务,包裹事务无意义
// - SQLite: DDL+DML 混合在 GORM 事务中会导致数据丢失
dbType := db.Dialector.Name()
if dbType == "postgres" {
err := db.Transaction(func(tx *gorm.DB) error {
return performHardMigration(tx, mappings)
})
if err != nil {
return err
}
} else {
// SQLite / MySQL: 不使用事务包裹
if err := performHardMigration(db, mappings); err != nil {
return err
}
}
// 3. 标记成功
return markMigrationSuccess(db)
}
func markMigrationSuccess(db *gorm.DB) error {
if !db.Migrator().HasTable(&models.Setting{}) {
return nil
}
var flag models.Setting
res := db.Where(&models.Setting{Section: "system", Key: "migration_v3_success"}).Limit(1).Find(&flag)
if res.Error != nil || res.RowsAffected == 0 {
// 创建或更新
flag = models.Setting{
ID: utils.GenerateID(),
Section: "system",
Key: "migration_v3_success",
Value: "true",
}
return db.Create(&flag).Error
}
return db.Model(&flag).Update("value", "true").Error
}
func performHardMigration(db *gorm.DB, mappings map[string]map[uint]string) error {
allTables := getMigrationTables()
// ---------------------------------------------------------
// 第一阶段:全量构建 ID 映射映射表 (Pass 1)
// ---------------------------------------------------------
for _, t := range allTables {
actualName := getTableName(db, t.Model)
if actualName == "" || !db.Migrator().HasTable(actualName) {
continue
}
mappings[t.EntityName] = make(map[uint]string)
oldTableName := actualName + "_v2_bak"
// 如果还没有备份表,说明这是第一次处理该表,先重命名
if !db.Migrator().HasTable(oldTableName) {
if isTableStringID(db, t.Model) {
continue // 已经是字符串 ID 且无备份,跳过
}
if err := db.Migrator().RenameTable(actualName, oldTableName); err != nil {
return fmt.Errorf("重命名表 %s 失败: %v", actualName, err)
}
// SQLite 重命名表后,索引名称不变,会导致 AutoMigrate 创建新表时索引冲突
// 需要先删除备份表上的旧索引
dropOldIndexes(db, oldTableName)
}
// 预先为该表所有记录生成新的 xid
var rows []map[string]interface{}
db.Table(oldTableName).Select("id").Find(&rows)
for _, row := range rows {
if val, ok := getValFromMap(row, "id"); ok {
uid := parseUint(val)
if uid > 0 {
mappings[t.EntityName][uid] = utils.GenerateID()
}
}
}
logger.Infof("[MigrationV3] Pass 1: 构建关键表 %s 的 ID 映射, 共 %d 条", actualName, len(mappings[t.EntityName]))
}
// 构建 fallback ID:FK 映射找不到时使用第一条记录的新 ID
fallbackIDs := make(map[string]string)
for _, t := range allTables {
// 优先从 mappings 中取第一个新 xid(本次迁移生成的)
if m, ok := mappings[t.EntityName]; ok && len(m) > 0 {
for _, newID := range m {
fallbackIDs[t.EntityName] = newID
break
}
continue
}
// 如果该表已经是 string ID(被跳过了),从实际表查第一条
actualName := getTableName(db, t.Model)
if actualName != "" && db.Migrator().HasTable(actualName) {
var row map[string]interface{}
if err := db.Table(actualName).Select("id").Order("id").Limit(1).Find(&row).Error; err == nil {
if id, ok := getValFromMap(row, "id"); ok && id != nil {
fallbackIDs[t.EntityName] = fmt.Sprintf("%v", id)
}
}
}
}
// ---------------------------------------------------------
// 第二、三阶段:正式转换数据并处理关联字段 (Pass 2 & 3)
// ---------------------------------------------------------
for _, t := range allTables {
actualName := getTableName(db, t.Model)
oldTableName := actualName + "_v2_bak"
if !db.Migrator().HasTable(oldTableName) {
// 虽然可能已经改过格式,但为了安全还是 AutoMigrate 一下
db.AutoMigrate(t.Model)
continue
}
logger.Infof("[MigrationV3] Pass 2&3: 正在转换数据并修复关联: %s", actualName)
db.AutoMigrate(t.Model)
// 获取新表的有效列名(小写)
columnTypes, _ := db.Migrator().ColumnTypes(t.Model)
validColumns := make(map[string]bool)
for _, ct := range columnTypes {
validColumns[strings.ToLower(ct.Name())] = true
}
var oldData []map[string]interface{}
if err := db.Table(oldTableName).Find(&oldData).Error; err != nil {
return err
}
for _, row := range oldData {
// 1. 处理主键 ID
if val, ok := getValFromMap(row, "id"); ok {
uid := parseUint(val)
if nid, exists := mappings[t.EntityName][uid]; exists {
row["id"] = nid
} else {
row["id"] = utils.GenerateID()
}
}
// 2. 处理单外键关联 (Phase 3)
for field, parentEntity := range t.FKs {
columnName := getColumnName(field)
if val, ok := getValFromMap(row, columnName); ok && val != nil {
strVal := fmt.Sprintf("%v", val)
if strVal == "" || strVal == "0" || strVal == "<nil>" {
row[columnName] = nil
continue
}
// 如果值已经是非数字字符串(例如已迁移过的 xid),直接保留
if !utils.IsNumeric(strVal) {
continue
}
// 数字型外键,查映射表转换
ufk := parseUint(val)
if ufk > 0 {
if nid, exists := mappings[parentEntity][ufk]; exists {
row[columnName] = nid
} else if fbID, hasFB := fallbackIDs[parentEntity]; hasFB {
logger.Warnf("[MigrationV3] 表 %s 的 %s=%d 在 %s 映射中未找到,使用首条记录 ID: %s", getTableName(db, t.Model), columnName, ufk, parentEntity, fbID)
row[columnName] = fbID
} else {
logger.Warnf("[MigrationV3] 表 %s 的 %s=%d 在 %s 映射中未找到且无 fallback,置空", getTableName(db, t.Model), columnName, ufk, parentEntity)
row[columnName] = nil
}
} else {
row[columnName] = nil
}
}
}
// 3. 处理复合多外键字段 (envs)
for field, parentEntity := range t.MultiFKs {
columnName := getColumnName(field)
if val, ok := getValFromMap(row, columnName); ok && val != nil {
if strVal, ok := val.(string); ok && strVal != "" {
row[columnName] = transformMultiIDs(strVal, parentEntity, mappings)
}
}
}
// 4. 过滤不存在的列并插入
filteredRow := make(map[string]interface{})
for k, v := range row {
if validColumns[strings.ToLower(k)] {
filteredRow[k] = v
}
}
if err := db.Table(actualName).Create(filteredRow).Error; err != nil {
return err
}
}
// 迁移完成,清理备份表
db.Migrator().DropTable(oldTableName)
}
return nil
}
// 辅助函数:解析各种数字 ID
func parseUint(val interface{}) uint {
if val == nil {
return 0
}
switch v := val.(type) {
case uint:
return v
case int64:
return uint(v)
case int:
return uint(v)
case uint64:
return uint(v)
case float64:
return uint(v)
case string:
var u uint
fmt.Sscanf(v, "%d", &u)
return u
}
return 0
}
// 辅助函数:字段名转列名
func getColumnName(field string) string {
switch field {
case "AgentID":
return "agent_id"
case "TaskID":
return "task_id"
case "UserID":
return "user_id"
case "LogID":
return "log_id"
case "Envs":
return "envs"
default:
return strings.ToLower(field)
}
}
// 辅助函数:处理逗号分隔的 ID 列表
func transformMultiIDs(oldStr string, parentEntity string, mappings map[string]map[uint]string) string {
parts := strings.Split(oldStr, ",")
var result []string
for _, p := range parts {
p = strings.TrimSpace(p)
if p == "" {
continue
}
if len(p) == 20 && !utils.IsNumeric(p) {
result = append(result, p) // 已经是 xid,保留
continue
}
uid := parseUint(p)
if uid > 0 {
if nid, exists := mappings[parentEntity][uid]; exists {
result = append(result, nid)
}
}
}
return strings.Join(result, ",")
}
// dropOldIndexes 删除备份表上的旧索引,防止 AutoMigrate 创建新表时索引名冲突
// SQLite 重命名表后索引名不变,MySQL/PostgreSQL 也可能存在类似问题
func dropOldIndexes(db *gorm.DB, tableName string) {
dbType := db.Dialector.Name()
var indexNames []string
switch dbType {
case "sqlite":
var indexes []struct{ Name string }
db.Raw("SELECT name FROM sqlite_master WHERE type='index' AND tbl_name=? AND name NOT LIKE 'sqlite_%'", tableName).Scan(&indexes)
for _, idx := range indexes {
indexNames = append(indexNames, idx.Name)
}
case "mysql":
var indexes []struct {
KeyName string `gorm:"column:Key_name"`
}
db.Raw("SHOW INDEX FROM `" + tableName + "`").Scan(&indexes)
seen := make(map[string]bool)
for _, idx := range indexes {
if idx.KeyName != "PRIMARY" && !seen[idx.KeyName] {
indexNames = append(indexNames, idx.KeyName)
seen[idx.KeyName] = true
}
}
case "postgres":
var indexes []struct {
IndexName string `gorm:"column:indexname"`
}
db.Raw("SELECT indexname FROM pg_indexes WHERE tablename=?", tableName).Scan(&indexes)
for _, idx := range indexes {
if !strings.HasSuffix(idx.IndexName, "_pkey") {
indexNames = append(indexNames, idx.IndexName)
}
}
}
for _, name := range indexNames {
var dropSQL string
switch dbType {
case "mysql":
dropSQL = fmt.Sprintf("DROP INDEX `%s` ON `%s`", name, tableName)
default:
dropSQL = fmt.Sprintf("DROP INDEX IF EXISTS \"%s\"", name)
}
if err := db.Exec(dropSQL).Error; err != nil {
logger.Warnf("[MigrationV3] 删除旧索引 %s 失败 (可忽略): %v", name, err)
} else {
logger.Infof("[MigrationV3] 已删除备份表旧索引: %s", name)
}
}
}