170 lines
3.9 KiB
Go
170 lines
3.9 KiB
Go
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...))
|
|
}
|