82 lines
1.7 KiB
Go
82 lines
1.7 KiB
Go
package common
|
|
|
|
import (
|
|
"fmt"
|
|
"log"
|
|
"time"
|
|
|
|
"gorm.io/driver/mysql"
|
|
"gorm.io/driver/postgres"
|
|
"gorm.io/driver/sqlite"
|
|
"gorm.io/gorm"
|
|
"gorm.io/gorm/logger"
|
|
)
|
|
|
|
var DB *gorm.DB
|
|
|
|
// InitDatabase 初始化数据库连接
|
|
func InitDatabase(config *DatabaseConfig) error {
|
|
var err error
|
|
var dialector gorm.Dialector
|
|
|
|
// 根据驱动类型选择相应的方言
|
|
switch config.Driver {
|
|
case "mysql":
|
|
dialector = mysql.Open(config.GetDSN())
|
|
case "postgres":
|
|
dialector = postgres.Open(config.GetDSN())
|
|
case "sqlite":
|
|
dialector = sqlite.Open(config.GetDSN())
|
|
default:
|
|
return fmt.Errorf("不支持的数据库驱动: %s", config.Driver)
|
|
}
|
|
|
|
// GORM 配置
|
|
gormConfig := &gorm.Config{
|
|
Logger: logger.Default.LogMode(logger.Info),
|
|
}
|
|
|
|
// 连接数据库
|
|
DB, err = gorm.Open(dialector, gormConfig)
|
|
if err != nil {
|
|
return fmt.Errorf("连接数据库失败: %v", err)
|
|
}
|
|
|
|
// 获取底层的 *sql.DB
|
|
sqlDB, err := DB.DB()
|
|
if err != nil {
|
|
return fmt.Errorf("获取数据库连接失败: %v", err)
|
|
}
|
|
|
|
// 设置连接池参数
|
|
sqlDB.SetMaxIdleConns(config.MaxIdleConns)
|
|
sqlDB.SetMaxOpenConns(config.MaxOpenConns)
|
|
sqlDB.SetConnMaxLifetime(time.Duration(config.ConnMaxLifetime) * time.Second)
|
|
|
|
// 测试连接
|
|
if err := sqlDB.Ping(); err != nil {
|
|
return fmt.Errorf("数据库连接测试失败: %v", err)
|
|
}
|
|
|
|
log.Println("数据库连接成功")
|
|
return nil
|
|
}
|
|
|
|
// GetDB 获取数据库实例
|
|
func GetDB() *gorm.DB {
|
|
return DB
|
|
}
|
|
|
|
// CloseDatabase 关闭数据库连接
|
|
func CloseDatabase() error {
|
|
sqlDB, err := DB.DB()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return sqlDB.Close()
|
|
}
|
|
|
|
// AutoMigrate 自动迁移数据表
|
|
func AutoMigrate(models ...interface{}) error {
|
|
return DB.AutoMigrate(models...)
|
|
} |