feat: 完善后端脚手架基础能力

This commit is contained in:
2026-08-09 01:06:21 +08:00
parent e311f416f2
commit f95311a107
61 changed files with 5630 additions and 521 deletions
+169
View File
@@ -0,0 +1,169 @@
package database
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"time"
"skeleton/config"
"go.uber.org/zap"
"gorm.io/driver/mysql"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
var DB *gorm.DB
// Init 根据配置的 driver 初始化数据库连接。
func Init(cfg *config.DatabaseConfig, log *zap.Logger) error {
if err := prepareSQLiteDirectory(cfg); err != nil {
return err
}
dialector, err := Dialector(cfg)
if err != nil {
return err
}
gormLogger := logger.New(
&GormZapWriter{Logger: log},
logger.Config{
SlowThreshold: time.Second,
LogLevel: logger.Info,
Colorful: false,
},
)
db, err := gorm.Open(dialector, &gorm.Config{Logger: gormLogger})
if err != nil {
return fmt.Errorf("连接 %s 数据库失败: %w", cfg.Driver, err)
}
sqlDB, err := db.DB()
if err != nil {
return fmt.Errorf("获取底层数据库连接失败: %w", err)
}
sqlDB.SetMaxIdleConns(cfg.MaxIdleConns)
sqlDB.SetMaxOpenConns(cfg.MaxOpenConns)
sqlDB.SetConnMaxLifetime(time.Duration(cfg.ConnMaxLifetime) * time.Minute)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := sqlDB.PingContext(ctx); err != nil {
_ = sqlDB.Close()
return fmt.Errorf("%s 数据库连接测试失败: %w", cfg.Driver, err)
}
DB = db
log.Info("数据库连接初始化成功",
zap.String("driver", cfg.Driver),
zap.String("database", databaseName(cfg)),
zap.Int("max_idle_conns", cfg.MaxIdleConns),
zap.Int("max_open_conns", cfg.MaxOpenConns),
)
return nil
}
func prepareSQLiteDirectory(cfg *config.DatabaseConfig) error {
driver := strings.ToLower(strings.TrimSpace(cfg.Driver))
if driver != "sqlite" && driver != "sqlite3" {
return nil
}
path := BuildDSN(cfg)
if path == "" || path == ":memory:" || strings.HasPrefix(path, "file:") {
return nil
}
dir := filepath.Dir(path)
if dir == "." {
return nil
}
if err := os.MkdirAll(dir, 0755); err != nil {
return fmt.Errorf("创建 SQLite 数据目录失败: %w", err)
}
return nil
}
// Dialector 创建对应数据库的 GORM 方言实例,便于应用和测试复用。
func Dialector(cfg *config.DatabaseConfig) (gorm.Dialector, error) {
driver := strings.ToLower(strings.TrimSpace(cfg.Driver))
dsn := BuildDSN(cfg)
switch driver {
case "postgres", "postgresql":
return postgres.Open(dsn), nil
case "mysql":
return mysql.Open(dsn), nil
case "sqlite", "sqlite3":
return sqliteDialector(dsn)
default:
return nil, fmt.Errorf("不支持的数据库驱动 %q,可选值: postgres, mysql, sqlite", cfg.Driver)
}
}
// BuildDSN 返回显式 DSN,或根据结构化配置构建对应方言的连接串。
func BuildDSN(cfg *config.DatabaseConfig) string {
if cfg.DSN != "" {
return cfg.DSN
}
switch strings.ToLower(strings.TrimSpace(cfg.Driver)) {
case "mysql":
return fmt.Sprintf("%s:%s@tcp(%s:%d)/%s?charset=utf8mb4&parseTime=True&loc=Local",
cfg.Username, cfg.Password, cfg.Host, cfg.Port, cfg.DBName)
case "sqlite", "sqlite3":
return cfg.SQLitePath
default:
return fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s sslmode=%s",
cfg.Host, cfg.Port, cfg.Username, cfg.Password, cfg.DBName, cfg.SSLMode)
}
}
func databaseName(cfg *config.DatabaseConfig) string {
if strings.HasPrefix(strings.ToLower(cfg.Driver), "sqlite") {
return BuildDSN(cfg)
}
return cfg.DBName
}
func Close() error {
if DB == nil {
return nil
}
sqlDB, err := DB.DB()
if err != nil {
return err
}
return sqlDB.Close()
}
func GetDB() *gorm.DB {
return DB
}
// Ping 检查关系数据库连接是否可用。
func Ping(ctx context.Context) error {
if DB == nil {
return fmt.Errorf("数据库尚未初始化")
}
sqlDB, err := DB.DB()
if err != nil {
return err
}
return sqlDB.PingContext(ctx)
}
type GormZapWriter struct {
Logger *zap.Logger
}
func (g *GormZapWriter) Printf(format string, args ...interface{}) {
g.Logger.Info(fmt.Sprintf(format, args...))
}
+54
View File
@@ -0,0 +1,54 @@
package database
import (
"strings"
"testing"
"skeleton/config"
)
func TestBuildDSN(t *testing.T) {
tests := []struct {
name string
cfg config.DatabaseConfig
want string
}{
{
name: "explicit DSN wins",
cfg: config.DatabaseConfig{Driver: "mysql", DSN: "custom-dsn"},
want: "custom-dsn",
},
{
name: "postgres",
cfg: config.DatabaseConfig{Driver: "postgres", Host: "db", Port: 5432,
Username: "app", Password: "secret", DBName: "demo", SSLMode: "require"},
want: "host=db port=5432 user=app password=secret dbname=demo sslmode=require",
},
{
name: "mysql",
cfg: config.DatabaseConfig{Driver: "mysql", Host: "db", Port: 3306,
Username: "app", Password: "secret", DBName: "demo"},
want: "app:secret@tcp(db:3306)/demo?charset=utf8mb4&parseTime=True&loc=Local",
},
{
name: "sqlite",
cfg: config.DatabaseConfig{Driver: "sqlite", SQLitePath: "data/demo.db"},
want: "data/demo.db",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := BuildDSN(&tt.cfg); got != tt.want {
t.Fatalf("BuildDSN() = %q, want %q", got, tt.want)
}
})
}
}
func TestDialectorRejectsUnknownDriver(t *testing.T) {
_, err := Dialector(&config.DatabaseConfig{Driver: "oracle"})
if err == nil || !strings.Contains(err.Error(), "不支持的数据库驱动") {
t.Fatalf("Dialector() error = %v", err)
}
}
+55
View File
@@ -0,0 +1,55 @@
//go:build integration
package database
import (
"fmt"
"os"
"testing"
"time"
"skeleton/config"
"skeleton/models"
"go.uber.org/zap"
)
func TestPostgresAndMySQL(t *testing.T) {
tests := []struct {
driver string
env string
}{
{driver: "postgres", env: "TEST_POSTGRES_DSN"},
{driver: "mysql", env: "TEST_MYSQL_DSN"},
}
for _, tt := range tests {
t.Run(tt.driver, func(t *testing.T) {
dsn := os.Getenv(tt.env)
if dsn == "" {
t.Skipf("%s 未设置", tt.env)
}
cfg := config.DatabaseConfig{
Driver: tt.driver, DSN: dsn,
MaxIdleConns: 1, MaxOpenConns: 2, ConnMaxLifetime: 1,
}
if err := Init(&cfg, zap.NewNop()); err != nil {
t.Fatalf("Init() error = %v", err)
}
t.Cleanup(func() { _ = Close(); DB = nil })
if err := DB.AutoMigrate(&models.User{}); err != nil {
t.Fatalf("AutoMigrate() error = %v", err)
}
user := models.User{
Username: fmt.Sprintf("integration_%s_%d", tt.driver, time.Now().UnixNano()),
Password: "test-hash",
}
if err := DB.Create(&user).Error; err != nil {
t.Fatalf("Create() error = %v", err)
}
if err := DB.Delete(&user).Error; err != nil {
t.Fatalf("Delete() error = %v", err)
}
})
}
}
+15 -5
View File
@@ -12,7 +12,8 @@ import (
// MigrationConfig 迁移配置
type MigrationConfig struct {
Environment string // local, production
Environment string // 例如 local_postgres、production_mysql
Directory string // 按数据库方言隔离的迁移目录
Timeout int // 超时时间(秒)
}
@@ -77,11 +78,15 @@ func (m *MigrationManager) GenerateMigration(name string) error {
}
// ApplyMigrations 应用迁移
func (m *MigrationManager) ApplyMigrations() error {
func (m *MigrationManager) ApplyMigrations(dryRun bool) error {
ctx, cancel := context.WithTimeout(context.Background(), time.Duration(m.config.Timeout)*time.Second)
defer cancel()
cmd := exec.CommandContext(ctx, "atlas", "migrate", "apply", "--env", m.config.Environment)
args := []string{"migrate", "apply", "--env", m.config.Environment}
if dryRun {
args = append(args, "--dry-run")
}
cmd := exec.CommandContext(ctx, "atlas", args...)
output, err := cmd.CombinedOutput()
if err != nil {
@@ -91,7 +96,9 @@ func (m *MigrationManager) ApplyMigrations() error {
return fmt.Errorf("应用迁移失败: %w", err)
}
m.logger.Info("迁移应用成功", zap.String("output", string(output)))
m.logger.Info("迁移应用成功",
zap.Bool("dry_run", dryRun),
zap.String("output", string(output)))
return nil
}
@@ -126,7 +133,10 @@ func EnsureAtlasInstalled() error {
// InitMigrationDirectory 初始化迁移目录
func (m *MigrationManager) InitMigrationDirectory() error {
migrationDir := "migrations"
migrationDir := m.config.Directory
if migrationDir == "" {
migrationDir = "migrations"
}
// 检查目录是否存在
if _, err := os.Stat(migrationDir); os.IsNotExist(err) {
-98
View File
@@ -1,98 +0,0 @@
package database
import (
"context"
"fmt"
"time"
"skeleton/config"
"go.uber.org/zap"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
var DB *gorm.DB
// Init 初始化数据库连接
func Init(cfg *config.DatabaseConfig, log *zap.Logger) error {
// 构建DSN
dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s sslmode=%s",
cfg.Host, cfg.Port, cfg.Username, cfg.Password, cfg.DBName, cfg.SSLMode)
// 配置GORM日志
gormLogger := logger.New(
&GormZapWriter{Logger: log},
logger.Config{
SlowThreshold: time.Second,
LogLevel: logger.Info,
Colorful: false,
},
)
// 打开数据库连接
db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{
Logger: gormLogger,
})
if err != nil {
return fmt.Errorf("连接数据库失败: %w", err)
}
// 获取底层的sql.DB以配置连接池
sqlDB, err := db.DB()
if err != nil {
return fmt.Errorf("获取底层数据库连接失败: %w", err)
}
// 配置连接池
sqlDB.SetMaxIdleConns(cfg.MaxIdleConns)
sqlDB.SetMaxOpenConns(cfg.MaxOpenConns)
sqlDB.SetConnMaxLifetime(time.Duration(cfg.ConnMaxLifetime) * time.Minute)
// 测试连接
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := sqlDB.PingContext(ctx); err != nil {
return fmt.Errorf("数据库连接测试失败: %w", err)
}
DB = db
log.Info("数据库连接初始化成功",
zap.String("host", cfg.Host),
zap.Int("port", cfg.Port),
zap.String("database", cfg.DBName),
zap.Int("max_idle_conns", cfg.MaxIdleConns),
zap.Int("max_open_conns", cfg.MaxOpenConns),
)
return nil
}
// Close 关闭数据库连接
func Close() error {
if DB != nil {
sqlDB, err := DB.DB()
if err != nil {
return err
}
return sqlDB.Close()
}
return nil
}
// GetDB 获取数据库实例
func GetDB() *gorm.DB {
return DB
}
// GormZapWriter GORM的Zap日志写入器
type GormZapWriter struct {
Logger *zap.Logger
}
func (g *GormZapWriter) Printf(format string, args ...interface{}) {
g.Logger.Info(fmt.Sprintf(format, args...))
}
+17
View File
@@ -21,6 +21,11 @@ var redisClient *RedisClient
// InitRedis 初始化Redis客户端
func InitRedis(cfg *config.Config, logger *zap.Logger) error {
if !cfg.Redis.Enabled {
redisClient = nil
logger.Info("Redis 已禁用")
return nil
}
// 创建Redis客户端配置
rdb := redis.NewClient(&redis.Options{
Addr: fmt.Sprintf("%s:%d", cfg.Redis.Host, cfg.Redis.Port),
@@ -55,6 +60,18 @@ func InitRedis(cfg *config.Config, logger *zap.Logger) error {
return nil
}
// Ping 检查 Redis 是否可用。
func (r *RedisClient) Ping(ctx context.Context) error {
if r == nil || r.client == nil {
return nil
}
return r.client.Ping(ctx).Err()
}
func RedisEnabled() bool {
return redisClient != nil
}
// GetRedisClient 获取Redis客户端实例
func GetRedisClient() *RedisClient {
return redisClient
+12
View File
@@ -0,0 +1,12 @@
//go:build cgo
package database
import (
"gorm.io/driver/sqlite"
"gorm.io/gorm"
)
func sqliteDialector(dsn string) (gorm.Dialector, error) {
return sqlite.Open(dsn), nil
}
+13
View File
@@ -0,0 +1,13 @@
//go:build !cgo
package database
import (
"fmt"
"gorm.io/gorm"
)
func sqliteDialector(string) (gorm.Dialector, error) {
return nil, fmt.Errorf("SQLite 驱动需要 CGO;请使用 CGO_ENABLED=1 在目标平台编译")
}
+30
View File
@@ -0,0 +1,30 @@
//go:build cgo
package database
import (
"testing"
"skeleton/config"
"go.uber.org/zap"
)
func TestInitSQLiteInMemory(t *testing.T) {
cfg := config.DatabaseConfig{
Driver: "sqlite",
DSN: ":memory:",
MaxIdleConns: 1,
MaxOpenConns: 1,
}
if err := Init(&cfg, zap.NewNop()); err != nil {
t.Fatalf("Init() error = %v", err)
}
t.Cleanup(func() {
_ = Close()
DB = nil
})
if GetDB() == nil {
t.Fatal("GetDB() returned nil")
}
}
+41
View File
@@ -0,0 +1,41 @@
package database
import (
"context"
"fmt"
"sync"
"time"
)
var revokedTokens sync.Map
func RevokeToken(ctx context.Context, tokenID string, ttl time.Duration) error {
if tokenID == "" || ttl <= 0 {
return nil
}
if redisClient != nil {
return redisClient.client.Set(ctx, "jwt:revoked:"+tokenID, "1", ttl).Err()
}
revokedTokens.Store(tokenID, time.Now().Add(ttl))
return nil
}
func IsTokenRevoked(ctx context.Context, tokenID string) (bool, error) {
if tokenID == "" {
return false, fmt.Errorf("token 缺少 jti")
}
if redisClient != nil {
count, err := redisClient.client.Exists(ctx, "jwt:revoked:"+tokenID).Result()
return count > 0, err
}
value, ok := revokedTokens.Load(tokenID)
if !ok {
return false, nil
}
expiresAt := value.(time.Time)
if time.Now().After(expiresAt) {
revokedTokens.Delete(tokenID)
return false, nil
}
return true, nil
}