201 lines
4.5 KiB
Go
201 lines
4.5 KiB
Go
package config
|
|
|
|
import (
|
|
"database/sql"
|
|
"fmt"
|
|
"log"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/fsnotify/fsnotify"
|
|
"github.com/spf13/viper"
|
|
|
|
_ "github.com/lib/pq" // PostgreSQL 驱动,根据实际数据库替换
|
|
)
|
|
|
|
// DBConfig 数据库配置结构体
|
|
type DBConfig struct {
|
|
Host string
|
|
Port int
|
|
Username string
|
|
Password string
|
|
Dbname string
|
|
MaxOpenConns int // 最大打开连接数
|
|
MaxIdleConns int // 最大空闲连接数
|
|
ConnMaxLifetime time.Duration // 连接最大生命周期
|
|
}
|
|
|
|
// 默认连接池参数
|
|
const (
|
|
defaultMaxOpenConns = 10
|
|
defaultMaxIdleConns = 5
|
|
defaultConnMaxLifetime = time.Hour
|
|
)
|
|
|
|
var (
|
|
dbConfigs map[string]DBConfig // 配置副本
|
|
dataSources = make(map[string]*sql.DB) // 连接池
|
|
configsDSN = make(map[string]string) // key -> dsn
|
|
mu sync.RWMutex
|
|
v *viper.Viper
|
|
)
|
|
|
|
// InitDBConfig 初始化并加载 db.yaml 配置,同时启动监听
|
|
func InitDBConfig(configPath string) error {
|
|
v = viper.New()
|
|
v.SetConfigFile(configPath)
|
|
v.SetConfigType("yaml")
|
|
|
|
if err := v.ReadInConfig(); err != nil {
|
|
return fmt.Errorf("读取数据库配置失败: %w", err)
|
|
}
|
|
|
|
if err := unmarshalConfigs(); err != nil {
|
|
return err
|
|
}
|
|
|
|
// 首次加载后,初始化数据源
|
|
if err := ReloadDataSources(GetDBConfigs()); err != nil {
|
|
log.Printf("首次加载数据源失败: %v\n", err)
|
|
return err
|
|
}
|
|
|
|
// 监听配置文件变化
|
|
v.WatchConfig()
|
|
v.OnConfigChange(func(e fsnotify.Event) {
|
|
log.Printf("数据库配置文件发生变化: %s\n", e.Name)
|
|
if err := unmarshalConfigs(); err != nil {
|
|
log.Printf("重新加载数据库配置失败: %v\n", err)
|
|
} else {
|
|
if err := ReloadDataSources(GetDBConfigs()); err != nil {
|
|
log.Printf("更新数据源失败: %v\n", err)
|
|
} else {
|
|
log.Println("数据库连接池已更新")
|
|
}
|
|
}
|
|
})
|
|
|
|
return nil
|
|
}
|
|
|
|
// unmarshalConfigs 解析配置到全局变量,线程安全
|
|
func unmarshalConfigs() error {
|
|
mu.Lock()
|
|
defer mu.Unlock()
|
|
|
|
temp := make(map[string]DBConfig)
|
|
if err := v.Unmarshal(&temp); err != nil {
|
|
return fmt.Errorf("解析数据库配置失败: %w", err)
|
|
}
|
|
dbConfigs = temp
|
|
return nil
|
|
}
|
|
|
|
// GetDBConfigs 并发安全地返回当前数据库配置副本
|
|
func GetDBConfigs() map[string]DBConfig {
|
|
mu.RLock()
|
|
defer mu.RUnlock()
|
|
|
|
res := make(map[string]DBConfig, len(dbConfigs))
|
|
for k, v := range dbConfigs {
|
|
res[k] = v
|
|
}
|
|
return res
|
|
}
|
|
|
|
// buildDSN 构建 PostgreSQL DSN
|
|
func buildDSN(cfg DBConfig) string {
|
|
return fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s sslmode=disable",
|
|
cfg.Host, cfg.Port, cfg.Username, cfg.Password, cfg.Dbname)
|
|
}
|
|
|
|
// ReloadDataSources 支持增量更新数据源,并验证数据库连接
|
|
func ReloadDataSources(newConfigs map[string]DBConfig) error {
|
|
mu.Lock()
|
|
defer mu.Unlock()
|
|
|
|
retain := make(map[string]bool)
|
|
|
|
for key, cfg := range newConfigs {
|
|
newDSN := buildDSN(cfg)
|
|
oldDSN, exists := configsDSN[key]
|
|
|
|
// 新增数据源或配置变更
|
|
if !exists || oldDSN != newDSN {
|
|
// 关闭旧连接(如果存在)
|
|
if exists {
|
|
if oldDB := dataSources[key]; oldDB != nil {
|
|
_ = oldDB.Close()
|
|
}
|
|
}
|
|
|
|
db, err := sql.Open("postgres", newDSN)
|
|
if err != nil {
|
|
log.Printf("[db] 数据源 %s Open 失败: %v", key, err)
|
|
continue
|
|
}
|
|
|
|
// 尝试实际连接数据库
|
|
if err := db.Ping(); err != nil {
|
|
log.Printf("[db] 数据源 %s 连接失败: %v", key, err)
|
|
_ = db.Close()
|
|
continue
|
|
}
|
|
|
|
applyDBConfig(db, cfg)
|
|
dataSources[key] = db
|
|
configsDSN[key] = newDSN
|
|
|
|
if !exists {
|
|
log.Printf("[db] 新增数据源 %s 成功", key)
|
|
} else {
|
|
log.Printf("[db] 更新数据源 %s 成功", key)
|
|
}
|
|
}
|
|
|
|
// 配置没变,不用操作
|
|
retain[key] = true
|
|
}
|
|
|
|
// 关闭并移除不再需要的数据源
|
|
for key, db := range dataSources {
|
|
if !retain[key] {
|
|
_ = db.Close()
|
|
delete(dataSources, key)
|
|
delete(configsDSN, key)
|
|
log.Printf("[db] 移除数据源 %s", key)
|
|
}
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// applyDBConfig 设置 sql.DB 连接池参数,并应用默认值
|
|
func applyDBConfig(db *sql.DB, cfg DBConfig) {
|
|
maxOpen := cfg.MaxOpenConns
|
|
if maxOpen <= 0 {
|
|
maxOpen = defaultMaxOpenConns
|
|
}
|
|
db.SetMaxOpenConns(maxOpen)
|
|
|
|
maxIdle := cfg.MaxIdleConns
|
|
if maxIdle <= 0 {
|
|
maxIdle = defaultMaxIdleConns
|
|
}
|
|
db.SetMaxIdleConns(maxIdle)
|
|
|
|
timeout := cfg.ConnMaxLifetime
|
|
if timeout <= 0 {
|
|
timeout = defaultConnMaxLifetime
|
|
}
|
|
db.SetConnMaxLifetime(timeout)
|
|
}
|
|
|
|
// GetDB 获取数据库连接
|
|
func GetDB(key string) (*sql.DB, bool) {
|
|
mu.RLock()
|
|
defer mu.RUnlock()
|
|
db, ok := dataSources[key]
|
|
return db, ok
|
|
}
|