- users/user_groups 模型与 CRUD 接口,组成员关系手动维护(规避 GORM 零值主键问题) - 迁移 v2~v4:用户组、用户表、初始 admin 用户 - 初始密码随机生成,仅终端打印一次并写入 data/admin_password.txt
176 lines
4.6 KiB
Go
176 lines
4.6 KiB
Go
package database
|
|
|
|
import (
|
|
"context"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"testing"
|
|
|
|
"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(&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)
|
|
}
|
|
|
|
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 _, 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)
|
|
}
|
|
|
|
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("两次生成的密码不应相同")
|
|
}
|
|
}
|