package database import ( "context" "os" "path/filepath" "strings" "testing" "time" "golang.org/x/crypto/bcrypt" "gorm.io/gorm" "rill/internal/config" "rill/internal/model" ) func testConfig(t *testing.T) *config.Config { t.Helper() return &config.Config{ Database: config.DatabaseConfig{ Driver: driverSQLite, ConnectTimeout: "5s", SQLite: config.SQLiteConfig{Path: filepath.Join(t.TempDir(), "test.db")}, }, } } func openTestDB(t *testing.T) *gorm.DB { t.Helper() db, err := Open(testConfig(t)) if err != nil { t.Fatalf("Open 失败: %v", err) } t.Cleanup(func() { if err := Close(db); err != nil { t.Errorf("Close 失败: %v", err) } }) return db } func TestOpenAndPing(t *testing.T) { db := openTestDB(t) if err := Ping(context.Background(), db); err != nil { t.Fatalf("Ping 失败: %v", err) } } func TestOpenUnsupportedDriver(t *testing.T) { cfg := testConfig(t) cfg.Database.Driver = "postgres" if _, err := Open(cfg); err == nil { t.Fatal("期望不支持的驱动返回错误") } } func TestMigrateIdempotentAndCRUD(t *testing.T) { db := openTestDB(t) ctx := context.Background() for i := 1; i <= 2; i++ { if err := Migrate(ctx, db); err != nil { t.Fatalf("第 %d 次 Migrate 失败: %v", i, err) } } if !db.Migrator().HasTable(&model.Note{}) { t.Error("notes 表未创建") } if !db.Migrator().HasTable(&model.User{}) { t.Error("users 表未创建") } if !db.Migrator().HasTable(&model.UserGroup{}) { t.Error("user_groups 表未创建") } if !db.Migrator().HasTable(&model.SiteSetting{}) { t.Error("site_settings 表未创建") } if !db.Migrator().HasTable(&schemaMigration{}) { t.Error("schema_migrations 表未创建") } var groups []model.UserGroup if err := db.WithContext(ctx).Order("id ASC").Find(&groups).Error; err != nil { t.Fatalf("查询内置组失败: %v", err) } if len(groups) != 2 || groups[0].ID != model.GroupIDAdmin || groups[1].ID != model.GroupIDUser { t.Errorf("内置组异常: %+v", groups) } if groups[0].Description != "Administrator group" || groups[1].Description != "Regular user group" { t.Errorf("内置组描述应为英文: %+v", groups) } var admin model.User if err := db.WithContext(ctx).Where("username = ?", adminUsername).First(&admin).Error; err != nil { t.Fatalf("初始管理员未创建: %v", err) } if admin.Status != 1 { t.Errorf("初始管理员状态 = %d, 期望 1", admin.Status) } if admin.Nickname != "Administrator" { t.Errorf("初始管理员昵称应为英文: %q", admin.Nickname) } for _, column := range []string{"gender", "birthday"} { if !db.Migrator().HasColumn(&model.User{}, column) { t.Errorf("users 表缺少列 %s", column) } } if admin.Gender != "" || !admin.Birthday.IsZero() { t.Errorf("初始管理员资料字段应为空: gender=%q birthday=%v", admin.Gender, admin.Birthday.Time) } if _, err := bcrypt.Cost([]byte(admin.PasswordHash)); err != nil { t.Errorf("初始管理员密码哈希无效: %v", err) } var membership int64 if err := db.WithContext(ctx).Model(&model.UserGroupMember{}). Where("user_id = ? AND user_group_id = ?", admin.ID, model.GroupIDAdmin). Count(&membership).Error; err != nil { t.Fatalf("查询初始管理员组关系失败: %v", err) } if membership != 1 { t.Errorf("初始管理员未加入 admin 组: count=%d", membership) } var setting model.SiteSetting if err := db.WithContext(ctx).First(&setting, model.SiteSettingID).Error; err != nil { t.Fatalf("默认站点信息未创建: %v", err) } if setting.SiteName != "Rill" || setting.Logo != "" || setting.Footer != "" { t.Errorf("默认站点信息异常: %+v", setting) } passwordFile := adminPasswordPath(db) raw, err := os.ReadFile(passwordFile) if err != nil { t.Fatalf("读取初始密码文件失败: %v", err) } if !strings.Contains(string(raw), "删除") { t.Error("密码文件缺少删除提醒") } password := "" for _, line := range strings.Split(string(raw), "\n") { if value, ok := strings.CutPrefix(line, "密码: "); ok { password = strings.TrimSpace(value) break } } if password == "" { t.Fatalf("密码文件缺少密码行: %s", raw) } if err := bcrypt.CompareHashAndPassword([]byte(admin.PasswordHash), []byte(password)); err != nil { t.Errorf("密码文件中的密码与数据库哈希不匹配: %v", err) } var count int64 if err := db.Model(&schemaMigration{}).Count(&count).Error; err != nil { t.Fatalf("统计迁移记录失败: %v", err) } if count != int64(len(migrations)) { t.Errorf("迁移记录数 = %d, 期望 %d", count, len(migrations)) } note := model.Note{Title: "标题", Content: "内容"} if err := db.WithContext(ctx).Create(¬e).Error; err != nil { t.Fatalf("创建记录失败: %v", err) } var got model.Note if err := db.WithContext(ctx).First(&got, note.ID).Error; err != nil { t.Fatalf("查询记录失败: %v", err) } if got.Title != note.Title || got.Content != note.Content { t.Errorf("查询结果 = %+v, 期望 Title=%q Content=%q", got, note.Title, note.Content) } } func TestGeneratePassword(t *testing.T) { password, err := generatePassword(adminPasswordLen) if err != nil { t.Fatalf("生成密码失败: %v", err) } if len(password) != adminPasswordLen { t.Errorf("密码长度 = %d, 期望 %d", len(password), adminPasswordLen) } for _, r := range password { if !strings.ContainsRune(adminPasswordCharset, r) { t.Errorf("密码包含非法字符 %q", r) } } other, err := generatePassword(adminPasswordLen) if err != nil { t.Fatalf("生成密码失败: %v", err) } if password == other { t.Error("两次生成的密码不应相同") } } func TestAddProfileFieldsMigration(t *testing.T) { db := openTestDB(t) ctx := context.Background() if err := Migrate(ctx, db); err != nil { t.Fatalf("执行迁移失败: %v", err) } migrator := db.Migrator() for _, column := range []string{"gender", "birthday"} { if err := migrator.DropColumn(&model.User{}, column); err != nil { t.Fatalf("删除列 %s 失败: %v", column, err) } } // 用原始 SQL 模拟旧版本数据(当时表里还没有 gender/birthday 列)。 now := time.Now() if err := db.WithContext(ctx).Exec( "INSERT INTO users (username, email, password_hash, nickname, avatar, status, created_at, updated_at) VALUES (?, ?, ?, '', '', 1, ?, ?)", "legacy", "legacy@example.com", "x", now, now, ).Error; err != nil { t.Fatalf("创建存量用户失败: %v", err) } var addProfile Migration for _, m := range migrations { if m.Version == 6 { addProfile = m } } if addProfile.Up == nil { t.Fatal("未找到 v6 迁移") } if err := addProfile.Up(db.WithContext(ctx)); err != nil { t.Fatalf("执行 v6 迁移失败: %v", err) } for _, column := range []string{"gender", "birthday"} { if !migrator.HasColumn(&model.User{}, column) { t.Errorf("v6 未补齐列 %s", column) } } var got model.User if err := db.WithContext(ctx).Where("username = ?", "legacy").First(&got).Error; err != nil { t.Fatalf("查询存量用户失败: %v", err) } if got.Gender != "" || !got.Birthday.IsZero() { t.Errorf("存量用户资料字段应默认空: gender=%q birthday=%v", got.Gender, got.Birthday.Time) } } func TestTranslateBuiltinDataMigration(t *testing.T) { db := openTestDB(t) ctx := context.Background() if err := Migrate(ctx, db); err != nil { t.Fatalf("执行迁移失败: %v", err) } var translate Migration for _, m := range migrations { if m.Version == 5 { translate = m } } if translate.Up == nil { t.Fatal("未找到 v5 翻译迁移") } // 还原为历史中文数据后重新执行 v5,应更新内置数据。 if err := db.WithContext(ctx).Model(&model.UserGroup{}). Where("id = ?", model.GroupIDAdmin).Update("description", "管理员组").Error; err != nil { t.Fatalf("还原 admin 组描述失败: %v", err) } if err := db.WithContext(ctx).Model(&model.UserGroup{}). Where("id = ?", model.GroupIDUser).Update("description", "普通用户组").Error; err != nil { t.Fatalf("还原 user 组描述失败: %v", err) } if err := db.WithContext(ctx).Model(&model.User{}). Where("username = ?", "admin").Update("nickname", "管理员").Error; err != nil { t.Fatalf("还原 admin 昵称失败: %v", err) } if err := translate.Up(db.WithContext(ctx)); err != nil { t.Fatalf("执行 v5 迁移失败: %v", err) } var admin model.User if err := db.WithContext(ctx).Where("username = ?", "admin").First(&admin).Error; err != nil { t.Fatalf("查询管理员失败: %v", err) } if admin.Nickname != "Administrator" { t.Errorf("管理员昵称 = %q, 期望 Administrator", admin.Nickname) } var adminGroup, userGroup model.UserGroup if err := db.WithContext(ctx).First(&adminGroup, model.GroupIDAdmin).Error; err != nil { t.Fatalf("查询 admin 组失败: %v", err) } if err := db.WithContext(ctx).First(&userGroup, model.GroupIDUser).Error; err != nil { t.Fatalf("查询 user 组失败: %v", err) } if adminGroup.Description != "Administrator group" || userGroup.Description != "Regular user group" { t.Errorf("内置组描述 = %q / %q", adminGroup.Description, userGroup.Description) } // 手工修改过的值不应被覆盖。 if err := db.WithContext(ctx).Model(&model.User{}). Where("username = ?", "admin").Update("nickname", "Boss").Error; err != nil { t.Fatalf("自定义昵称失败: %v", err) } if err := translate.Up(db.WithContext(ctx)); err != nil { t.Fatalf("再次执行 v5 迁移失败: %v", err) } if err := db.WithContext(ctx).Where("username = ?", "admin").First(&admin).Error; err != nil { t.Fatalf("查询管理员失败: %v", err) } if admin.Nickname != "Boss" { t.Errorf("自定义昵称被覆盖: %q", admin.Nickname) } }