Files

128 lines
2.5 KiB
Go
Raw Permalink Normal View History

package db
import (
"log"
"time"
"d:\项目\dobaochet\backend\config"
"d:\项目\dobaochet\backend\models"
"d:\项目\dobaochet\backend\utils"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// DB 全局数据库连接
var DB *gorm.DB
// InitDB 初始化数据库连接
func InitDB(cfg *config.Config) (*gorm.DB, error) {
// 配置GORM日志
newLogger := logger.New(
log.New(log.Writer(), "\r\n", log.LstdFlags),
logger.Config{
SlowThreshold: time.Second, // 慢SQL阈值
LogLevel: logger.Info, // 日志级别
Colorful: true, // 彩色日志
},
)
// 连接数据库
var err error
switch {
case len(cfg.Database.DSN) > 6 && cfg.Database.DSN[:6] == "sqlite":
// SQLite连接
DB, err = gorm.Open(sqlite.Open(cfg.Database.DSN[9:]), &gorm.Config{
Logger: newLogger,
})
default:
// 默认使用SQLite
DB, err = gorm.Open(sqlite.Open("./chat.db"), &gorm.Config{
Logger: newLogger,
})
}
if err != nil {
return nil, err
}
// 设置连接池
sqlDB, err := DB.DB()
if err != nil {
return nil, err
}
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns)
sqlDB.SetMaxOpenConns(cfg.Database.MaxOpenConns)
sqlDB.SetConnMaxLifetime(time.Hour)
// 自动迁移模型
err = migrateModels()
if err != nil {
return nil, err
}
// 初始化数据
err = seedData(cfg)
if err != nil {
return nil, err
}
return DB, nil
}
// migrateModels 自动迁移模型
func migrateModels() error {
return DB.AutoMigrate(
&models.User{},
&models.Quota{},
&models.Conversation{},
&models.Message{},
)
}
// seedData 初始化数据
func seedData(cfg *config.Config) error {
// 检查是否已有管理员用户
var count int64
DB.Model(&models.User{}).Where("is_admin = ?", true).Count(&count)
if count == 0 {
// 创建管理员用户
password, err := utils.HashPassword("admin123")
if err != nil {
return err
}
adminUser := models.User{
Username: "admin",
Email: "admin@example.com",
Password: password,
IsAdmin: true,
IsActive: true,
LastLoginAt: time.Now(),
Quota: models.Quota{
TotalTokens: cfg.Quota.DefaultTotalTokens * 10, // 管理员配额10倍
UsedTokens: 0,
ResetAt: time.Now().Add(cfg.Quota.ResetInterval),
TokenLimit: cfg.Quota.DefaultTokenLimit * 10,
},
}
result := DB.Create(&adminUser)
if result.Error != nil {
return result.Error
}
log.Println("已创建管理员用户: admin / admin123")
}
return nil
}
// GetDB 获取数据库连接
func GetDB() *gorm.DB {
return DB
}