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() }