This commit is contained in:
oneao committed 2026-04-24 22:37:52 +08:00
1 parent 46cbb2f50d
commit e3d765a9d4
15 files changed
+903 -190

No files matched your search

+22 -14
View File
@@ -10,33 +10,41 @@ import (
var defaultClient *Client
type Executor interface {
Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error)
Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error)
QueryRow(ctx context.Context, sql string, args ...any) pgx.Row
}
type Client struct {
pool *pgxpool.Pool
tx pgx.Tx
}
// Init 初始化
// Init 初始化(只调用一次)
func Init(pool *pgxpool.Pool) {
defaultClient = &Client{pool: pool}
}
// New 创建会话
// New 获取全局 client
func New() *Client {
return &Client{
pool: defaultClient.pool,
if defaultClient == nil {
panic("db not initialized, call db.Init(pool) first")
}
return defaultClient
}
// 内部获取执行器(关键)
func (c *Client) exec() Executor {
func (c *Client) Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error) {
if c.tx != nil {
return c.tx
return c.tx.Exec(ctx, sql, args...)
}
return c.pool
return c.pool.Exec(ctx, sql, args...)
}
func (c *Client) Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) {
if c.tx != nil {
return c.tx.Query(ctx, sql, args...)
}
return c.pool.Query(ctx, sql, args...)
}
func (c *Client) QueryRow(ctx context.Context, sql string, args ...any) pgx.Row {
if c.tx != nil {
return c.tx.QueryRow(ctx, sql, args...)
}
return c.pool.QueryRow(ctx, sql, args...)
}
+8 -15
View File
@@ -7,32 +7,24 @@ import (
)
func (c *Client) Delete(ctx context.Context, table, keyField string, data map[string]any) error {
_, err := validateData(table, keyField, data)
keyVal, err := validateData(table, keyField, data)
if err != nil {
return err
}
where := make([]string, 0, len(data))
args := make([]any, 0, len(data))
i := 1
for k, v := range data {
where = append(where, fmt.Sprintf("%s = $%d", quoteCol(k), i))
args = append(args, v)
i++
}
sql := fmt.Sprintf(
"DELETE FROM %s WHERE %s",
"DELETE FROM %s WHERE %s = $1",
quoteTable(table),
strings.Join(where, " AND "),
quoteCol(keyField),
)
_, err = c.exec().Exec(ctx, sql, args...)
_, err = c.pool.Exec(ctx, sql, keyVal) // ⭐修复 exec
return err
}
func (c *Client) BatchDelete(ctx context.Context, table, keyField string, list []map[string]any) error {
if len(list) == 0 {
return fmt.Errorf("empty data")
}
@@ -44,6 +36,7 @@ func (c *Client) BatchDelete(ctx context.Context, table, keyField string, list [
)
for _, row := range list {
v, err := validateData(table, keyField, row)
if err != nil {
return err
@@ -61,6 +54,6 @@ func (c *Client) BatchDelete(ctx context.Context, table, keyField string, list [
strings.Join(in, ", "),
)
_, err := c.exec().Exec(ctx, sql, args...)
_, err := c.pool.Exec(ctx, sql, args...) // ⭐修复 exec
return err
}
+10 -4
View File
@@ -31,6 +31,7 @@ func buildInsertSQL(table string, data map[string]any) (string, []any) {
}
func (c *Client) Insert(ctx context.Context, table, keyField string, data map[string]any) error {
if _, err := validateData(table, keyField, data); err != nil {
return err
}
@@ -38,11 +39,13 @@ func (c *Client) Insert(ctx context.Context, table, keyField string, data map[st
data = applyMetaFields(ctx, table, data, true)
sql, args := buildInsertSQL(table, data)
_, err := c.exec().Exec(ctx, sql, args...)
_, err := c.pool.Exec(ctx, sql, args...)
return err
}
func (c *Client) BatchInsert(ctx context.Context, table, keyField string, list []map[string]any) error {
if len(list) == 0 {
return fmt.Errorf("empty data")
}
@@ -54,6 +57,7 @@ func (c *Client) BatchInsert(ctx context.Context, table, keyField string, list [
tableSQL := quoteTable(table)
// 固定字段顺序(稳定性关键)
var cols []string
for k := range first {
cols = append(cols, k)
@@ -67,9 +71,11 @@ func (c *Client) BatchInsert(ctx context.Context, table, keyField string, list [
)
for _, row := range list {
row = applyMetaFields(ctx, table, row, true)
var place []string
for _, col := range cols {
v, ok := row[col]
if !ok {
@@ -85,8 +91,8 @@ func (c *Client) BatchInsert(ctx context.Context, table, keyField string, list [
}
var quotedCols []string
for _, c := range cols {
quotedCols = append(quotedCols, quoteCol(c))
for _, col := range cols {
quotedCols = append(quotedCols, quoteCol(col))
}
sql := fmt.Sprintf(
@@ -96,6 +102,6 @@ func (c *Client) BatchInsert(ctx context.Context, table, keyField string, list [
strings.Join(values, ", "),
)
_, err := c.exec().Exec(ctx, sql, args...)
_, err := c.pool.Exec(ctx, sql, args...) // ⭐关键修复
return err
}
+13 -10
View File
@@ -34,7 +34,8 @@ func (c *Client) LoadData(
order = " ORDER BY " + orderBy
}
sql := fmt.Sprintf("SELECT %s FROM %s%s%s",
sql := fmt.Sprintf(
"SELECT %s FROM %s%s%s",
selectCols,
viewName,
where,
@@ -54,26 +55,25 @@ func (c *Client) LoadDataBySQL(
sql string,
args []any,
) ([]map[string]any, error) {
logger.FromCtx(ctx).Info("LoadDataBySQL",
zap.String("sql", sql),
zap.Any("args", args),
)
return c.query(ctx, sql, args...)
}
// ==========================
// 核心查询方法(已升级)
// ==========================
func (c *Client) query(ctx context.Context, sql string, args ...any) ([]map[string]any, error) {
sql = convertPlaceholder(sql)
rows, err := c.exec().Query(ctx, sql, args...)
rows, err := c.pool.Query(ctx, sql, args...)
if err != nil {
return nil, err
}
defer rows.Close()
// 获取字段名
fields := rows.FieldDescriptions()
result := make([]map[string]any, 0, 16)
@@ -86,19 +86,22 @@ func (c *Client) query(ctx context.Context, sql string, args ...any) ([]map[stri
row := make(map[string]any, len(values))
for i, f := range fields {
for i := range fields {
if i < len(values) {
row[string(f.Name)] = values[i]
row[string(fields[i].Name)] = values[i]
}
}
result = append(result, row)
}
return result, rows.Err()
if err := rows.Err(); err != nil {
return nil, err
}
return result, nil
}
// 高效 placeholder 转换
func convertPlaceholder(sql string) string {
var sb strings.Builder
sb.Grow(len(sql))
+14 -10
View File
@@ -2,12 +2,10 @@ package db
import (
"context"
"github.com/jackc/pgx/v5"
)
// WithTx 开启事务(在当前 client 上)
func (c *Client) WithTx(ctx context.Context, fn func(tx *Client) error) error {
tx, err := c.pool.Begin(ctx)
if err != nil {
return err
@@ -19,16 +17,22 @@ func (c *Client) WithTx(ctx context.Context, fn func(tx *Client) error) error {
tx: tx,
}
defer func(tx pgx.Tx, ctx context.Context) {
err := tx.Rollback(ctx)
if err != nil {
}
}(tx, ctx)
// rollback 兜底(防 panic / 提前 return)
defer func() {
_ = tx.Rollback(ctx)
}()
// 执行业务
if err := fn(txClient); err != nil {
_ = tx.Rollback(ctx)
return err
}
return tx.Commit(ctx)
// commit
if err := tx.Commit(ctx); err != nil {
_ = tx.Rollback(ctx)
return err
}
return nil
}
+12 -11
View File
@@ -8,6 +8,7 @@ import (
)
func (c *Client) Update(ctx context.Context, table, keyField string, data map[string]any) error {
keyVal, err := validateData(table, keyField, data)
if err != nil {
return err
@@ -15,9 +16,6 @@ func (c *Client) Update(ctx context.Context, table, keyField string, data map[st
data = applyMetaFields(ctx, table, data, false)
tableSQL := quoteTable(table)
keySQL := quoteCol(keyField)
var (
set []string
args []any
@@ -38,17 +36,18 @@ func (c *Client) Update(ctx context.Context, table, keyField string, data map[st
sql := fmt.Sprintf(
"UPDATE %s SET %s WHERE %s=$%d",
tableSQL,
quoteTable(table),
strings.Join(set, ", "),
keySQL,
quoteCol(keyField),
i,
)
_, err = c.exec().Exec(ctx, sql, args...)
_, err = c.pool.Exec(ctx, sql, args...)
return err
}
func (c *Client) BatchUpdate(ctx context.Context, table, keyField string, list []map[string]any) error {
if len(list) == 0 {
return fmt.Errorf("empty data")
}
@@ -75,6 +74,7 @@ func (c *Client) BatchUpdate(ctx context.Context, table, keyField string, list [
sets []string
)
// CASE 构建
for _, col := range cols {
colSQL := quoteCol(col)
@@ -84,7 +84,7 @@ func (c *Client) BatchUpdate(ctx context.Context, table, keyField string, list [
for _, row := range list {
row = applyMetaFields(ctx, table, row, false)
keyVal, _ := row[keyField]
keyVal := row[keyField]
val := row[col]
caseSQL.WriteString(fmt.Sprintf(
@@ -101,9 +101,10 @@ func (c *Client) BatchUpdate(ctx context.Context, table, keyField string, list [
sets = append(sets, caseSQL.String())
}
var where []string
// ⭐修复 IN 写法(关键)
var inPlaceholders []string
for _, row := range list {
where = append(where, fmt.Sprintf("$%d", argIndex))
inPlaceholders = append(inPlaceholders, fmt.Sprintf("$%d", argIndex))
args = append(args, row[keyField])
argIndex++
}
@@ -113,9 +114,9 @@ func (c *Client) BatchUpdate(ctx context.Context, table, keyField string, list [
tableSQL,
strings.Join(sets, ", "),
keySQL,
strings.Join(where, ", "),
strings.Join(inPlaceholders, ", "),
)
_, err := c.exec().Exec(ctx, sql, args...)
_, err := c.pool.Exec(ctx, sql, args...) // ⭐修复点
return err
}