Files
2026-03-11 16:09:17 -05:00

162 lines
4.6 KiB
Go
Executable File

package config
import (
"context"
"fmt"
"time"
"github.com/oarkflow/log"
"gorm.io/gorm/logger"
"gorm.io/driver/mysql"
"gorm.io/driver/postgres"
"gorm.io/driver/sqlserver" // SQL Server driver
"gorm.io/gorm"
"gorm.io/plugin/dbresolver"
)
type DatabaseDriver struct {
Driver string `yaml:"driver" env:"DB_DRIVER"`
Host string `yaml:"host" env:"DB_HOST"`
Username string `yaml:"username" env:"DB_USER"`
Password string `yaml:"password" env:"DB_PASS"`
DBName string `yaml:"db_name" env:"DB_NAME"`
Port int `yaml:"port" env:"DB_PORT"`
Connections int `yaml:"connections" env:"DB_CONNECTIONS"`
}
type DatabaseConfig struct {
*gorm.DB
Drivers map[string]DatabaseDriver `yaml:"drivers"`
Default DatabaseDriver `yaml:"default"`
}
func (d *DatabaseConfig) Setup() error {
var err error
connectionString := ""
if d.DB != nil {
return nil
}
gormLogger := New(&log.DefaultLogger, logger.Config{
LogLevel: 0,
}, false)
newLogger := gormLogger.LogMode(logger.Info)
switch d.Default.Driver {
case "postgres":
connectionString = fmt.Sprintf("host=%s port=%d user=%s dbname=%s password=%s", d.Default.Host, d.Default.Port, d.Default.Username, d.Default.DBName, d.Default.Password)
d.DB, err = gorm.Open(postgres.Open(connectionString), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
Logger: newLogger,
})
case "mysql":
connectionString = fmt.Sprintf("%s:%s@tcp(%s:%d)/%s?charset=utf8&parseTime=True&loc=Local", d.Default.Username, d.Default.Password, d.Default.Host, d.Default.Port, d.Default.DBName)
d.DB, err = gorm.Open(mysql.Open(connectionString), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
Logger: newLogger,
})
case "sqlserver":
connectionString = fmt.Sprintf(
"sqlserver://%s:%s@%s:%d?database=%s&charset=utf8mb4",
d.Default.Username,
d.Default.Password,
d.Default.Host,
d.Default.Port,
d.Default.DBName,
)
d.DB, err = gorm.Open(sqlserver.Open(connectionString), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
Logger: newLogger,
})
default:
return fmt.Errorf("unsupported database driver: %s", d.Default.Driver)
}
if err != nil {
fmt.Println(d.Default)
panic(err)
}
d.DB.Use(
dbresolver.Register(dbresolver.Config{}).
SetConnMaxLifetime(24 * time.Hour).
SetMaxIdleConns(100).
SetMaxOpenConns(100),
)
return nil
}
func New(logger *log.Logger, config logger.Config, slient bool) logger.Interface {
return &gormLogger{
Log: logger,
Config: config,
Slient: slient,
}
}
type gormLogger struct {
Log *log.Logger
Config logger.Config
Slient bool
}
func (l *gormLogger) LogMode(level logger.LogLevel) logger.Interface {
var newLogger = gormLogger{Log: l.Log}
switch level {
case logger.Silent:
newLogger.Slient = true
case logger.Error:
newLogger.Log.SetLevel(log.ErrorLevel)
case logger.Warn:
newLogger.Log.SetLevel(log.WarnLevel)
case logger.Info:
newLogger.Log.SetLevel(log.InfoLevel)
}
return &newLogger
}
func (l *gormLogger) Info(ctx context.Context, format string, args ...interface{}) {
l.Log.Info().Msgf(format, args...)
}
func (l *gormLogger) Warn(ctx context.Context, format string, args ...interface{}) {
l.Log.Warn().Msgf(format, args...)
}
func (l *gormLogger) Error(ctx context.Context, format string, args ...interface{}) {
l.Log.Error().Msgf(format, args...)
}
func (l *gormLogger) Trace(ctx context.Context, begin time.Time, fc func() (string, int64), err error) {
if l.Slient {
return
}
elapsed := time.Since(begin)
switch {
case err != nil && l.Log.Level >= log.ErrorLevel:
sql, rows := fc()
if rows == -1 {
l.Log.Error().Caller(1).Err(err).Dur("elapsed", elapsed).Str("sql", sql).Msg("")
} else {
l.Log.Error().Caller(1).Err(err).Dur("elapsed", elapsed).Str("sql", sql).Int64("rows", rows).Msg("")
}
case elapsed > l.Config.SlowThreshold && l.Config.SlowThreshold != 0 && l.Log.Level >= log.WarnLevel:
sql, rows := fc()
if rows == -1 {
l.Log.Warn().Caller(1).Err(err).Dur("elapsed", elapsed).Str("sql", sql).Msgf("SLOW SQL >= %v", l.Config.SlowThreshold)
} else {
l.Log.Warn().Caller(1).Err(err).Dur("elapsed", elapsed).Str("sql", sql).Int64("rows", rows).Msgf("SLOW SQL >= %v", l.Config.SlowThreshold)
}
case l.Log.Level == log.InfoLevel:
sql, rows := fc()
if rows == -1 {
l.Log.Info().Caller(1).Err(err).Dur("elapsed", elapsed).Str("sql", sql).Msg("")
} else {
l.Log.Info().Caller(1).Err(err).Dur("elapsed", elapsed).Str("sql", sql).Int64("rows", rows).Msg("")
}
}
}