Add internal/database/database.go
This commit is contained in:
157
internal/database/database.go
Normal file
157
internal/database/database.go
Normal file
@@ -0,0 +1,157 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"log"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"gotest/internal/config"
|
||||||
|
"gotest/internal/model"
|
||||||
|
|
||||||
|
"github.com/glebarez/sqlite"
|
||||||
|
"golang.org/x/crypto/bcrypt"
|
||||||
|
"gorm.io/driver/mysql"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DB 封装 gorm.DB,便于扩展
|
||||||
|
type DB struct {
|
||||||
|
*gorm.DB
|
||||||
|
}
|
||||||
|
|
||||||
|
// Init 初始化数据库连接,根据驱动选择 SQLite 或 MySQL
|
||||||
|
func Init(cfg *config.DatabaseConfig) (*DB, error) {
|
||||||
|
var db *gorm.DB
|
||||||
|
var err error
|
||||||
|
|
||||||
|
switch cfg.Driver {
|
||||||
|
case "mysql":
|
||||||
|
db, err = gorm.Open(mysql.Open(cfg.DSN), &gorm.Config{})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("连接 MySQL 失败: %w", err)
|
||||||
|
}
|
||||||
|
// 配置连接池
|
||||||
|
sqlDB, derr := db.DB()
|
||||||
|
if derr == nil {
|
||||||
|
sqlDB.SetMaxIdleConns(cfg.MySQL.MaxIdleConns)
|
||||||
|
sqlDB.SetMaxOpenConns(cfg.MySQL.MaxOpenConns)
|
||||||
|
}
|
||||||
|
log.Printf("已连接 MySQL: %s:%d/%s", cfg.MySQL.Host, cfg.MySQL.Port, cfg.MySQL.Database)
|
||||||
|
case "sqlite":
|
||||||
|
fallthrough
|
||||||
|
default:
|
||||||
|
db, err = gorm.Open(sqlite.Open(cfg.DSN), &gorm.Config{})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("连接 SQLite 失败: %w", err)
|
||||||
|
}
|
||||||
|
// SQLite 连接池优化:WAL 模式 + 忙等待 + 连接复用
|
||||||
|
sqlDB, derr := db.DB()
|
||||||
|
if derr == nil {
|
||||||
|
sqlDB.SetMaxIdleConns(10)
|
||||||
|
sqlDB.SetMaxOpenConns(100)
|
||||||
|
sqlDB.SetConnMaxLifetime(time.Hour)
|
||||||
|
// 启用 WAL 模式提升并发读写性能
|
||||||
|
db.Exec("PRAGMA journal_mode=WAL")
|
||||||
|
db.Exec("PRAGMA busy_timeout=5000")
|
||||||
|
db.Exec("PRAGMA foreign_keys=ON")
|
||||||
|
}
|
||||||
|
log.Printf("已连接 SQLite: %s (WAL 模式)", cfg.DSN)
|
||||||
|
}
|
||||||
|
|
||||||
|
// auto_migrate 自动建表/更新表结构
|
||||||
|
if cfg.AutoMigrate {
|
||||||
|
if err := db.AutoMigrate(
|
||||||
|
&model.User{}, &model.Config{},
|
||||||
|
&model.Company{}, &model.EnterpriseNews{}, &model.EnterpriseService{}, &model.TeamMember{},
|
||||||
|
&model.ProductCategory{}, &model.Product{},
|
||||||
|
&model.RestaurantCategory{}, &model.Restaurant{}, &model.MenuCategory{}, &model.MenuItem{},
|
||||||
|
&model.GameCategory{}, &model.Game{},
|
||||||
|
&model.AppCategory{}, &model.App{},
|
||||||
|
&model.MusicCategory{}, &model.Song{},
|
||||||
|
); err != nil {
|
||||||
|
return nil, fmt.Errorf("自动建表失败: %w", err)
|
||||||
|
}
|
||||||
|
log.Println("数据库表结构迁移完成")
|
||||||
|
}
|
||||||
|
|
||||||
|
// init_data 数据初始化(仅在表为空时执行,不会覆盖已有数据)
|
||||||
|
if cfg.InitData {
|
||||||
|
if err := seedData(db); err != nil {
|
||||||
|
return nil, fmt.Errorf("数据初始化失败: %w", err)
|
||||||
|
}
|
||||||
|
if err := seedEnterpriseData(db); err != nil {
|
||||||
|
return nil, fmt.Errorf("企业数据初始化失败: %w", err)
|
||||||
|
}
|
||||||
|
if err := seedShopData(db); err != nil {
|
||||||
|
return nil, fmt.Errorf("电商数据初始化失败: %w", err)
|
||||||
|
}
|
||||||
|
if err := seedFoodData(db); err != nil {
|
||||||
|
return nil, fmt.Errorf("外卖数据初始化失败: %w", err)
|
||||||
|
}
|
||||||
|
if err := seedGameData(db); err != nil {
|
||||||
|
return nil, fmt.Errorf("游戏数据初始化失败: %w", err)
|
||||||
|
}
|
||||||
|
if err := seedAppData(db); err != nil {
|
||||||
|
return nil, fmt.Errorf("应用数据初始化失败: %w", err)
|
||||||
|
}
|
||||||
|
if err := seedMusicData(db); err != nil {
|
||||||
|
return nil, fmt.Errorf("音乐数据初始化失败: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &DB{db}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// seedData 初始化基础数据,仅在用户表为空时执行
|
||||||
|
func seedData(db *gorm.DB) error {
|
||||||
|
// 检查是否已有用户,避免重复初始化
|
||||||
|
var count int64
|
||||||
|
if err := db.Model(&model.User{}).Count(&count).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if count > 0 {
|
||||||
|
log.Println("检测到已有用户数据,跳过初始化")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 创建默认管理员账号 admin / admin123
|
||||||
|
hashedPassword, err := bcrypt.GenerateFromPassword([]byte("admin123"), bcrypt.DefaultCost)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("生成管理员密码哈希失败: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
admin := model.User{
|
||||||
|
Username: "admin",
|
||||||
|
Password: string(hashedPassword),
|
||||||
|
Role: "admin",
|
||||||
|
Status: 1,
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := db.Create(&admin).Error; err != nil {
|
||||||
|
return fmt.Errorf("创建默认管理员失败: %w", err)
|
||||||
|
}
|
||||||
|
log.Println("默认管理员账号已创建: admin / admin123(请及时修改密码)")
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// HealthCheck 检查数据库连接是否正常
|
||||||
|
func (d *DB) HealthCheck() error {
|
||||||
|
sqlDB, err := d.DB.DB()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("获取底层 DB 失败: %w", err)
|
||||||
|
}
|
||||||
|
if err := sqlDB.Ping(); err != nil {
|
||||||
|
return fmt.Errorf("数据库 Ping 失败: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close 关闭数据库连接
|
||||||
|
func (d *DB) Close() error {
|
||||||
|
sqlDB, err := d.DB.DB()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return sqlDB.Close()
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user