package database import ( "context" "fmt" "os" "path/filepath" "strings" "time" "skeleton/config" "go.uber.org/zap" "gorm.io/driver/mysql" "gorm.io/driver/postgres" "gorm.io/gorm" "gorm.io/gorm/logger" ) var DB *gorm.DB // Init 根据配置的 driver 初始化数据库连接。 func Init(cfg *config.DatabaseConfig, log *zap.Logger) error { if err := prepareSQLiteDirectory(cfg); err != nil { return err } dialector, err := Dialector(cfg) if err != nil { return err } gormLogger := logger.New( &GormZapWriter{Logger: log}, logger.Config{ SlowThreshold: time.Second, LogLevel: logger.Info, Colorful: false, }, ) db, err := gorm.Open(dialector, &gorm.Config{Logger: gormLogger}) if err != nil { return fmt.Errorf("连接 %s 数据库失败: %w", cfg.Driver, err) } sqlDB, err := db.DB() if err != nil { return fmt.Errorf("获取底层数据库连接失败: %w", err) } sqlDB.SetMaxIdleConns(cfg.MaxIdleConns) sqlDB.SetMaxOpenConns(cfg.MaxOpenConns) sqlDB.SetConnMaxLifetime(time.Duration(cfg.ConnMaxLifetime) * time.Minute) ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() if err := sqlDB.PingContext(ctx); err != nil { _ = sqlDB.Close() return fmt.Errorf("%s 数据库连接测试失败: %w", cfg.Driver, err) } DB = db log.Info("数据库连接初始化成功", zap.String("driver", cfg.Driver), zap.String("database", databaseName(cfg)), zap.Int("max_idle_conns", cfg.MaxIdleConns), zap.Int("max_open_conns", cfg.MaxOpenConns), ) return nil } func prepareSQLiteDirectory(cfg *config.DatabaseConfig) error { driver := strings.ToLower(strings.TrimSpace(cfg.Driver)) if driver != "sqlite" && driver != "sqlite3" { return nil } path := BuildDSN(cfg) if path == "" || path == ":memory:" || strings.HasPrefix(path, "file:") { return nil } dir := filepath.Dir(path) if dir == "." { return nil } if err := os.MkdirAll(dir, 0755); err != nil { return fmt.Errorf("创建 SQLite 数据目录失败: %w", err) } return nil } // Dialector 创建对应数据库的 GORM 方言实例,便于应用和测试复用。 func Dialector(cfg *config.DatabaseConfig) (gorm.Dialector, error) { driver := strings.ToLower(strings.TrimSpace(cfg.Driver)) dsn := BuildDSN(cfg) switch driver { case "postgres", "postgresql": return postgres.Open(dsn), nil case "mysql": return mysql.Open(dsn), nil case "sqlite", "sqlite3": return sqliteDialector(dsn) default: return nil, fmt.Errorf("不支持的数据库驱动 %q,可选值: postgres, mysql, sqlite", cfg.Driver) } } // BuildDSN 返回显式 DSN,或根据结构化配置构建对应方言的连接串。 func BuildDSN(cfg *config.DatabaseConfig) string { if cfg.DSN != "" { return cfg.DSN } switch strings.ToLower(strings.TrimSpace(cfg.Driver)) { case "mysql": return fmt.Sprintf("%s:%s@tcp(%s:%d)/%s?charset=utf8mb4&parseTime=True&loc=Local", cfg.Username, cfg.Password, cfg.Host, cfg.Port, cfg.DBName) case "sqlite", "sqlite3": return cfg.SQLitePath default: return fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s sslmode=%s", cfg.Host, cfg.Port, cfg.Username, cfg.Password, cfg.DBName, cfg.SSLMode) } } func databaseName(cfg *config.DatabaseConfig) string { if strings.HasPrefix(strings.ToLower(cfg.Driver), "sqlite") { return BuildDSN(cfg) } return cfg.DBName } func Close() error { if DB == nil { return nil } sqlDB, err := DB.DB() if err != nil { return err } return sqlDB.Close() } func GetDB() *gorm.DB { return DB } // Ping 检查关系数据库连接是否可用。 func Ping(ctx context.Context) error { if DB == nil { return fmt.Errorf("数据库尚未初始化") } sqlDB, err := DB.DB() if err != nil { return err } return sqlDB.PingContext(ctx) } type GormZapWriter struct { Logger *zap.Logger } func (g *GormZapWriter) Printf(format string, args ...interface{}) { g.Logger.Info(fmt.Sprintf(format, args...)) }