diff --git a/internal/database/database.go b/internal/database/database.go new file mode 100644 index 0000000..5904979 --- /dev/null +++ b/internal/database/database.go @@ -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() +}