20260507213409
This commit is contained in:
1 parent
3070eac71b
commit
4c7f516a03
204 files changed
+3623
No files matched your search
@@ -0,0 +1,50 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
var defaultClient *Client
|
||||
|
||||
type Client struct {
|
||||
pool *pgxpool.Pool
|
||||
tx pgx.Tx
|
||||
}
|
||||
|
||||
// Init 初始化(只调用一次)
|
||||
func Init(pool *pgxpool.Pool) {
|
||||
defaultClient = &Client{pool: pool}
|
||||
}
|
||||
|
||||
// New 获取全局 client
|
||||
func New() *Client {
|
||||
if defaultClient == nil {
|
||||
panic("db not initialized, call db.Init(pool) first")
|
||||
}
|
||||
return defaultClient
|
||||
}
|
||||
|
||||
func (c *Client) Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error) {
|
||||
if c.tx != nil {
|
||||
return c.tx.Exec(ctx, sql, args...)
|
||||
}
|
||||
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...)
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func (c *Client) Delete(ctx context.Context, table, keyField string, data map[string]any) error {
|
||||
|
||||
keyVal, err := validateData(table, keyField, data)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"DELETE FROM %s WHERE %s = $1",
|
||||
quoteTable(table),
|
||||
quoteCol(keyField),
|
||||
)
|
||||
|
||||
_, 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")
|
||||
}
|
||||
|
||||
var (
|
||||
args []any
|
||||
in []string
|
||||
i = 1
|
||||
)
|
||||
|
||||
for _, row := range list {
|
||||
|
||||
v, err := validateData(table, keyField, row)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
in = append(in, fmt.Sprintf("$%d", i))
|
||||
args = append(args, v)
|
||||
i++
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"DELETE FROM %s WHERE %s IN (%s)",
|
||||
quoteTable(table),
|
||||
quoteCol(keyField),
|
||||
strings.Join(in, ", "),
|
||||
)
|
||||
|
||||
_, err := c.pool.Exec(ctx, sql, args...) // ⭐修复 exec
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func buildInsertSQL(table string, data map[string]any) (string, []any) {
|
||||
var (
|
||||
cols []string
|
||||
vals []string
|
||||
args []any
|
||||
i = 1
|
||||
)
|
||||
|
||||
for k, v := range data {
|
||||
cols = append(cols, quoteCol(k))
|
||||
vals = append(vals, fmt.Sprintf("$%d", i))
|
||||
args = append(args, v)
|
||||
i++
|
||||
}
|
||||
|
||||
return fmt.Sprintf(
|
||||
"INSERT INTO %s (%s) VALUES (%s)",
|
||||
quoteTable(table),
|
||||
strings.Join(cols, ", "),
|
||||
strings.Join(vals, ", "),
|
||||
), args
|
||||
}
|
||||
|
||||
func (c *Client) Insert(ctx context.Context, table, keyField string, data map[string]any) error {
|
||||
if _, err := validateInsertData(table, keyField, data); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
ensureID(data, keyField)
|
||||
|
||||
data = applyMetaFields(ctx, table, data, true)
|
||||
|
||||
sql, args := buildInsertSQL(table, data)
|
||||
|
||||
_, 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")
|
||||
}
|
||||
|
||||
first := list[0]
|
||||
if _, err := validateInsertData(table, keyField, first); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
tableSQL := quoteTable(table)
|
||||
|
||||
// 固定字段顺序(稳定性关键)
|
||||
var cols []string
|
||||
for k := range first {
|
||||
cols = append(cols, k)
|
||||
}
|
||||
sort.Strings(cols)
|
||||
|
||||
var (
|
||||
args []any
|
||||
values []string
|
||||
argIndex = 1
|
||||
)
|
||||
|
||||
for _, row := range list {
|
||||
ensureID(row, keyField)
|
||||
|
||||
row = applyMetaFields(ctx, table, row, true)
|
||||
|
||||
var place []string
|
||||
|
||||
for _, col := range cols {
|
||||
v, ok := row[col]
|
||||
if !ok {
|
||||
return fmt.Errorf("missing field: %s", col)
|
||||
}
|
||||
|
||||
place = append(place, fmt.Sprintf("$%d", argIndex))
|
||||
args = append(args, v)
|
||||
argIndex++
|
||||
}
|
||||
|
||||
values = append(values, fmt.Sprintf("(%s)", strings.Join(place, ",")))
|
||||
}
|
||||
|
||||
var quotedCols []string
|
||||
for _, col := range cols {
|
||||
quotedCols = append(quotedCols, quoteCol(col))
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"INSERT INTO %s (%s) VALUES %s",
|
||||
tableSQL,
|
||||
strings.Join(quotedCols, ", "),
|
||||
strings.Join(values, ", "),
|
||||
)
|
||||
|
||||
_, err := c.pool.Exec(ctx, sql, args...) // ⭐关键修复
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"allapp-go/internal/middleware"
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
var auditExcludeTables = map[string]bool{
|
||||
"b_user": true,
|
||||
"b_user_oauth": true,
|
||||
}
|
||||
|
||||
// 返回新 map(避免修改入参)
|
||||
func applyMetaFields(ctx context.Context, table string, data map[string]any, isInsert bool) map[string]any {
|
||||
if auditExcludeTables[table] {
|
||||
return data
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
newData := make(map[string]any, len(data)+2)
|
||||
|
||||
for k, v := range data {
|
||||
newData[k] = v
|
||||
}
|
||||
|
||||
if isInsert {
|
||||
newData["create_time"] = now
|
||||
newData["update_time"] = now
|
||||
} else {
|
||||
newData["update_time"] = now
|
||||
}
|
||||
|
||||
if userID, ok := middleware.GetUserID(ctx); ok {
|
||||
if isInsert {
|
||||
newData["create_by"] = userID
|
||||
}
|
||||
newData["update_by"] = userID
|
||||
}
|
||||
|
||||
return newData
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"allapp-go/pkg/logger"
|
||||
"context"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
func (c *Client) LoadData(
|
||||
ctx context.Context,
|
||||
viewName string,
|
||||
searchCondition string,
|
||||
orderBy string,
|
||||
searchColumns []string,
|
||||
args []any,
|
||||
) ([]map[string]any, error) {
|
||||
|
||||
selectCols := "*"
|
||||
if len(searchColumns) > 0 {
|
||||
selectCols = strings.Join(searchColumns, ", ")
|
||||
}
|
||||
|
||||
where := ""
|
||||
if searchCondition != "" {
|
||||
where = " WHERE " + searchCondition
|
||||
}
|
||||
|
||||
order := ""
|
||||
if orderBy != "" {
|
||||
order = " ORDER BY " + orderBy
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"SELECT %s FROM %s%s%s",
|
||||
selectCols,
|
||||
viewName,
|
||||
where,
|
||||
order,
|
||||
)
|
||||
|
||||
logger.FromCtx(ctx).Info("LoadData",
|
||||
zap.String("sql", sql),
|
||||
zap.Any("args", args),
|
||||
)
|
||||
|
||||
return c.query(ctx, sql, args...)
|
||||
}
|
||||
|
||||
func (c *Client) LoadDataBySQL(
|
||||
ctx context.Context,
|
||||
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.pool.Query(ctx, sql, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
fields := rows.FieldDescriptions()
|
||||
|
||||
result := make([]map[string]any, 0, 16)
|
||||
|
||||
for rows.Next() {
|
||||
values, err := rows.Values()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
row := make(map[string]any, len(values))
|
||||
|
||||
for i := range fields {
|
||||
if i < len(values) {
|
||||
row[string(fields[i].Name)] = values[i]
|
||||
}
|
||||
}
|
||||
|
||||
result = append(result, row)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func convertPlaceholder(sql string) string {
|
||||
var sb strings.Builder
|
||||
sb.Grow(len(sql))
|
||||
|
||||
argIndex := 1
|
||||
|
||||
for i := 0; i < len(sql); i++ {
|
||||
if sql[i] == '?' {
|
||||
sb.WriteByte('$')
|
||||
sb.WriteString(strconv.Itoa(argIndex))
|
||||
argIndex++
|
||||
} else {
|
||||
sb.WriteByte(sql[i])
|
||||
}
|
||||
}
|
||||
|
||||
return sb.String()
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
)
|
||||
|
||||
func (c *Client) WithTx(ctx context.Context, fn func(tx *Client) error) error {
|
||||
|
||||
tx, err := c.pool.Begin(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 创建事务 client
|
||||
txClient := &Client{
|
||||
pool: c.pool,
|
||||
tx: tx,
|
||||
}
|
||||
|
||||
// rollback 兜底(防 panic / 提前 return)
|
||||
defer func() {
|
||||
_ = tx.Rollback(ctx)
|
||||
}()
|
||||
|
||||
// 执行业务
|
||||
if err := fn(txClient); err != nil {
|
||||
_ = tx.Rollback(ctx)
|
||||
return err
|
||||
}
|
||||
|
||||
// commit
|
||||
if err := tx.Commit(ctx); err != nil {
|
||||
_ = tx.Rollback(ctx)
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
data = applyMetaFields(ctx, table, data, false)
|
||||
|
||||
var (
|
||||
set []string
|
||||
args []any
|
||||
i = 1
|
||||
)
|
||||
|
||||
for k, v := range data {
|
||||
if k == keyField {
|
||||
continue
|
||||
}
|
||||
|
||||
set = append(set, fmt.Sprintf("%s=$%d", quoteCol(k), i))
|
||||
args = append(args, v)
|
||||
i++
|
||||
}
|
||||
|
||||
args = append(args, keyVal)
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"UPDATE %s SET %s WHERE %s=$%d",
|
||||
quoteTable(table),
|
||||
strings.Join(set, ", "),
|
||||
quoteCol(keyField),
|
||||
i,
|
||||
)
|
||||
|
||||
_, 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")
|
||||
}
|
||||
|
||||
first := list[0]
|
||||
if _, err := validateData(table, keyField, first); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
tableSQL := quoteTable(table)
|
||||
keySQL := quoteCol(keyField)
|
||||
|
||||
var cols []string
|
||||
for k := range first {
|
||||
if k != keyField {
|
||||
cols = append(cols, k)
|
||||
}
|
||||
}
|
||||
sort.Strings(cols)
|
||||
|
||||
var (
|
||||
args []any
|
||||
argIndex = 1
|
||||
sets []string
|
||||
)
|
||||
|
||||
// CASE 构建
|
||||
for _, col := range cols {
|
||||
colSQL := quoteCol(col)
|
||||
|
||||
var caseSQL strings.Builder
|
||||
caseSQL.WriteString(fmt.Sprintf("%s = CASE %s ", colSQL, keySQL))
|
||||
|
||||
for _, row := range list {
|
||||
row = applyMetaFields(ctx, table, row, false)
|
||||
|
||||
keyVal := row[keyField]
|
||||
val := row[col]
|
||||
|
||||
caseSQL.WriteString(fmt.Sprintf(
|
||||
"WHEN $%d THEN $%d ",
|
||||
argIndex,
|
||||
argIndex+1,
|
||||
))
|
||||
|
||||
args = append(args, keyVal, val)
|
||||
argIndex += 2
|
||||
}
|
||||
|
||||
caseSQL.WriteString("END")
|
||||
sets = append(sets, caseSQL.String())
|
||||
}
|
||||
|
||||
// ⭐修复 IN 写法(关键)
|
||||
var inPlaceholders []string
|
||||
for _, row := range list {
|
||||
inPlaceholders = append(inPlaceholders, fmt.Sprintf("$%d", argIndex))
|
||||
args = append(args, row[keyField])
|
||||
argIndex++
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"UPDATE %s SET %s WHERE %s IN (%s)",
|
||||
tableSQL,
|
||||
strings.Join(sets, ", "),
|
||||
keySQL,
|
||||
strings.Join(inPlaceholders, ", "),
|
||||
)
|
||||
|
||||
_, err := c.pool.Exec(ctx, sql, args...) // ⭐修复点
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"allapp-go/pkg/uniqueid"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var identRe = regexp.MustCompile(`^[a-zA-Z0-9_]+$`)
|
||||
|
||||
// 安全标识符校验
|
||||
func safeIdent(name string) bool {
|
||||
return identRe.MatchString(name)
|
||||
}
|
||||
|
||||
// 表名引用(支持 schema)
|
||||
func quoteTable(table string) string {
|
||||
parts := strings.Split(table, ".")
|
||||
for i, p := range parts {
|
||||
if !safeIdent(p) {
|
||||
panic(fmt.Sprintf("invalid table: %s", p))
|
||||
}
|
||||
parts[i] = `"` + p + `"`
|
||||
}
|
||||
return strings.Join(parts, ".")
|
||||
}
|
||||
|
||||
// 字段引用
|
||||
func quoteCol(col string) string {
|
||||
if !safeIdent(col) {
|
||||
panic(fmt.Sprintf("invalid column: %s", col))
|
||||
}
|
||||
return `"` + col + `"`
|
||||
}
|
||||
|
||||
// 校验数据
|
||||
func validateData(table string, keyField string, data map[string]any) (any, error) {
|
||||
if table == "" {
|
||||
return nil, fmt.Errorf("表名不能为空")
|
||||
}
|
||||
|
||||
if len(data) == 0 {
|
||||
return nil, fmt.Errorf("数据不能为空")
|
||||
}
|
||||
|
||||
if keyField == "" {
|
||||
return nil, fmt.Errorf("主键字段不能为空")
|
||||
}
|
||||
|
||||
val, ok := data[keyField]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("缺少主键字段: %s", keyField)
|
||||
}
|
||||
|
||||
if val == nil {
|
||||
return nil, fmt.Errorf("主键字段值不能为空")
|
||||
}
|
||||
|
||||
return val, nil
|
||||
}
|
||||
|
||||
// 校验数据
|
||||
func validateInsertData(table string, keyField string, data map[string]any) (any, error) {
|
||||
if table == "" {
|
||||
return nil, fmt.Errorf("表名不能为空")
|
||||
}
|
||||
|
||||
if len(data) == 0 {
|
||||
return nil, fmt.Errorf("数据不能为空")
|
||||
}
|
||||
|
||||
if keyField == "" {
|
||||
return nil, fmt.Errorf("主键字段不能为空")
|
||||
}
|
||||
|
||||
val, _ := data[keyField]
|
||||
|
||||
return val, nil
|
||||
}
|
||||
|
||||
func ensureID(data map[string]any, keyField string) {
|
||||
v, ok := data[keyField]
|
||||
|
||||
if !ok || v == nil || v == "" {
|
||||
data[keyField] = uniqueid.NextId()
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user