release: v1.0.3 bug fix release
v1.0.3 定位为 bug fix release,收口 v1.0.2 引入的破坏性清理并修复 4 个轻量 bug + 依赖复查 + 版本号治理 + 文档对齐。同时包含此前未提交 的 v1.0.2 全部工作。 v1.0.3 变更: Fixed: - generateJTI 忽略 rand.Read 错误(jwt)— 改为 (string, error) 并传播 - QueryBuilder.Page Count 受残留 Limit 截断(repository)— countDB 加 Limit(-1).Offset(-1) - OSS/本地存储文件名冲突(storage)— 新增 uniqueFilename 加 8 字节随机后缀 - 数据库重试策略对不可恢复错误无效(database)— 新增 isTransientDBError Dependencies: - go mod tidy 补全 postgres 方言传递依赖;安全补丁升级 x/crypto v0.49→v0.53、golang-jwt/jwt/v5 v5.2.1→v5.3.1、gorilla/websocket v1.5.1→v1.5.3 Removed: - Breaking: 清理 v1.0.2 兼容别名(InitMySQL* / driverName) - 死代码 database.DBResolver - 代码中 292 行"评分/理由"自夸注释(#26,23 个文件) Changed: - Breaking: 错误码体系重构 CodeSuccess 1→0、CodeFail 0→1,删除 CodeInvalidParams,加编译期防撞码 - database/mysql.go → manager.go;Logger 拆分三独立 core 修复 Tee 重复写入 - 版本号常量化:app.go 新增 const Version 作为唯一来源,CLI/脚手架模板引用之,不再散落字面量 Security: - CORS 中间件按 W3C 规范修复 Allow-Credentials / Vary: Origin / 非白名单不回显 Added: - console 包显式 level 控制(SetLevel/WithLevel/LevelSilent,atomic.Int32 并发安全) - 新增测试 jwt/repository/storage/database/router/middleware/app - 文档:新增 CHANGELOG.md、Version_v1.0.2_report.md;更新 README/GUIDE 顶部更新日志 - 文档对齐:删除对不存在的 config.DefaultManager() getter 的描述(保持 API 与文档一致) 详见 CHANGELOG.md 与 README.md 更新日志。 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,139 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/EthanCodeCraft/xlgo-core/config"
|
||||
"gorm.io/driver/mysql"
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// 内置驱动常量(更多驱动可通过 RegisterDialect 扩展)
|
||||
const (
|
||||
DriverMySQL = config.DriverMySQL
|
||||
DriverPostgres = config.DriverPostgres
|
||||
)
|
||||
|
||||
// DialectorFactory 根据 DSN 返回 GORM Dialector
|
||||
type DialectorFactory func(dsn string) gorm.Dialector
|
||||
|
||||
// DialectSpec 描述一种数据库方言:如何建立连接 + 如何拼接 DSN
|
||||
type DialectSpec struct {
|
||||
// Name 驱动主名称(如 "mysql"、"postgres"、"sqlite"),大小写不敏感
|
||||
Name string
|
||||
// Aliases 驱动别名(如 postgres 的 "postgresql"、"pg")
|
||||
Aliases []string
|
||||
// Dialector 由 DSN 构造 GORM Dialector
|
||||
Dialector DialectorFactory
|
||||
// DSN 由 DatabaseConfig 拼接连接字符串。可选——
|
||||
// 不提供时使用 cfg.MySQLDSN() 兜底(适合自定义驱动通过 CustomDSN 指定连接串的场景)
|
||||
DSN config.DSNBuilder
|
||||
}
|
||||
|
||||
var (
|
||||
dialectsMu sync.RWMutex
|
||||
dialects = map[string]DialectorFactory{}
|
||||
)
|
||||
|
||||
// RegisterDialect 注册一种数据库方言。
|
||||
// 同时把 DSN 构建器登记到 config 包,使 cfg.Database.DSN() 也能识别新驱动。
|
||||
// 已注册的同名驱动会被覆盖。
|
||||
//
|
||||
// 用法示例(接入 SQLite):
|
||||
//
|
||||
// import "gorm.io/driver/sqlite"
|
||||
//
|
||||
// database.RegisterDialect(database.DialectSpec{
|
||||
// Name: "sqlite",
|
||||
// Dialector: func(dsn string) gorm.Dialector { return sqlite.Open(dsn) },
|
||||
// DSN: func(c *config.DatabaseConfig) string { return c.Name }, // 文件路径
|
||||
// })
|
||||
func RegisterDialect(spec DialectSpec) {
|
||||
if spec.Dialector == nil || strings.TrimSpace(spec.Name) == "" {
|
||||
return
|
||||
}
|
||||
|
||||
dialectsMu.Lock()
|
||||
for _, n := range append([]string{spec.Name}, spec.Aliases...) {
|
||||
key := normalizeDriver(n)
|
||||
if key != "" {
|
||||
dialects[key] = spec.Dialector
|
||||
}
|
||||
}
|
||||
dialectsMu.Unlock()
|
||||
|
||||
if spec.DSN != nil {
|
||||
config.RegisterDSNBuilder(spec.Name, spec.DSN, spec.Aliases...)
|
||||
}
|
||||
}
|
||||
|
||||
// LookupDialect 查找已注册的 Dialector 工厂
|
||||
func LookupDialect(driver string) (DialectorFactory, bool) {
|
||||
key := normalizeDriver(driver)
|
||||
dialectsMu.RLock()
|
||||
defer dialectsMu.RUnlock()
|
||||
f, ok := dialects[key]
|
||||
return f, ok
|
||||
}
|
||||
|
||||
// RegisteredDialects 返回所有已注册的驱动名(用于诊断)
|
||||
func RegisteredDialects() []string {
|
||||
dialectsMu.RLock()
|
||||
defer dialectsMu.RUnlock()
|
||||
names := make([]string, 0, len(dialects))
|
||||
for k := range dialects {
|
||||
names = append(names, k)
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// Dialector 根据配置返回 GORM Dialector。
|
||||
// 驱动由 cfg.Database.Driver 决定,未指定或未注册时按 MySQL 兜底(向后兼容)。
|
||||
func Dialector(cfg *config.Config) gorm.Dialector {
|
||||
return dialectorForDSN(cfg.Database.Driver, cfg.Database.DSN())
|
||||
}
|
||||
|
||||
// dialectorForDSN 根据驱动名和 DSN 返回 Dialector
|
||||
func dialectorForDSN(driver, dsn string) gorm.Dialector {
|
||||
if f, ok := LookupDialect(driver); ok {
|
||||
return f(dsn)
|
||||
}
|
||||
// 未注册时回退到 MySQL,与 config.DSN() 的回退保持一致
|
||||
return mysql.Open(dsn)
|
||||
}
|
||||
|
||||
// normalizeDriver 规范化驱动名(小写、去空白)
|
||||
func normalizeDriver(name string) string {
|
||||
return strings.ToLower(strings.TrimSpace(name))
|
||||
}
|
||||
|
||||
// driverDescription 返回带别名提示的驱动描述(用于错误信息和日志)
|
||||
func driverDescription(driver string) string {
|
||||
key := normalizeDriver(driver)
|
||||
if key == "" {
|
||||
return DriverMySQL + " (default)"
|
||||
}
|
||||
if _, ok := LookupDialect(key); ok {
|
||||
return key
|
||||
}
|
||||
return fmt.Sprintf("%s (unregistered, fallback=%s)", key, DriverMySQL)
|
||||
}
|
||||
|
||||
func init() {
|
||||
// 内置 MySQL
|
||||
RegisterDialect(DialectSpec{
|
||||
Name: DriverMySQL,
|
||||
Dialector: func(dsn string) gorm.Dialector { return mysql.Open(dsn) },
|
||||
DSN: func(c *config.DatabaseConfig) string { return c.MySQLDSN() },
|
||||
})
|
||||
// 内置 PostgreSQL
|
||||
RegisterDialect(DialectSpec{
|
||||
Name: DriverPostgres,
|
||||
Aliases: []string{"postgresql", "pg"},
|
||||
Dialector: func(dsn string) gorm.Dialector { return postgres.Open(dsn) },
|
||||
DSN: func(c *config.DatabaseConfig) string { return c.PostgresDSN() },
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,473 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/EthanCodeCraft/xlgo-core/config"
|
||||
"github.com/EthanCodeCraft/xlgo-core/logger"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
gormlogger "gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
type dbModeContextKey struct{}
|
||||
|
||||
const (
|
||||
dbModeMaster = "master"
|
||||
dbModeReplica = "replica"
|
||||
)
|
||||
|
||||
// ReplicaPicker 从库选择策略
|
||||
type ReplicaPicker interface {
|
||||
Pick(replicas []*gorm.DB) *gorm.DB
|
||||
}
|
||||
|
||||
// RoundRobinPicker 轮询选择从库
|
||||
type RoundRobinPicker struct {
|
||||
counter uint64
|
||||
}
|
||||
|
||||
// Pick 轮询选择一个从库
|
||||
func (p *RoundRobinPicker) Pick(replicas []*gorm.DB) *gorm.DB {
|
||||
if len(replicas) == 0 {
|
||||
return nil
|
||||
}
|
||||
n := atomic.AddUint64(&p.counter, 1)
|
||||
return replicas[int(n-1)%len(replicas)]
|
||||
}
|
||||
|
||||
// RandomPicker 随机选择从库
|
||||
type RandomPicker struct{}
|
||||
|
||||
// Pick 随机选择一个从库
|
||||
func (p *RandomPicker) Pick(replicas []*gorm.DB) *gorm.DB {
|
||||
if len(replicas) == 0 {
|
||||
return nil
|
||||
}
|
||||
return replicas[rand.Intn(len(replicas))]
|
||||
}
|
||||
|
||||
// Manager 数据库管理器,持有主库与从库连接实例
|
||||
type Manager struct {
|
||||
cfg *config.Config
|
||||
master *gorm.DB
|
||||
replicas []*gorm.DB
|
||||
picker ReplicaPicker
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// NewManager 创建数据库管理器
|
||||
func NewManager(cfg *config.Config) *Manager {
|
||||
return &Manager{cfg: cfg, picker: &RandomPicker{}}
|
||||
}
|
||||
|
||||
// SetPicker 设置从库选择策略
|
||||
func (m *Manager) SetPicker(p ReplicaPicker) {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
m.mu.Lock()
|
||||
m.picker = p
|
||||
m.mu.Unlock()
|
||||
}
|
||||
|
||||
// Picker 返回当前从库选择策略
|
||||
func (m *Manager) Picker() ReplicaPicker {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
return m.picker
|
||||
}
|
||||
|
||||
// Master 返回主库实例
|
||||
func (m *Manager) Master() *gorm.DB {
|
||||
return m.master
|
||||
}
|
||||
|
||||
// Replicas 返回所有从库实例
|
||||
func (m *Manager) Replicas() []*gorm.DB {
|
||||
return m.replicas
|
||||
}
|
||||
|
||||
// Replica 按策略选择一个从库;无从库时返回主库
|
||||
func (m *Manager) Replica() *gorm.DB {
|
||||
if len(m.replicas) == 0 {
|
||||
return m.master
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.picker != nil {
|
||||
if db := m.picker.Pick(m.replicas); db != nil {
|
||||
return db
|
||||
}
|
||||
}
|
||||
return m.replicas[0]
|
||||
}
|
||||
|
||||
// FromContext 根据上下文选择数据库
|
||||
func (m *Manager) FromContext(ctx context.Context) *gorm.DB {
|
||||
mode, ok := ctx.Value(dbModeContextKey{}).(string)
|
||||
if !ok {
|
||||
return m.Replica()
|
||||
}
|
||||
switch mode {
|
||||
case dbModeMaster:
|
||||
return m.master
|
||||
case dbModeReplica:
|
||||
return m.Replica()
|
||||
default:
|
||||
return m.Replica()
|
||||
}
|
||||
}
|
||||
|
||||
// Open 打开主库连接
|
||||
func (m *Manager) Open(ctx context.Context) error {
|
||||
if m.cfg == nil {
|
||||
return errors.New("数据库配置未设置")
|
||||
}
|
||||
return m.InitDB(m.cfg)
|
||||
}
|
||||
|
||||
// OpenWithReplicas 打开主库与从库连接
|
||||
func (m *Manager) OpenWithReplicas(ctx context.Context, replicaDSNs []string) error {
|
||||
if m.cfg == nil {
|
||||
return errors.New("数据库配置未设置")
|
||||
}
|
||||
return m.InitDBWithReplicas(m.cfg, replicaDSNs)
|
||||
}
|
||||
|
||||
// Close 关闭主库与全部从库连接
|
||||
func (m *Manager) Close() error {
|
||||
var errs []error
|
||||
|
||||
if m.master != nil {
|
||||
sqlDB, err := m.master.DB()
|
||||
if err != nil {
|
||||
errs = append(errs, err)
|
||||
} else if err := sqlDB.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
for _, replica := range m.replicas {
|
||||
if replica == nil {
|
||||
continue
|
||||
}
|
||||
sqlDB, err := replica.DB()
|
||||
if err != nil {
|
||||
errs = append(errs, err)
|
||||
continue
|
||||
}
|
||||
if err := sqlDB.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
m.master = nil
|
||||
m.replicas = nil
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
// HealthCheck 健康检查,主库不可达时返回错误
|
||||
func (m *Manager) HealthCheck(ctx context.Context) error {
|
||||
if m.master == nil {
|
||||
return errors.New("database master not initialized")
|
||||
}
|
||||
sqlDB, err := m.master.DB()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return sqlDB.PingContext(ctx)
|
||||
}
|
||||
|
||||
// DefaultManager 默认数据库管理器
|
||||
var DefaultManager = &Manager{picker: &RandomPicker{}}
|
||||
|
||||
// InitDB 初始化数据库连接(带重试机制),驱动由配置决定
|
||||
func (m *Manager) InitDB(cfg *config.Config) error {
|
||||
var err error
|
||||
m.cfg = cfg
|
||||
|
||||
// GORM 日志配置
|
||||
var gormLogLevel gormlogger.LogLevel
|
||||
if cfg.IsDevelopment() {
|
||||
gormLogLevel = gormlogger.Info
|
||||
} else {
|
||||
gormLogLevel = gormlogger.Warn
|
||||
}
|
||||
|
||||
gormConfig := &gorm.Config{
|
||||
Logger: gormlogger.Default.LogMode(gormLogLevel),
|
||||
}
|
||||
|
||||
// 重试配置
|
||||
maxRetries := 5
|
||||
retryDelay := time.Second
|
||||
|
||||
var lastErr error
|
||||
for i := range maxRetries {
|
||||
// 连接主库
|
||||
m.master, err = gorm.Open(Dialector(cfg), gormConfig)
|
||||
if err == nil {
|
||||
sqlDB, err := m.master.DB()
|
||||
if err == nil {
|
||||
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns)
|
||||
sqlDB.SetMaxOpenConns(cfg.Database.MaxOpenConns)
|
||||
sqlDB.SetConnMaxLifetime(time.Hour)
|
||||
|
||||
if err := sqlDB.Ping(); err == nil {
|
||||
logger.Info("数据库主库连接成功",
|
||||
zap.String("driver", driverDescription(cfg.Database.Driver)),
|
||||
zap.String("host", cfg.Database.Host),
|
||||
zap.Int("port", cfg.Database.Port))
|
||||
return nil
|
||||
} else {
|
||||
// Ping 失败(如服务端暂时不可达)视作可重试
|
||||
lastErr = err
|
||||
}
|
||||
} else {
|
||||
lastErr = err
|
||||
}
|
||||
} else {
|
||||
lastErr = err
|
||||
// 不可恢复的错误(认证失败、未知数据库、DSN 非法等)直接返回,不必重试
|
||||
if !isTransientDBError(err) {
|
||||
return fmt.Errorf("数据库连接失败(不可恢复): %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
logger.Warnf("数据库连接失败,第 %d/%d 次重试: %v", i+1, maxRetries, lastErr)
|
||||
time.Sleep(retryDelay)
|
||||
retryDelay *= 2
|
||||
if retryDelay > 30*time.Second {
|
||||
retryDelay = 30 * time.Second
|
||||
}
|
||||
}
|
||||
|
||||
return fmt.Errorf("数据库连接失败(重试 %d 次): %w", maxRetries, lastErr)
|
||||
}
|
||||
|
||||
// isTransientDBError 判断数据库连接错误是否值得重试。
|
||||
// 认证失败、未知数据库、非法 DSN/驱动等属于配置类错误,重试无意义,直接返回更友好。
|
||||
func isTransientDBError(err error) bool {
|
||||
if err == nil {
|
||||
return true
|
||||
}
|
||||
msg := err.Error()
|
||||
nonTransient := []string{
|
||||
"Access denied", // MySQL 认证失败(用户名/密码错误)
|
||||
"authentication plugin", // MySQL 认证插件不支持
|
||||
"Unknown database", // MySQL 目标库不存在
|
||||
"invalid DSN", // DSN 语法错误
|
||||
"unknown driver", // 驱动未注册
|
||||
"unsupported driver", // 驱动不支持
|
||||
}
|
||||
for _, sub := range nonTransient {
|
||||
if strings.Contains(msg, sub) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// InitDBWithReplicas 初始化数据库主从连接,驱动由配置决定
|
||||
// replicaDSNs: 从库连接字符串列表(需与主库驱动匹配)
|
||||
func (m *Manager) InitDBWithReplicas(cfg *config.Config, replicaDSNs []string) error {
|
||||
// 先初始化主库
|
||||
if err := m.InitDB(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
m.replicas = nil
|
||||
|
||||
// 初始化从库
|
||||
if len(replicaDSNs) > 0 {
|
||||
var gormLogLevel gormlogger.LogLevel
|
||||
if cfg.IsDevelopment() {
|
||||
gormLogLevel = gormlogger.Info
|
||||
} else {
|
||||
gormLogLevel = gormlogger.Warn
|
||||
}
|
||||
|
||||
gormConfig := &gorm.Config{
|
||||
Logger: gormlogger.Default.LogMode(gormLogLevel),
|
||||
}
|
||||
|
||||
for i, dsn := range replicaDSNs {
|
||||
replicaDB, err := gorm.Open(dialectorForDSN(cfg.Database.Driver, dsn), gormConfig)
|
||||
if err != nil {
|
||||
logger.Warnf("数据库从库 %d 连接失败: %v", i+1, err)
|
||||
continue
|
||||
}
|
||||
|
||||
sqlDB, err := replicaDB.DB()
|
||||
if err != nil {
|
||||
logger.Warnf("数据库从库 %d 获取连接池失败: %v", i+1, err)
|
||||
continue
|
||||
}
|
||||
|
||||
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns)
|
||||
sqlDB.SetMaxOpenConns(cfg.Database.MaxOpenConns / 2) // 从库连接数可适当减少
|
||||
sqlDB.SetConnMaxLifetime(time.Hour)
|
||||
|
||||
if err := sqlDB.Ping(); err != nil {
|
||||
logger.Warnf("数据库从库 %d Ping 失败: %v", i+1, err)
|
||||
continue
|
||||
}
|
||||
|
||||
m.replicas = append(m.replicas, replicaDB)
|
||||
logger.Info("数据库从库连接成功", zap.Int("index", i+1))
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// InitDB 初始化数据库连接(带重试机制),驱动由配置决定
|
||||
func InitDB(cfg *config.Config) error {
|
||||
return DefaultManager.InitDB(cfg)
|
||||
}
|
||||
|
||||
// InitDBWithReplicas 初始化数据库主从连接,驱动由配置决定
|
||||
func InitDBWithReplicas(cfg *config.Config, replicaDSNs []string) error {
|
||||
return DefaultManager.InitDBWithReplicas(cfg, replicaDSNs)
|
||||
}
|
||||
|
||||
// GetReadDB 获取读库实例(按策略选择从库)
|
||||
func GetReadDB() *gorm.DB {
|
||||
return DefaultManager.Replica()
|
||||
}
|
||||
|
||||
// GetWriteDB 获取写库实例(主库)
|
||||
func GetWriteDB() *gorm.DB {
|
||||
return DefaultManager.Master()
|
||||
}
|
||||
|
||||
// GetDB 获取数据库实例(默认主库,兼容旧代码)
|
||||
func GetDB() *gorm.DB {
|
||||
return DefaultManager.Master()
|
||||
}
|
||||
|
||||
// GetReplicas 获取所有从库实例
|
||||
func GetReplicas() []*gorm.DB {
|
||||
return DefaultManager.Replicas()
|
||||
}
|
||||
|
||||
// SetReplicaPicker 设置默认管理器的从库选择策略
|
||||
func SetReplicaPicker(p ReplicaPicker) {
|
||||
DefaultManager.SetPicker(p)
|
||||
}
|
||||
|
||||
// UseMaster 强制使用主库(用于事务或需要实时数据的场景)
|
||||
func UseMaster(ctx context.Context) context.Context {
|
||||
return context.WithValue(ctx, dbModeContextKey{}, dbModeMaster)
|
||||
}
|
||||
|
||||
// UseReplica 强制使用从库(用于报表查询等场景)
|
||||
func UseReplica(ctx context.Context) context.Context {
|
||||
return context.WithValue(ctx, dbModeContextKey{}, dbModeReplica)
|
||||
}
|
||||
|
||||
// GetDBFromContext 根据上下文选择数据库
|
||||
func GetDBFromContext(ctx context.Context) *gorm.DB {
|
||||
return DefaultManager.FromContext(ctx)
|
||||
}
|
||||
|
||||
// AutoMigrate 自动迁移数据库表结构(由应用通过 WithMigrator/WithModels 注册)
|
||||
func AutoMigrate() error {
|
||||
logger.Info("数据库表结构迁移完成")
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close 关闭主库连接(兼容旧代码,从库连接请使用 CloseAll)
|
||||
func Close() error {
|
||||
if DefaultManager.master == nil {
|
||||
return nil
|
||||
}
|
||||
sqlDB, err := DefaultManager.master.DB()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = sqlDB.Close()
|
||||
DefaultManager.master = nil
|
||||
return err
|
||||
}
|
||||
|
||||
// CloseAll 关闭所有数据库连接(包括从库)
|
||||
func CloseAll() error {
|
||||
return DefaultManager.Close()
|
||||
}
|
||||
|
||||
// Transaction 事务操作(自动使用主库)
|
||||
func Transaction(fn func(tx *gorm.DB) error) error {
|
||||
if DefaultManager.master == nil {
|
||||
return errors.New("数据库未初始化")
|
||||
}
|
||||
return DefaultManager.master.Transaction(fn)
|
||||
}
|
||||
|
||||
// TransactionWithContext 带上下文的事务操作
|
||||
func TransactionWithContext(ctx context.Context, fn func(tx *gorm.DB) error) error {
|
||||
if DefaultManager.master == nil {
|
||||
return errors.New("数据库未初始化")
|
||||
}
|
||||
return DefaultManager.master.WithContext(ctx).Transaction(fn)
|
||||
}
|
||||
|
||||
// ReadQuery 读查询(自动路由到从库)
|
||||
func ReadQuery(ctx context.Context, model any, query string, args ...any) error {
|
||||
db := GetDBFromContext(ctx)
|
||||
if db == nil {
|
||||
return errors.New("数据库未初始化")
|
||||
}
|
||||
return db.WithContext(ctx).Where(query, args...).Find(model).Error
|
||||
}
|
||||
|
||||
// WriteQuery 写查询(强制使用主库)
|
||||
func WriteQuery(ctx context.Context, model any, query string, args ...any) error {
|
||||
if DefaultManager.master == nil {
|
||||
return errors.New("数据库未初始化")
|
||||
}
|
||||
return DefaultManager.master.WithContext(ctx).Where(query, args...).Find(model).Error
|
||||
}
|
||||
|
||||
// HealthCheck 健康检查
|
||||
func HealthCheck() map[string]bool {
|
||||
result := make(map[string]bool)
|
||||
|
||||
// 检查主库
|
||||
if DefaultManager.master != nil {
|
||||
sqlDB, err := DefaultManager.master.DB()
|
||||
if err == nil && sqlDB.Ping() == nil {
|
||||
result["master"] = true
|
||||
} else {
|
||||
result["master"] = false
|
||||
}
|
||||
} else {
|
||||
result["master"] = false
|
||||
}
|
||||
|
||||
// 检查从库
|
||||
for i, replica := range DefaultManager.replicas {
|
||||
if replica != nil {
|
||||
sqlDB, err := replica.DB()
|
||||
if err == nil && sqlDB.Ping() == nil {
|
||||
result[fmt.Sprintf("replica_%d", i+1)] = true
|
||||
} else {
|
||||
result[fmt.Sprintf("replica_%d", i+1)] = false
|
||||
}
|
||||
} else {
|
||||
result[fmt.Sprintf("replica_%d", i+1)] = false
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,216 @@
|
||||
package database_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/EthanCodeCraft/xlgo-core/config"
|
||||
"github.com/EthanCodeCraft/xlgo-core/database"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
"gorm.io/gorm/schema"
|
||||
)
|
||||
|
||||
func TestCloseAllWithoutInit(t *testing.T) {
|
||||
if err := database.CloseAll(); err != nil {
|
||||
t.Fatalf("CloseAll without init should not error: %v", err)
|
||||
}
|
||||
if database.GetDB() != nil {
|
||||
t.Fatal("expected DB nil")
|
||||
}
|
||||
if database.GetReadDB() != nil {
|
||||
t.Fatal("expected read DB nil")
|
||||
}
|
||||
if len(database.GetReplicas()) != 0 {
|
||||
t.Fatal("expected replicas empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDBContextHelpersWithoutInit(t *testing.T) {
|
||||
ctx := database.UseMaster(context.Background())
|
||||
if db := database.GetDBFromContext(ctx); db != nil {
|
||||
t.Fatal("expected nil DB without init")
|
||||
}
|
||||
|
||||
ctx = database.UseReplica(context.Background())
|
||||
if db := database.GetDBFromContext(ctx); db != nil {
|
||||
t.Fatal("expected nil read DB without init")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRoundRobinPicker(t *testing.T) {
|
||||
replicas := []*gorm.DB{{}, {}, {}}
|
||||
p := &database.RoundRobinPicker{}
|
||||
|
||||
first := p.Pick(replicas)
|
||||
second := p.Pick(replicas)
|
||||
third := p.Pick(replicas)
|
||||
fourth := p.Pick(replicas)
|
||||
|
||||
if first == nil || second == nil || third == nil {
|
||||
t.Fatal("Picker returned nil for non-empty replicas")
|
||||
}
|
||||
if first != replicas[0] || second != replicas[1] || third != replicas[2] {
|
||||
t.Fatal("RoundRobinPicker should cycle through replicas in order")
|
||||
}
|
||||
if fourth != replicas[0] {
|
||||
t.Fatal("RoundRobinPicker should wrap around to the first replica")
|
||||
}
|
||||
if p.Pick(nil) != nil || p.Pick([]*gorm.DB{}) != nil {
|
||||
t.Fatal("Picker should return nil for empty replicas")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRandomPicker(t *testing.T) {
|
||||
replicas := []*gorm.DB{{}, {}}
|
||||
p := &database.RandomPicker{}
|
||||
|
||||
picked := p.Pick(replicas)
|
||||
if picked == nil {
|
||||
t.Fatal("RandomPicker returned nil for non-empty replicas")
|
||||
}
|
||||
if picked != replicas[0] && picked != replicas[1] {
|
||||
t.Fatal("RandomPicker returned a replica not in the slice")
|
||||
}
|
||||
if p.Pick(nil) != nil || p.Pick([]*gorm.DB{}) != nil {
|
||||
t.Fatal("RandomPicker should return nil for empty replicas")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerReplicaFallbackToMaster(t *testing.T) {
|
||||
mgr := database.NewManager(&config.Config{})
|
||||
if mgr.Master() != nil {
|
||||
t.Fatal("expected nil master before init")
|
||||
}
|
||||
if mgr.Replicas() != nil {
|
||||
t.Fatal("expected nil replicas before init")
|
||||
}
|
||||
// 无从库时 Replica 应返回 master(此处均为 nil)
|
||||
if mgr.Replica() != nil {
|
||||
t.Fatal("expected Replica to fall back to master when no replicas")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerSetPicker(t *testing.T) {
|
||||
mgr := database.NewManager(&config.Config{})
|
||||
rr := &database.RoundRobinPicker{}
|
||||
mgr.SetPicker(rr)
|
||||
if mgr.Picker() != rr {
|
||||
t.Fatal("SetPicker did not install the picker")
|
||||
}
|
||||
// nil 不应覆盖已有 picker
|
||||
mgr.SetPicker(nil)
|
||||
if mgr.Picker() != rr {
|
||||
t.Fatal("SetPicker(nil) should not clear the existing picker")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultManagerHealthCheckWithoutInit(t *testing.T) {
|
||||
if err := database.DefaultManager.HealthCheck(context.Background()); err == nil {
|
||||
t.Fatal("expected error when health checking uninitialized master")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDialectorSelectsByDriver(t *testing.T) {
|
||||
mysqlCfg := &config.Config{Database: config.DatabaseConfig{
|
||||
Driver: config.DriverMySQL, Host: "localhost", Port: 3306,
|
||||
User: "root", Password: "pass", Name: "db",
|
||||
}}
|
||||
if name := database.Dialector(mysqlCfg).Name(); name != "mysql" {
|
||||
t.Fatalf("expected mysql dialector, got %q", name)
|
||||
}
|
||||
|
||||
pgCfg := &config.Config{Database: config.DatabaseConfig{
|
||||
Driver: config.DriverPostgres, Host: "localhost", Port: 5432,
|
||||
User: "postgres", Password: "pass", Name: "db",
|
||||
}}
|
||||
if name := database.Dialector(pgCfg).Name(); name != "postgres" {
|
||||
t.Fatalf("expected postgres dialector, got %q", name)
|
||||
}
|
||||
|
||||
// 别名也应解析为 postgres
|
||||
pgAliasCfg := &config.Config{Database: config.DatabaseConfig{
|
||||
Driver: "PG", Host: "localhost", Port: 5432,
|
||||
User: "postgres", Password: "pass", Name: "db",
|
||||
}}
|
||||
if name := database.Dialector(pgAliasCfg).Name(); name != "postgres" {
|
||||
t.Fatalf("expected postgres dialector via alias, got %q", name)
|
||||
}
|
||||
|
||||
// 未指定 Driver 时默认 mysql
|
||||
defaultCfg := &config.Config{Database: config.DatabaseConfig{
|
||||
Host: "localhost", Port: 3306, User: "root", Password: "pass", Name: "db",
|
||||
}}
|
||||
if name := database.Dialector(defaultCfg).Name(); name != "mysql" {
|
||||
t.Fatalf("expected default mysql dialector, got %q", name)
|
||||
}
|
||||
}
|
||||
|
||||
// stubDialector 是一个用于测试 RegisterDialect 的占位 Dialector。
|
||||
type stubDialector struct{ name string }
|
||||
|
||||
func (s stubDialector) Name() string { return s.name }
|
||||
func (s stubDialector) Initialize(_ *gorm.DB) error { return nil }
|
||||
func (s stubDialector) Migrator(db *gorm.DB) gorm.Migrator { return nil }
|
||||
func (s stubDialector) DataTypeOf(*schema.Field) string { return "" }
|
||||
func (s stubDialector) DefaultValueOf(*schema.Field) clause.Expression { return nil }
|
||||
func (s stubDialector) BindVarTo(writer clause.Writer, _ *gorm.Statement, _ any) {}
|
||||
func (s stubDialector) QuoteTo(writer clause.Writer, str string) { _, _ = writer.WriteString(str) }
|
||||
func (s stubDialector) Explain(sql string, _ ...any) string { return sql }
|
||||
|
||||
func TestRegisterDialectAndCustomDriver(t *testing.T) {
|
||||
const driver = "stubdb"
|
||||
|
||||
database.RegisterDialect(database.DialectSpec{
|
||||
Name: driver,
|
||||
Aliases: []string{"stub"},
|
||||
Dialector: func(dsn string) gorm.Dialector { return stubDialector{name: "stubdb"} },
|
||||
DSN: func(c *config.DatabaseConfig) string {
|
||||
return "stub://" + c.Host
|
||||
},
|
||||
})
|
||||
|
||||
// Dialector 工厂可以解析主名和别名
|
||||
if _, ok := database.LookupDialect(driver); !ok {
|
||||
t.Fatal("expected stubdb dialector to be registered")
|
||||
}
|
||||
if _, ok := database.LookupDialect("STUB"); !ok {
|
||||
t.Fatal("expected stub alias to be registered (case-insensitive)")
|
||||
}
|
||||
|
||||
cfg := &config.Config{Database: config.DatabaseConfig{
|
||||
Driver: driver, Host: "localhost",
|
||||
}}
|
||||
if name := database.Dialector(cfg).Name(); name != "stubdb" {
|
||||
t.Fatalf("expected stubdb dialector, got %q", name)
|
||||
}
|
||||
|
||||
// config.DSN() 应使用注册的 DSN 构建器
|
||||
if dsn := cfg.Database.DSN(); dsn != "stub://localhost" {
|
||||
t.Fatalf("expected DSN built by registered builder, got %q", dsn)
|
||||
}
|
||||
|
||||
// 未知驱动回退到 mysql
|
||||
unknownCfg := &config.Config{Database: config.DatabaseConfig{
|
||||
Driver: "no-such-driver", Host: "localhost", Port: 3306,
|
||||
User: "root", Password: "pass", Name: "db",
|
||||
}}
|
||||
if name := database.Dialector(unknownCfg).Name(); name != "mysql" {
|
||||
t.Fatalf("expected fallback to mysql for unknown driver, got %q", name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisteredDialectsContainsBuiltins(t *testing.T) {
|
||||
registered := database.RegisteredDialects()
|
||||
want := map[string]bool{"mysql": false, "postgres": false, "pg": false}
|
||||
for _, n := range registered {
|
||||
if _, ok := want[n]; ok {
|
||||
want[n] = true
|
||||
}
|
||||
}
|
||||
for k, found := range want {
|
||||
if !found {
|
||||
t.Errorf("expected %q to be registered by default", k)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,315 +0,0 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/EthanCodeCraft/xlgo-core/config"
|
||||
"github.com/EthanCodeCraft/xlgo-core/logger"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/driver/mysql"
|
||||
"gorm.io/gorm"
|
||||
gormlogger "gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
var (
|
||||
// DB 主库实例(写操作)
|
||||
DB *gorm.DB
|
||||
// DBRead 读库实例(读操作)
|
||||
DBRead *gorm.DB
|
||||
// replicas 从库列表
|
||||
replicas []*gorm.DB
|
||||
// replicaMutex 从库选择锁
|
||||
replicaMutex sync.Mutex
|
||||
)
|
||||
|
||||
// InitMySQL 初始化 MySQL 连接(带重试机制)
|
||||
func InitMySQL(cfg *config.Config) error {
|
||||
var err error
|
||||
|
||||
// GORM 日志配置
|
||||
var gormLogLevel gormlogger.LogLevel
|
||||
if cfg.IsDevelopment() {
|
||||
gormLogLevel = gormlogger.Info
|
||||
} else {
|
||||
gormLogLevel = gormlogger.Warn
|
||||
}
|
||||
|
||||
gormConfig := &gorm.Config{
|
||||
Logger: gormlogger.Default.LogMode(gormLogLevel),
|
||||
}
|
||||
|
||||
// 重试配置
|
||||
maxRetries := 5
|
||||
retryDelay := time.Second
|
||||
|
||||
for i := range maxRetries {
|
||||
// 连接主库
|
||||
DB, err = gorm.Open(mysql.Open(cfg.Database.DSN()), gormConfig)
|
||||
if err == nil {
|
||||
sqlDB, err := DB.DB()
|
||||
if err == nil {
|
||||
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns)
|
||||
sqlDB.SetMaxOpenConns(cfg.Database.MaxOpenConns)
|
||||
sqlDB.SetConnMaxLifetime(time.Hour)
|
||||
|
||||
if err := sqlDB.Ping(); err == nil {
|
||||
logger.Info("MySQL 主库连接成功", zap.String("host", cfg.Database.Host), zap.Int("port", cfg.Database.Port))
|
||||
return nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
logger.Warnf("MySQL 连接失败,第 %d/%d 次重试: %v", i+1, maxRetries, err)
|
||||
time.Sleep(retryDelay)
|
||||
retryDelay *= 2
|
||||
if retryDelay > 30*time.Second {
|
||||
retryDelay = 30 * time.Second
|
||||
}
|
||||
}
|
||||
|
||||
return fmt.Errorf("MySQL 连接失败(重试 %d 次): %w", maxRetries, err)
|
||||
}
|
||||
|
||||
// InitMySQLWithReplicas 初始化 MySQL 主从连接
|
||||
// masterDSN: 主库连接字符串
|
||||
// replicaDSNs: 从库连接字符串列表
|
||||
func InitMySQLWithReplicas(cfg *config.Config, replicaDSNs []string) error {
|
||||
// 先初始化主库
|
||||
if err := InitMySQL(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 初始化从库
|
||||
if len(replicaDSNs) > 0 {
|
||||
var gormLogLevel gormlogger.LogLevel
|
||||
if cfg.IsDevelopment() {
|
||||
gormLogLevel = gormlogger.Info
|
||||
} else {
|
||||
gormLogLevel = gormlogger.Warn
|
||||
}
|
||||
|
||||
gormConfig := &gorm.Config{
|
||||
Logger: gormlogger.Default.LogMode(gormLogLevel),
|
||||
}
|
||||
|
||||
for i, dsn := range replicaDSNs {
|
||||
replicaDB, err := gorm.Open(mysql.Open(dsn), gormConfig)
|
||||
if err != nil {
|
||||
logger.Warnf("MySQL 从库 %d 连接失败: %v", i+1, err)
|
||||
continue
|
||||
}
|
||||
|
||||
sqlDB, err := replicaDB.DB()
|
||||
if err != nil {
|
||||
logger.Warnf("MySQL 从库 %d 获取连接池失败: %v", i+1, err)
|
||||
continue
|
||||
}
|
||||
|
||||
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns)
|
||||
sqlDB.SetMaxOpenConns(cfg.Database.MaxOpenConns / 2) // 从库连接数可适当减少
|
||||
sqlDB.SetConnMaxLifetime(time.Hour)
|
||||
|
||||
if err := sqlDB.Ping(); err != nil {
|
||||
logger.Warnf("MySQL 从库 %d Ping 失败: %v", i+1, err)
|
||||
continue
|
||||
}
|
||||
|
||||
replicas = append(replicas, replicaDB)
|
||||
logger.Info("MySQL 从库连接成功", zap.Int("index", i+1))
|
||||
}
|
||||
|
||||
// 设置默认读库
|
||||
if len(replicas) > 0 {
|
||||
DBRead = replicas[0]
|
||||
} else {
|
||||
DBRead = DB // 无从库时使用主库
|
||||
}
|
||||
} else {
|
||||
DBRead = DB // 无从库配置时使用主库
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetReadDB 获取读库实例(自动选择从库)
|
||||
func GetReadDB() *gorm.DB {
|
||||
if len(replicas) == 0 {
|
||||
return DB
|
||||
}
|
||||
|
||||
replicaMutex.Lock()
|
||||
defer replicaMutex.Unlock()
|
||||
|
||||
// 随机选择一个从库
|
||||
idx := rand.Intn(len(replicas))
|
||||
return replicas[idx]
|
||||
}
|
||||
|
||||
// GetWriteDB 获取写库实例(主库)
|
||||
func GetWriteDB() *gorm.DB {
|
||||
return DB
|
||||
}
|
||||
|
||||
// GetDB 获取数据库实例(默认主库,兼容旧代码)
|
||||
func GetDB() *gorm.DB {
|
||||
return DB
|
||||
}
|
||||
|
||||
// GetReplicas 获取所有从库实例
|
||||
func GetReplicas() []*gorm.DB {
|
||||
return replicas
|
||||
}
|
||||
|
||||
// UseMaster 强制使用主库(用于事务或需要实时数据的场景)
|
||||
func UseMaster(ctx context.Context) context.Context {
|
||||
return context.WithValue(ctx, "db_mode", "master")
|
||||
}
|
||||
|
||||
// UseReplica 强制使用从库(用于报表查询等场景)
|
||||
func UseReplica(ctx context.Context) context.Context {
|
||||
return context.WithValue(ctx, "db_mode", "replica")
|
||||
}
|
||||
|
||||
// GetDBFromContext 根据上下文选择数据库
|
||||
func GetDBFromContext(ctx context.Context) *gorm.DB {
|
||||
mode, ok := ctx.Value("db_mode").(string)
|
||||
if !ok {
|
||||
return GetReadDB()
|
||||
}
|
||||
|
||||
switch mode {
|
||||
case "master":
|
||||
return DB
|
||||
case "replica":
|
||||
return GetReadDB()
|
||||
default:
|
||||
return GetReadDB()
|
||||
}
|
||||
}
|
||||
|
||||
// DBResolver 数据库解析器(用于 GORM 钩子)
|
||||
type DBResolver struct{}
|
||||
|
||||
// BeforeQuery 查询前钩子,自动路由到从库
|
||||
func (r *DBResolver) BeforeQuery(db *gorm.DB) {
|
||||
// 如果在事务中,使用主库
|
||||
if db.Statement.ConnPool != nil {
|
||||
return
|
||||
}
|
||||
|
||||
// 检查上下文
|
||||
ctx := db.Statement.Context
|
||||
if ctx != nil {
|
||||
mode, ok := ctx.Value("db_mode").(string)
|
||||
if ok && mode == "master" {
|
||||
return // 强制主库
|
||||
}
|
||||
}
|
||||
|
||||
// 读操作路由到从库
|
||||
if len(replicas) > 0 && DBRead != nil {
|
||||
db.Statement.ConnPool = DBRead.Statement.ConnPool
|
||||
}
|
||||
}
|
||||
|
||||
// AutoMigrate 自动迁移数据库表结构(由应用重写)
|
||||
func AutoMigrate() error {
|
||||
logger.Info("数据库表结构迁移完成")
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close 关闭数据库连接
|
||||
func Close() error {
|
||||
if DB == nil {
|
||||
return nil
|
||||
}
|
||||
sqlDB, err := DB.DB()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return sqlDB.Close()
|
||||
}
|
||||
|
||||
// CloseAll 关闭所有数据库连接(包括从库)
|
||||
func CloseAll() error {
|
||||
// 关闭主库
|
||||
if err := Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 关闭从库
|
||||
for _, replica := range replicas {
|
||||
if replica == nil {
|
||||
continue
|
||||
}
|
||||
sqlDB, err := replica.DB()
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
sqlDB.Close()
|
||||
}
|
||||
|
||||
replicas = nil
|
||||
DBRead = nil
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Transaction 事务操作(自动使用主库)
|
||||
func Transaction(fn func(tx *gorm.DB) error) error {
|
||||
return DB.Transaction(fn)
|
||||
}
|
||||
|
||||
// TransactionWithContext 带上下文的事务操作
|
||||
func TransactionWithContext(ctx context.Context, fn func(tx *gorm.DB) error) error {
|
||||
return DB.WithContext(ctx).Transaction(fn)
|
||||
}
|
||||
|
||||
// ReadQuery 读查询(自动路由到从库)
|
||||
func ReadQuery(ctx context.Context, model any, query string, args ...any) error {
|
||||
db := GetDBFromContext(ctx)
|
||||
return db.WithContext(ctx).Where(query, args...).Find(model).Error
|
||||
}
|
||||
|
||||
// WriteQuery 写查询(强制使用主库)
|
||||
func WriteQuery(ctx context.Context, model any, query string, args ...any) error {
|
||||
return DB.WithContext(ctx).Where(query, args...).Find(model).Error
|
||||
}
|
||||
|
||||
// HealthCheck 健康检查
|
||||
func HealthCheck() map[string]bool {
|
||||
result := make(map[string]bool)
|
||||
|
||||
// 检查主库
|
||||
if DB != nil {
|
||||
sqlDB, err := DB.DB()
|
||||
if err == nil && sqlDB.Ping() == nil {
|
||||
result["master"] = true
|
||||
} else {
|
||||
result["master"] = false
|
||||
}
|
||||
} else {
|
||||
result["master"] = false
|
||||
}
|
||||
|
||||
// 检查从库
|
||||
for i, replica := range replicas {
|
||||
if replica != nil {
|
||||
sqlDB, err := replica.DB()
|
||||
if err == nil && sqlDB.Ping() == nil {
|
||||
result[fmt.Sprintf("replica_%d", i+1)] = true
|
||||
} else {
|
||||
result[fmt.Sprintf("replica_%d", i+1)] = false
|
||||
}
|
||||
} else {
|
||||
result[fmt.Sprintf("replica_%d", i+1)] = false
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
+13
-3
@@ -36,10 +36,20 @@ func InitRedis(cfg *config.Config) error {
|
||||
|
||||
// CloseRedis 关闭 Redis 连接
|
||||
func CloseRedis() error {
|
||||
if RedisClient != nil {
|
||||
return RedisClient.Close()
|
||||
if RedisClient == nil {
|
||||
return nil
|
||||
}
|
||||
return nil
|
||||
err := RedisClient.Close()
|
||||
RedisClient = nil
|
||||
return err
|
||||
}
|
||||
|
||||
// HealthCheckRedis Redis 健康检查
|
||||
func HealthCheckRedis(ctx context.Context) error {
|
||||
if RedisClient == nil {
|
||||
return fmt.Errorf("Redis 未初始化")
|
||||
}
|
||||
return RedisClient.Ping(ctx).Err()
|
||||
}
|
||||
|
||||
// GetRedis 获取 Redis 客户端
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
package database_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/EthanCodeCraft/xlgo-core/database"
|
||||
)
|
||||
|
||||
func TestCloseRedisWithoutInit(t *testing.T) {
|
||||
if err := database.CloseRedis(); err != nil {
|
||||
t.Fatalf("CloseRedis without init should not error: %v", err)
|
||||
}
|
||||
if database.GetRedis() != nil {
|
||||
t.Fatal("expected Redis client nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHealthCheckRedisWithoutInit(t *testing.T) {
|
||||
if err := database.HealthCheckRedis(context.Background()); err == nil {
|
||||
t.Fatal("expected health check error without Redis init")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestIsTransientDBError(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
want bool
|
||||
}{
|
||||
{"nil", nil, true},
|
||||
{"access denied", errors.New("Error 1045: Access denied for user 'root'@'localhost'"), false},
|
||||
{"auth plugin", errors.New("authentication plugin 'caching_sha2_password' cannot be loaded"), false},
|
||||
{"unknown database", errors.New("Error 1049: Unknown database 'foo'"), false},
|
||||
{"invalid DSN", errors.New("invalid DSN: missing the slash separating the database name"), false},
|
||||
{"unknown driver", errors.New("sql: unknown driver \"foobar\" (forgotten import?)"), false},
|
||||
{"connection refused (transient)", errors.New("dial tcp 127.0.0.1:3306: connect: connection refused"), true},
|
||||
{"i/o timeout (transient)", errors.New("dial tcp 10.0.0.1:3306: i/o timeout"), true},
|
||||
{"empty msg", errors.New(""), true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := isTransientDBError(tt.err); got != tt.want {
|
||||
t.Errorf("isTransientDBError(%v) = %v, want %v", tt.err, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user