feat: 完善后端脚手架基础能力
This commit is contained in:
@@ -0,0 +1,169 @@
|
||||
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...))
|
||||
}
|
||||
Reference in New Issue
Block a user