feat(db): 支持 MySQL 迁移(修复模型外键冲突 + 大文本类型 + 迁移工具)

- Attachment 模型移除 Message 关联:其 foreignKey 名 MessageID 与
  Message.MessageID 字符串字段冲突,GORM AutoMigrate 会生成错误外键
  (messages.message_id → attachments.id 且强转 bigint),MySQL 下建表
  直接失败(SQLite 因动态类型侥幸可用)
- Message.TextBody/HtmlBody 改 mediumtext:MySQL TEXT 仅 64KB,
  大 HTML 邮件会写入失败
- 新增 cmd/migrate:SQLite → MySQL 一次性迁移工具(GORM 模型读源、
  批量写目标、时间统一 UTC、ban_entries 零值时间转 NULL、逐表校验)
This commit is contained in:
dsh
2026-08-19 11:26:32 -04:00
parent f8fbe1ebdb
commit 07e81fc328
3 changed files with 219 additions and 3 deletions
+215
View File
@@ -0,0 +1,215 @@
// migrate 一次性工具:把 SQLite 数据迁移到 MySQLmailgo 库)。
// 用法:go run ./cmd/migrate -from /srv/mail_go/mail.db -dsn "mailgo:密码@tcp(127.0.0.1:3306)/mailgo?charset=utf8mb4&parseTime=True&loc=UTC"
package main
import (
"flag"
"fmt"
"log"
"time"
"mail_go/config"
"mail_go/internal/db"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
var (
fromDSN = flag.String("from", "/srv/mail_go/mail.db", "SQLite 数据库路径")
mysqlDSN = flag.String("dsn", "", "MySQL DSN(目标库,需已创建 mailgo 库与用户)")
)
func main() {
flag.Parse()
if *mysqlDSN == "" {
log.Fatal("缺少 -dsn")
}
// 目标:MySQLInitDB 内含 AutoMigrate,按当前模型建表)
mdb, err := db.InitDB(config.DatabaseConfig{Driver: "mysql", DSN: *mysqlDSN}, config.StorageConfig{BaseDir: "/srv/mail_go/"})
if err != nil {
log.Fatalf("连接 MySQL 失败: %v", err)
}
log.Println("MySQL 建表完成(AutoMigrate")
// 源:SQLite(只读)
sdb, err := db.InitDB(config.DatabaseConfig{Driver: "sqlite", DSN: *fromDSN}, config.StorageConfig{BaseDir: "/srv/mail_go/"})
if err != nil {
log.Fatalf("连接 SQLite 失败: %v", err)
}
sdb.Logger = logger.Default.LogMode(logger.Silent)
// 关闭 GORM 自动时间戳(保留原始 CreatedAt/UpdatedAt
mw := mdb.Session(&gorm.Session{SkipHooks: true})
// 按外键依赖顺序复制:domains → users → messages → attachments → 其余
// 所有时间统一 UTCMySQL DATETIME 无时区)。
utc := func(t time.Time) time.Time {
if t.IsZero() {
// MySQL DATETIME 最小年份 1000;零值由调用方转 NULL
return t
}
return t.UTC()
}
_ = utc
// ---- domains ----
var domains []db.Domain
if err := sdb.Order("id").Find(&domains).Error; err != nil {
log.Fatalf("读 domains: %v", err)
}
for i := range domains {
domains[i].CreatedAt = domains[i].CreatedAt.UTC()
domains[i].UpdatedAt = domains[i].UpdatedAt.UTC()
}
if err := mw.Create(&domains).Error; err != nil {
log.Fatalf("写 domains: %v", err)
}
log.Printf("domains: %d", len(domains))
// ---- users ----
var users []db.User
if err := sdb.Order("id").Find(&users).Error; err != nil {
log.Fatalf("读 users: %v", err)
}
for i := range users {
users[i].CreatedAt = users[i].CreatedAt.UTC()
users[i].UpdatedAt = users[i].UpdatedAt.UTC()
}
if err := mw.Create(&users).Error; err != nil {
log.Fatalf("写 users: %v", err)
}
log.Printf("users: %d", len(users))
// ---- messages ----
var msgs []db.Message
if err := sdb.Order("id").Find(&msgs).Error; err != nil {
log.Fatalf("读 messages: %v", err)
}
for i := range msgs {
msgs[i].Date = msgs[i].Date.UTC()
msgs[i].CreatedAt = msgs[i].CreatedAt.UTC()
}
if err := mw.Create(&msgs).Error; err != nil {
log.Fatalf("写 messages: %v", err)
}
log.Printf("messages: %d", len(msgs))
// ---- attachments ----
var atts []db.Attachment
if err := sdb.Order("id").Find(&atts).Error; err != nil {
log.Fatalf("读 attachments: %v", err)
}
for i := range atts {
atts[i].CreatedAt = atts[i].CreatedAt.UTC()
}
if err := mw.Create(&atts).Error; err != nil {
log.Fatalf("写 attachments: %v", err)
}
log.Printf("attachments: %d", len(atts))
// ---- outbound_messages(原样,含时间转 UTC----
var outs []db.OutboundMessage
if err := sdb.Order("id").Find(&outs).Error; err != nil {
log.Fatalf("读 outbound_messages: %v", err)
}
for i := range outs {
outs[i].NextAttemptAt = outs[i].NextAttemptAt.UTC()
if outs[i].CompletedAt != nil && !outs[i].CompletedAt.IsZero() {
u := outs[i].CompletedAt.UTC()
outs[i].CompletedAt = &u
}
outs[i].CreatedAt = outs[i].CreatedAt.UTC()
outs[i].UpdatedAt = outs[i].UpdatedAt.UTC()
}
if err := mw.Create(&outs).Error; err != nil {
log.Fatalf("写 outbound_messages: %v", err)
}
log.Printf("outbound_messages: %d", len(outs))
// ---- ban_entriesexpires_at 零值 → NULL----
rows, err := sdb.Raw("SELECT id, ip_address, reason, fail_count, ban_count, expires_at, created_at, updated_at FROM ban_entries ORDER BY id").Rows()
if err != nil {
log.Fatalf("读 ban_entries: %v", err)
}
defer rows.Close()
bans := 0
for rows.Next() {
var (
id uint
ip string
reason *string
failCount int
banCount int
expires *time.Time
created *time.Time
updated *time.Time
)
if err := rows.Scan(&id, &ip, &reason, &failCount, &banCount, &expires, &created, &updated); err != nil {
log.Fatalf("扫 ban_entries: %v", err)
}
norm := func(t *time.Time) *time.Time {
if t == nil || t.IsZero() {
return nil
}
u := t.UTC()
return &u
}
if err := mdb.Exec("INSERT INTO ban_entries (id, ip_address, reason, fail_count, ban_count, expires_at, created_at, updated_at) VALUES (?,?,?,?,?,?,?,?)",
id, ip, reason, failCount, banCount, norm(expires), norm(created), norm(updated)).Error; err != nil {
log.Fatalf("写 ban_entries id=%d: %v", id, err)
}
bans++
}
log.Printf("ban_entries: %d", bans)
// ---- protocol_logs ----
var logs []db.ProtocolLog
if err := sdb.Order("id").Find(&logs).Error; err != nil {
log.Fatalf("读 protocol_logs: %v", err)
}
for i := range logs {
logs[i].CreatedAt = logs[i].CreatedAt.UTC()
}
if err := mw.Create(&logs).Error; err != nil {
log.Fatalf("写 protocol_logs: %v", err)
}
log.Printf("protocol_logs: %d", len(logs))
// ---- mailbox_states ----
var states []db.MailboxState
if err := sdb.Order("user_id ASC, folder ASC").Find(&states).Error; err != nil {
log.Fatalf("读 mailbox_states: %v", err)
}
for i := range states {
states[i].CreatedAt = states[i].CreatedAt.UTC()
states[i].UpdatedAt = states[i].UpdatedAt.UTC()
}
if err := mw.Create(&states).Error; err != nil {
log.Fatalf("写 mailbox_states: %v", err)
}
log.Printf("mailbox_states: %d", len(states))
// ---- 校验 ----
check := func(table string, want int64) {
var got int64
if err := mdb.Table(table).Count(&got).Error; err != nil {
log.Fatalf("校验 %s: %v", table, err)
}
if got != want {
log.Fatalf("校验 %s 失败: got %d want %d", table, got, want)
}
fmt.Printf("校验 %s: %d/%d ✓\n", table, got, want)
}
check("domains", int64(len(domains)))
check("users", int64(len(users)))
check("messages", int64(len(msgs)))
check("attachments", int64(len(atts)))
check("outbound_messages", int64(len(outs)))
check("ban_entries", int64(bans))
check("protocol_logs", int64(len(logs)))
check("mailbox_states", int64(len(states)))
log.Println("迁移完成 ✅")
}