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:
杭州明婳科技
2026-06-22 21:40:10 +08:00
parent 2cc8c70960
commit dcfd24b624
61 changed files with 5604 additions and 1299 deletions
+139
View File
@@ -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() },
})
}
+473
View File
@@ -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
}
+216
View File
@@ -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)
}
}
}
-315
View File
@@ -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
View File
@@ -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 客户端
+23
View File
@@ -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")
}
}
+31
View File
@@ -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)
}
})
}
}