71 lines
1.5 KiB
Go
71 lines
1.5 KiB
Go
package db
|
|
|
|
import (
|
|
"allapp/db/repo"
|
|
"context"
|
|
"fmt"
|
|
|
|
"github.com/jackc/pgx/v5"
|
|
"github.com/jackc/pgx/v5/pgxpool"
|
|
)
|
|
|
|
var (
|
|
Pool *pgxpool.Pool
|
|
Queries *repo.Queries
|
|
)
|
|
|
|
type SQLCBatch interface {
|
|
Exec(func(int, error))
|
|
}
|
|
|
|
// WithTx 执行事务,支持 panic 和 error 自动回滚
|
|
// 默认使用全局 DB,调用更简洁
|
|
func WithTx(ctx context.Context, fn func(q *repo.Queries) error) (err error) {
|
|
return WithTxPool(ctx, Pool, 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
|
|
}
|