package database import ( "fmt" "log" "os" "path/filepath" "runtime" "github.com/pressly/goose/v3" "go.uber.org/zap" "gorm.io/driver/mysql" "gorm.io/driver/sqlite" "gorm.io/gorm" "gorm.io/gorm/logger" "nex/backend/internal/config" ) func Init(cfg *config.DatabaseConfig, zapLogger *zap.Logger) (*gorm.DB, error) { db, err := initDB(cfg) if err != nil { return nil, fmt.Errorf("初始化数据库失败: %w", err) } if err := runMigrations(db, cfg.Driver); err != nil { return nil, fmt.Errorf("数据库迁移失败: %w", err) } configurePool(db, cfg) return db, nil } func Close(db *gorm.DB) { sqlDB, err := db.DB() if err != nil { return } sqlDB.Close() } func initDB(cfg *config.DatabaseConfig) (*gorm.DB, error) { gormConfig := &gorm.Config{ Logger: logger.Default.LogMode(logger.Info), } switch cfg.Driver { case "mysql": dsn := fmt.Sprintf("%s:%s@tcp(%s:%d)/%s?charset=utf8mb4&parseTime=true&loc=Local", cfg.User, cfg.Password, cfg.Host, cfg.Port, cfg.DBName) return gorm.Open(mysql.Open(dsn), gormConfig) default: dbDir := filepath.Dir(cfg.Path) if err := os.MkdirAll(dbDir, 0755); err != nil { return nil, fmt.Errorf("创建数据库目录失败: %w", err) } return gorm.Open(sqlite.Open(cfg.Path), gormConfig) } } func runMigrations(db *gorm.DB, driver string) error { sqlDB, err := db.DB() if err != nil { return err } migrationsDir := getMigrationsDir(driver) if _, err := os.Stat(migrationsDir); os.IsNotExist(err) { return fmt.Errorf("迁移目录不存在: %s", migrationsDir) } gooseDialect := "sqlite3" migrationsSubDir := "sqlite" if driver == "mysql" { gooseDialect = "mysql" migrationsSubDir = "mysql" } goose.SetDialect(gooseDialect) if err := goose.Up(sqlDB, migrationsDir); err != nil { return err } log.Printf("使用 %s 方言执行迁移,目录: %s", gooseDialect, migrationsSubDir) return nil } func configurePool(db *gorm.DB, cfg *config.DatabaseConfig) { if cfg.Driver == "sqlite" { if err := db.Exec("PRAGMA journal_mode=WAL").Error; err != nil { log.Printf("警告: 启用 WAL 模式失败: %v", err) } } sqlDB, err := db.DB() if err != nil { return } sqlDB.SetMaxIdleConns(cfg.MaxIdleConns) sqlDB.SetMaxOpenConns(cfg.MaxOpenConns) sqlDB.SetConnMaxLifetime(cfg.ConnMaxLifetime) log.Printf("数据库连接池配置: MaxIdle=%d, MaxOpen=%d, MaxLifetime=%v", cfg.MaxIdleConns, cfg.MaxOpenConns, cfg.ConnMaxLifetime) } func getMigrationsDir(driver string) string { _, filename, _, ok := runtime.Caller(0) if ok { subDir := "sqlite" if driver == "mysql" { subDir = "mysql" } dir := filepath.Join(filepath.Dir(filename), "..", "..", "migrations", subDir) if abs, err := filepath.Abs(dir); err == nil { return abs } } return "./migrations" } func BuildDSN(cfg *config.DatabaseConfig) string { return fmt.Sprintf("%s:%s@tcp(%s:%d)/%s?charset=utf8mb4&parseTime=true&loc=Local", cfg.User, cfg.Password, cfg.Host, cfg.Port, cfg.DBName) }