package db import ( "cn/oneao/base-go/conf" "cn/oneao/base-go/db/repo" "context" "fmt" "time" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) var ( DB *pgxpool.Pool Queries *repo.Queries ) type SQLCBatch interface { Exec(func(int, error)) } func InitDB() { cfg := conf.GetConf().Pgsql sslMode := "disable" if cfg.SllMode { sslMode = "require" } dsn := fmt.Sprintf( "postgres://%s:%s@%s:%d/%s?sslmode=%s&TimeZone=%s", cfg.User, cfg.Password, cfg.Host, cfg.Port, cfg.Dbname, sslMode, cfg.TimeZone, ) ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() // 使用可修改配置 cfgPool, err := pgxpool.ParseConfig(dsn) if err != nil { panic(err) } // 给连接池参数设置默认值 if cfg.MaxOpenConns <= 0 { cfg.MaxOpenConns = 10 } if cfg.MaxIdleConns <= 0 { cfg.MaxIdleConns = 5 } if cfg.ConnMaxLifetime <= 0 { cfg.ConnMaxLifetime = 30 * time.Minute } cfgPool.MaxConns = cfg.MaxOpenConns cfgPool.MinConns = cfg.MaxIdleConns cfgPool.MaxConnLifetime = cfg.ConnMaxLifetime pool, err := pgxpool.NewWithConfig(ctx, cfgPool) if err != nil { panic(err) } // 测试连接 if err := pool.Ping(ctx); err != nil { panic(err) } DB = pool Queries = repo.New(DB) } // WithTx 执行事务,支持 panic 和 error 自动回滚 // 默认使用全局 DB,调用更简洁 func WithTx(ctx context.Context, fn func(q *repo.Queries) error) (err error) { return WithTxPool(ctx, DB, fn) } // WithTxPool 支持自定义连接池 func WithTxPool(ctx context.Context, pool *pgxpool.Pool, fn func(q *repo.Queries) error) (err error) { tx, err := pool.BeginTx(ctx, pgx.TxOptions{}) if err != nil { return err } defer func() { if p := recover(); p != nil { // panic 时回滚事务 if rbErr := tx.Rollback(ctx); rbErr != nil { fmt.Printf("rollback failed during panic: %v\n", rbErr) } panic(p) } else if err != nil { // 回滚事务,并捕获 rollback 错误 if rbErr := tx.Rollback(ctx); rbErr != nil { err = fmt.Errorf("rollback failed: %v, original error: %w", rbErr, err) } } else { // 提交事务,并捕获 commit 错误 if commitErr := tx.Commit(ctx); commitErr != nil { err = fmt.Errorf("commit failed: %w", commitErr) } } }() q := Queries.WithTx(tx) err = fn(q) return err } // RunBatch 通用批处理执行器(适配所有 sqlc Batch) func RunBatch(ctx context.Context, batch SQLCBatch) error { var firstErr error batch.Exec(func(i int, err error) { if err != nil && firstErr == nil { firstErr = err } }) return firstErr }