u
This commit is contained in:
1 parent
835958886e
commit
d04d22723b
155 files changed
+543
-9793
No files matched your search
@@ -1,123 +0,0 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"allapp/conf"
|
||||
"allapp/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
|
||||
}
|
||||
Reference in new issue
Block a user