package model import ( "crypto/sha256" "encoding/hex" "fmt" "log" "gorm.io/driver/postgres" "gorm.io/gorm" "gorm.io/gorm/logger" ) // DB 全局数据库实例 var DB *gorm.DB // InitDB 连接 PostgreSQL 并自动迁移 func InitDB(dsn string) error { db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{ Logger: logger.Default.LogMode(logger.Warn), }) if err != nil { return fmt.Errorf("连接 PostgreSQL 失败: %w", err) } // 旧表 refresh_tokens.token(明文)→ token_hash 体系:必须在 AutoMigrate // 创建 NOT NULL 列/唯一索引之前完成回填 if err := prepareRefreshTokenMigration(db); err != nil { return fmt.Errorf("refresh token 旧数据迁移失败: %w", err) } if err := db.AutoMigrate( &User{}, &Board{}, &Post{}, &Comment{}, &RefreshToken{}, &Like{}, &Notification{}, &Checkin{}, &Announcement{}, &SiteSetting{}, ); err != nil { return fmt.Errorf("自动迁移失败: %w", err) } // 新表结构就位后删除遗留的明文列 if err := dropLegacyRefreshTokenColumn(db); err != nil { return fmt.Errorf("refresh token 旧列清理失败: %w", err) } DB = db seedDefaultBoards(db) log.Println("[model] PostgreSQL 数据库初始化完成") return nil } // PingDB 检测数据库连接 func PingDB() error { if DB == nil { return fmt.Errorf("数据库未初始化") } sqlDB, err := DB.DB() if err != nil { return err } return sqlDB.Ping() } // prepareRefreshTokenMigration 旧版 refresh_tokens 表把明文存在 token 列, // 新版改为 token_hash(SHA-256,NOT NULL+唯一索引)。在 AutoMigrate 之前: // 1. 新增可带默认值的 token_hash 列(避免对存量行加 NOT NULL 列失败) // 2. 用存量明文回填哈希 // 3. AutoMigrate 随后补唯一索引/其余新列 // // 旧行无法回填 TokenCipher(密钥在 service 层),仅影响该行下一次轮转的 // 并发重放宽限,属一次性边界;轮转后即完全进入新体系。 func prepareRefreshTokenMigration(db *gorm.DB) error { var tableCount int64 if err := db.Raw(`SELECT count(1) FROM information_schema.tables WHERE table_name = 'refresh_tokens'`). Scan(&tableCount).Error; err != nil { return err } if tableCount == 0 { return nil // 全新数据库,AutoMigrate 直接建新表 } var hasHashCol int64 if err := db.Raw(`SELECT count(1) FROM information_schema.columns WHERE table_name = 'refresh_tokens' AND column_name = 'token_hash'`). Scan(&hasHashCol).Error; err != nil { return err } if hasHashCol > 0 { return nil // 已是新结构 } var hasLegacyCol int64 if err := db.Raw(`SELECT count(1) FROM information_schema.columns WHERE table_name = 'refresh_tokens' AND column_name = 'token'`). Scan(&hasLegacyCol).Error; err != nil { return err } if hasLegacyCol == 0 { return nil } if err := db.Exec(`DELETE FROM refresh_tokens WHERE token IS NULL OR token = ''`).Error; err != nil { return err } if err := db.Exec(`ALTER TABLE refresh_tokens ADD COLUMN token_hash varchar(64) NOT NULL DEFAULT ''`).Error; err != nil { return err } type legacyRow struct { ID uint Token string } var rows []legacyRow if err := db.Raw(`SELECT id, token FROM refresh_tokens`).Scan(&rows).Error; err != nil { return err } for _, r := range rows { sum := sha256.Sum256([]byte(r.Token)) if err := db.Exec( `UPDATE refresh_tokens SET token_hash = ? WHERE id = ?`, hex.EncodeToString(sum[:]), r.ID, ).Error; err != nil { return err } } log.Printf("[model] refresh_tokens 已回填 %d 行 token_hash", len(rows)) return nil } // dropLegacyRefreshTokenColumn 新结构就位后删除明文 token 列(PostgreSQL // 会连带删除该列上的旧唯一索引) func dropLegacyRefreshTokenColumn(db *gorm.DB) error { var hasLegacyCol int64 if err := db.Raw(`SELECT count(1) FROM information_schema.columns WHERE table_name = 'refresh_tokens' AND column_name = 'token'`). Scan(&hasLegacyCol).Error; err != nil { return err } if hasLegacyCol == 0 { return nil } return db.Exec(`ALTER TABLE refresh_tokens DROP COLUMN token`).Error } // seedDefaultBoards 写入默认板块 func seedDefaultBoards(db *gorm.DB) { defaults := []Board{ {Name: "综合讨论", Description: "什么都可以聊", Icon: "message-circle", SortOrder: 1}, {Name: "技术分享", Description: "分享技术心得与问题", Icon: "code", SortOrder: 2}, {Name: "问答求助", Description: "提问与解答", Icon: "help-circle", SortOrder: 3}, {Name: "闲聊灌水", Description: "轻松闲聊", Icon: "coffee", SortOrder: 4}, } for _, b := range defaults { var count int64 db.Model(&Board{}).Where("name = ?", b.Name).Count(&count) if count == 0 { _ = db.Create(&b).Error } } }