package database import ( "context" "fmt" "os" "path/filepath" "github.com/pressly/goose/v3" "go.uber.org/zap" "gorm.io/driver/mysql" "gorm.io/driver/sqlite" "gorm.io/gorm" "nex/backend/internal/config" "nex/backend/migrations" pkglogger "nex/backend/pkg/logger" ) func Init(cfg *config.DatabaseConfig, zapLogger *zap.Logger) (*gorm.DB, error) { moduleLogger := pkglogger.WithModule(zapLogger, "database") db, err := initDB(cfg, moduleLogger) if err != nil { return nil, fmt.Errorf("初始化数据库失败: %w", err) } if err := runMigrations(db, cfg.Driver, moduleLogger); err != nil { return nil, fmt.Errorf("数据库迁移失败: %w", err) } configurePool(db, cfg, moduleLogger) return db, nil } func Close(db *gorm.DB) { sqlDB, err := db.DB() if err != nil { return } sqlDB.Close() } func initDB(cfg *config.DatabaseConfig, zapLogger *zap.Logger) (*gorm.DB, error) { gormLogger := pkglogger.NewGormLogger(zapLogger) gormConfig := &gorm.Config{ Logger: gormLogger, } 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) if zapLogger != nil { zapLogger.Info("连接 MySQL 数据库", zap.String("host", cfg.Host), zap.Int("port", cfg.Port), zap.String("database", cfg.DBName)) } return gorm.Open(mysql.Open(dsn), gormConfig) default: dbDir := filepath.Dir(cfg.Path) if err := os.MkdirAll(dbDir, 0o755); err != nil { return nil, fmt.Errorf("创建数据库目录失败: %w", err) } if zapLogger != nil { zapLogger.Info("连接 SQLite 数据库", zap.String("path", cfg.Path)) } return gorm.Open(sqlite.Open(cfg.Path), gormConfig) } } func runMigrations(db *gorm.DB, driver string, zapLogger *zap.Logger) error { sqlDB, err := db.DB() if err != nil { return err } dialect, fsys, err := migrations.ForDriver(driver) if err != nil { return err } if zapLogger != nil { zapLogger.Info("执行数据库迁移", zap.String("dialect", string(dialect)), zap.String("driver", driver)) } provider, err := goose.NewProvider(dialect, sqlDB, fsys) if err != nil { return fmt.Errorf("创建迁移提供者失败: %w", err) } if _, err := provider.Up(context.Background()); err != nil { return fmt.Errorf("执行迁移失败: %w", err) } return nil } func configurePool(db *gorm.DB, cfg *config.DatabaseConfig, zapLogger *zap.Logger) { if cfg.Driver == "sqlite" { if err := db.Exec("PRAGMA journal_mode=WAL").Error; err != nil { if zapLogger != nil { zapLogger.Warn("启用 WAL 模式失败", zap.Error(err)) } } } sqlDB, err := db.DB() if err != nil { return } sqlDB.SetMaxIdleConns(cfg.MaxIdleConns) sqlDB.SetMaxOpenConns(cfg.MaxOpenConns) sqlDB.SetConnMaxLifetime(cfg.ConnMaxLifetime) if zapLogger != nil { zapLogger.Info("数据库连接池配置", zap.Int("max_idle_conns", cfg.MaxIdleConns), zap.Int("max_open_conns", cfg.MaxOpenConns), zap.Duration("conn_max_lifetime", cfg.ConnMaxLifetime)) } } 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) }