Files
2026-09-20 00:19:08 +08:00

356 lines
9.9 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 全量迁移:把业务表中的老 uid(yz_system_tenant_user.uid)统一替换为
// 认证中心的 identity_id(yz_auth_identity.id)。
//
// 背景:统一认证上线后,令牌里的 user_id 是 identity_id;而业务表的
// uid / user_id / create_user_id / owner_user_id / uploader_id 等字段存的仍是
// 老表 uid。不迁移的话用户会「找不到自己的数据」,因此做全量替换而非兼容层。
//
// 用法(需在 go/ 目录运行,读取 conf/app.conf):
//
// go run scripts/uidmigrate/migrate_uid.go -check 预检:列出受影响的表/列/行数
// go run scripts/uidmigrate/migrate_uid.go -apply 执行迁移(记录变更日志,可回滚)
// go run scripts/uidmigrate/migrate_uid.go -rollback 按变更日志回滚
//
// 说明:本脚本直接使用 database/sql,避免 beego ORM 的 Raw 取值差异。
// 执行前请先自行备份数据库(mysqldump)。
package main
import (
"database/sql"
"flag"
"fmt"
"log"
"os"
"strings"
"time"
beego "github.com/beego/beego/v2/server/web"
_ "github.com/go-sql-driver/mysql"
)
// tempOffset 中间值偏移:先把旧值搬到远离目标值的区间,避免新旧值重叠导致错改
const tempOffset = 1000000000
// 需要处理的列名
var uidColumns = []string{
"uid", "user_id", "create_user_id", "update_user_id",
"owner_user_id", "uploader_id", "operator_id", "tuid",
}
// 不参与迁移的表
var skipTables = map[string]bool{
"yz_system_tenant_user": true,
"yz_users": true,
"yz_uid_migration_log": true,
}
type colRef struct {
Table string
Col string
}
var db *sql.DB
func main() {
check := flag.Bool("check", false, "预检,不写数据")
apply := flag.Bool("apply", false, "执行迁移")
rollback := flag.Bool("rollback", false, "按变更日志回滚")
flag.Parse()
if !*check && !*apply && !*rollback {
fmt.Println("请指定 -check / -apply / -rollback")
os.Exit(1)
}
if err := beego.LoadAppConfig("ini", "conf/app.conf"); err != nil {
log.Printf("加载 conf/app.conf 失败(若已自动加载可忽略): %v", err)
}
var err error
db, err = openDB()
if err != nil {
log.Fatalf("连接数据库失败: %v", err)
}
defer db.Close()
switch {
case *check:
runCheck()
case *apply:
runApply()
case *rollback:
runRollback()
}
}
func openDB() (*sql.DB, error) {
user, _ := beego.AppConfig.String("mysqluser")
pass, _ := beego.AppConfig.String("mysqlpass")
urls, _ := beego.AppConfig.String("mysqlurls")
name, _ := beego.AppConfig.String("mysqldb")
if user == "" || urls == "" || name == "" {
return nil, fmt.Errorf("数据库配置(mysqluser/mysqlurls/mysqldb) 未正确设置")
}
dsn := fmt.Sprintf("%s:%s@tcp(%s)/%s?charset=utf8mb4&parseTime=True&loc=Local", user, pass, urls, name)
return sql.Open("mysql", dsn)
}
// buildMapping 建立「老 uid → identity_id」映射:同企业(tid) 下账号/手机/邮箱 一致。
//
// 注意:老表是 utf8mb4_0900_ai_ci、新表是 utf8mb4_unicode_ci,直接 JOIN 比较字符串
// 会触发 Illegal mix of collations,因此改为在 Go 侧匹配。
func buildMapping() (map[uint64]uint64, error) {
legacy := make([]struct {
uid uint64
tid uint64
account, phone string
email string
}, 0)
rows, err := db.Query(
"SELECT uid, tid, IFNULL(account,''), IFNULL(phone,''), IFNULL(email,'') FROM yz_system_tenant_user WHERE delete_time IS NULL")
if err != nil {
return nil, err
}
for rows.Next() {
var it struct {
uid uint64
tid uint64
account, phone string
email string
}
if err := rows.Scan(&it.uid, &it.tid, &it.account, &it.phone, &it.email); err != nil {
continue
}
legacy = append(legacy, it)
}
rows.Close()
index := map[string]uint64{}
rows2, err := db.Query(
"SELECT identity_id, tid, IFNULL(account,''), IFNULL(phone,''), IFNULL(email,'') FROM yz_auth_tenant_user WHERE delete_time IS NULL")
if err != nil {
return nil, err
}
for rows2.Next() {
var identity, tid uint64
var account, phone, email string
if err := rows2.Scan(&identity, &tid, &account, &phone, &email); err != nil {
continue
}
if key := mergeKey(tid, account, phone, email); key != "" {
index[key] = identity
}
}
rows2.Close()
mapping := map[uint64]uint64{}
for _, it := range legacy {
key := mergeKey(it.tid, it.account, it.phone, it.email)
if key == "" {
continue
}
if id, ok := index[key]; ok && id > 0 {
mapping[it.uid] = id
}
}
return mapping, nil
}
// mergeKey 归并键:企业ID + (手机号 > 邮箱 > 账号)
func mergeKey(tid uint64, account, phone, email string) string {
base := fmt.Sprintf("%d:", tid)
if v := strings.TrimSpace(phone); v != "" {
return base + "p:" + v
}
if v := strings.TrimSpace(email); v != "" {
return base + "e:" + v
}
if v := strings.TrimSpace(account); v != "" {
return base + "a:" + v
}
return ""
}
func findColumns() ([]colRef, error) {
name, _ := beego.AppConfig.String("mysqldb")
quoted := make([]string, 0, len(uidColumns))
for _, c := range uidColumns {
quoted = append(quoted, "'"+c+"'")
}
rows, err := db.Query(fmt.Sprintf(`
SELECT TABLE_NAME, COLUMN_NAME
FROM information_schema.COLUMNS
WHERE TABLE_SCHEMA = '%s' AND COLUMN_NAME IN (%s)
ORDER BY TABLE_NAME, COLUMN_NAME`, name, strings.Join(quoted, ",")))
if err != nil {
return nil, err
}
defer rows.Close()
out := make([]colRef, 0)
total := 0
for rows.Next() {
var t, c string
if err := rows.Scan(&t, &c); err != nil {
continue
}
total++
if skipTables[t] || strings.Contains(t, "_bak") || strings.HasPrefix(t, "yz_auth_") {
continue
}
out = append(out, colRef{Table: t, Col: c})
}
log.Printf("扫描到 %d 个承载用户ID的列(跳过 %d 个)", len(out), total-len(out))
return out, rows.Err()
}
func runCheck() {
mapping, err := buildMapping()
if err != nil {
log.Fatalf("建立映射失败: %v", err)
}
if len(mapping) == 0 {
log.Fatal("没有建立任何 uid → identity_id 映射,请确认认证中心数据是否已迁移")
}
log.Printf("映射关系: %d 个老 uid", len(mapping))
for old, newUID := range mapping {
log.Printf(" uid %d → identity %d", old, newUID)
}
cols, err := findColumns()
if err != nil {
log.Fatalf("扫描列失败: %v", err)
}
in := joinKeys(mapping)
total := int64(0)
affected := 0
for _, c := range cols {
var cnt int64
err := db.QueryRow(fmt.Sprintf(
"SELECT COUNT(*) FROM `%s` WHERE `%s` IN (%s)", c.Table, c.Col, in)).Scan(&cnt)
if err != nil {
continue
}
if cnt > 0 {
log.Printf(" %-42s %-18s %d 行", c.Table, c.Col, cnt)
total += cnt
affected++
}
}
log.Printf("合计:%d 张表的列需要更新,约 %d 行(预检完成,未写入任何数据)", affected, total)
}
func runApply() {
mapping, err := buildMapping()
if err != nil {
log.Fatalf("建立映射失败: %v", err)
}
cols, err := findColumns()
if err != nil {
log.Fatalf("扫描列失败: %v", err)
}
ensureLogTable()
batch := time.Now().Format("20060102150405")
changed := 0
for _, c := range cols {
for oldUID, newUID := range mapping {
if oldUID == newUID {
continue
}
// 阶段一:旧值 → 临时值;阶段二:临时值 → 新值
cnt := execUpdate(c.Table, c.Col, oldUID, oldUID+tempOffset)
if cnt == 0 {
continue
}
execUpdate(c.Table, c.Col, oldUID+tempOffset, newUID)
saveLog(batch, c.Table, c.Col, oldUID, newUID, cnt)
changed++
log.Printf(" %s.%s: %d → %d(%d 行)", c.Table, c.Col, oldUID, newUID, cnt)
}
}
log.Printf("迁移完成:%d 处替换,批次号 %s", changed, batch)
log.Println("如需回滚:go run scripts/uidmigrate/migrate_uid.go -rollback")
}
func runRollback() {
rows, err := db.Query(
"SELECT table_name, column_name, old_value, new_value FROM yz_uid_migration_log ORDER BY id DESC")
if err != nil {
log.Fatalf("读取变更日志失败: %v", err)
}
defer rows.Close()
type item struct {
table, col string
oldV, newV uint64
}
list := make([]item, 0)
for rows.Next() {
var it item
if err := rows.Scan(&it.table, &it.col, &it.oldV, &it.newV); err != nil {
continue
}
list = append(list, it)
}
if len(list) == 0 {
log.Println("没有可回滚的记录")
return
}
for _, it := range list {
// 反向:新值 → 临时值 → 旧值
execUpdate(it.table, it.col, it.newV, it.newV+tempOffset)
cnt := execUpdate(it.table, it.col, it.newV+tempOffset, it.oldV)
log.Printf(" 回滚 %s.%s: %d → %d(%d 行)", it.table, it.col, it.newV, it.oldV, cnt)
}
log.Println("回滚完成")
}
// ---------------------------------------------------------------- 工具
func execUpdate(table, col string, from, to uint64) int64 {
res, err := db.Exec(fmt.Sprintf(
"UPDATE `%s` SET `%s` = ? WHERE `%s` = ?", table, col, col), to, from)
if err != nil {
log.Printf("更新失败 %s.%s (%d→%d): %v", table, col, from, to, err)
return 0
}
n, _ := res.RowsAffected()
return n
}
func saveLog(batch, table, col string, oldV, newV uint64, rows int64) {
_, err := db.Exec(
"INSERT INTO yz_uid_migration_log (batch_no, table_name, column_name, old_value, new_value, row_count, create_time) VALUES (?,?,?,?,?,?,NOW())",
batch, table, col, oldV, newV, rows)
if err != nil {
log.Printf("记录变更日志失败: %v", err)
}
}
func ensureLogTable() {
_, _ = db.Exec(`CREATE TABLE IF NOT EXISTS yz_uid_migration_log (
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT,
batch_no VARCHAR(32) NOT NULL,
table_name VARCHAR(128) NOT NULL,
column_name VARCHAR(64) NOT NULL,
old_value BIGINT UNSIGNED NOT NULL,
new_value BIGINT UNSIGNED NOT NULL,
row_count BIGINT NOT NULL DEFAULT 0,
create_time DATETIME NOT NULL,
PRIMARY KEY (id),
KEY idx_batch (batch_no)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`)
}
func joinKeys(m map[uint64]uint64) string {
parts := make([]string, 0, len(m))
for k := range m {
parts = append(parts, fmt.Sprintf("%d", k))
}
if len(parts) == 0 {
return "0"
}
return strings.Join(parts, ",")
}