feat: 完善后端脚手架基础能力
This commit is contained in:
@@ -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...))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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) {
|
||||
|
||||
@@ -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...))
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 在目标平台编译")
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user