u
This commit is contained in:
1 parent
89aaef1d6b
commit
6d54a9e407
192 files changed
+971
-51708
No files matched your search
@@ -1,52 +0,0 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
// 高效 placeholder 转换
|
||||
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()
|
||||
}
|
||||
@@ -1,42 +0,0 @@
|
||||
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 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 初始化
|
||||
func Init(pool *pgxpool.Pool) {
|
||||
defaultClient = &Client{pool: pool}
|
||||
}
|
||||
|
||||
// New 创建会话
|
||||
func New() *Client {
|
||||
return &Client{
|
||||
pool: defaultClient.pool,
|
||||
}
|
||||
}
|
||||
|
||||
// 内部获取执行器(关键)
|
||||
func (c *Client) exec() Executor {
|
||||
if c.tx != nil {
|
||||
return c.tx
|
||||
}
|
||||
return c.pool
|
||||
}
|
||||
@@ -1,83 +0,0 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func (c *Client) Delete(
|
||||
ctx context.Context,
|
||||
table string,
|
||||
keyField string,
|
||||
data map[string]any,
|
||||
) error {
|
||||
if keyField == "" {
|
||||
return fmt.Errorf("keyField 不能为空")
|
||||
}
|
||||
|
||||
keyVal, err := validateData(table, keyField, data)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"DELETE FROM %s WHERE %s = $1",
|
||||
table,
|
||||
keyField,
|
||||
)
|
||||
|
||||
_, err = c.exec().Exec(ctx, sql, keyVal)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *Client) BatchDelete(
|
||||
ctx context.Context,
|
||||
table string,
|
||||
keyField string,
|
||||
dataList []map[string]any,
|
||||
) error {
|
||||
|
||||
if table == "" {
|
||||
return fmt.Errorf("表名不能为空")
|
||||
}
|
||||
|
||||
if keyField == "" {
|
||||
return fmt.Errorf("主键字段不能为空")
|
||||
}
|
||||
|
||||
if len(dataList) == 0 {
|
||||
return fmt.Errorf("数据不能为空")
|
||||
}
|
||||
|
||||
var (
|
||||
placeholders []string
|
||||
args []any
|
||||
argIndex = 1
|
||||
)
|
||||
|
||||
for _, data := range dataList {
|
||||
keyVal, err := validateData(table, keyField, data)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if keyVal == nil {
|
||||
return fmt.Errorf("主键字段值不能为空")
|
||||
}
|
||||
|
||||
placeholders = append(placeholders, fmt.Sprintf("$%d", argIndex))
|
||||
args = append(args, keyVal)
|
||||
argIndex++
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"DELETE FROM %s WHERE %s IN (%s)",
|
||||
table,
|
||||
keyField,
|
||||
strings.Join(placeholders, ", "),
|
||||
)
|
||||
|
||||
_, err := c.exec().Exec(ctx, sql, args...)
|
||||
return err
|
||||
}
|
||||
@@ -1,120 +0,0 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// 构建 INSERT SQL(单条)
|
||||
func buildInsertSQL(table string, data map[string]any) (string, []any) {
|
||||
var (
|
||||
columns []string
|
||||
placeholders []string
|
||||
args []any
|
||||
)
|
||||
|
||||
i := 1
|
||||
for col, val := range data {
|
||||
columns = append(columns, col)
|
||||
placeholders = append(placeholders, fmt.Sprintf("$%d", i))
|
||||
args = append(args, val)
|
||||
i++
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"INSERT INTO %s (%s) VALUES (%s)",
|
||||
table,
|
||||
strings.Join(columns, ", "),
|
||||
strings.Join(placeholders, ", "),
|
||||
)
|
||||
|
||||
return sql, args
|
||||
}
|
||||
|
||||
func (c *Client) Insert(
|
||||
ctx context.Context,
|
||||
table string,
|
||||
keyField string,
|
||||
data map[string]any,
|
||||
) error {
|
||||
if _, err := validateData(table, keyField, data); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
sql, args := buildInsertSQL(table, data)
|
||||
|
||||
_, err := c.exec().Exec(ctx, sql, args...)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *Client) BatchInsert(
|
||||
ctx context.Context,
|
||||
table string,
|
||||
keyField string,
|
||||
dataList []map[string]any,
|
||||
) error {
|
||||
if len(dataList) == 0 {
|
||||
return fmt.Errorf("数据不能为空")
|
||||
}
|
||||
|
||||
// 用第一条确定列顺序
|
||||
first := dataList[0]
|
||||
if _, err := validateData(table, keyField, first); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var columns []string
|
||||
for col := range first {
|
||||
columns = append(columns, col)
|
||||
}
|
||||
|
||||
sort.Strings(columns)
|
||||
|
||||
var (
|
||||
valueStrings []string
|
||||
args []any
|
||||
argIndex = 1
|
||||
)
|
||||
|
||||
for _, data := range dataList {
|
||||
// 统一校验
|
||||
if _, err := validateData(table, keyField, data); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 字段数量检查
|
||||
if len(data) != len(columns) {
|
||||
return fmt.Errorf("批量插入失败:数据字段不一致")
|
||||
}
|
||||
|
||||
var placeholders []string
|
||||
|
||||
for _, col := range columns {
|
||||
val, ok := data[col]
|
||||
if !ok {
|
||||
return fmt.Errorf("缺少字段: %s", col)
|
||||
}
|
||||
|
||||
placeholders = append(placeholders, fmt.Sprintf("$%d", argIndex))
|
||||
args = append(args, val)
|
||||
argIndex++
|
||||
}
|
||||
|
||||
valueStrings = append(
|
||||
valueStrings,
|
||||
fmt.Sprintf("(%s)", strings.Join(placeholders, ", ")),
|
||||
)
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"INSERT INTO %s (%s) VALUES %s",
|
||||
table,
|
||||
strings.Join(columns, ", "),
|
||||
strings.Join(valueStrings, ", "),
|
||||
)
|
||||
|
||||
_, err := c.exec().Exec(ctx, sql, args...)
|
||||
return err
|
||||
}
|
||||
@@ -1,82 +0,0 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func (c *Client) LoadData(
|
||||
ctx context.Context,
|
||||
viewName string,
|
||||
searchCondition string,
|
||||
orderBy string,
|
||||
searchColumns []string,
|
||||
args []any,
|
||||
) ([][]any, error) {
|
||||
|
||||
var where string
|
||||
|
||||
// WHERE 构造
|
||||
if searchCondition != "" && len(searchColumns) > 0 {
|
||||
var conditions []string
|
||||
|
||||
for _, col := range searchColumns {
|
||||
conditions = append(conditions, fmt.Sprintf("%s LIKE ?", col))
|
||||
}
|
||||
|
||||
where = " WHERE (" + strings.Join(conditions, " OR ") + ")"
|
||||
|
||||
// 自动加 %
|
||||
for i := range args {
|
||||
if s, ok := args[i].(string); ok {
|
||||
args[i] = "%" + s + "%"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ORDER BY
|
||||
var order string
|
||||
if orderBy != "" {
|
||||
order = " ORDER BY " + orderBy
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf("SELECT * FROM %s%s%s", viewName, where, order)
|
||||
|
||||
return c.query(ctx, sql, args...)
|
||||
}
|
||||
|
||||
func (c *Client) LoadDataBySQL(
|
||||
ctx context.Context,
|
||||
sql string,
|
||||
args []any,
|
||||
) ([][]any, error) {
|
||||
return c.query(ctx, sql, args...)
|
||||
}
|
||||
|
||||
// 统一查询方法(核心优化)
|
||||
func (c *Client) query(ctx context.Context, sql string, args ...any) ([][]any, error) {
|
||||
sql = convertPlaceholder(sql)
|
||||
|
||||
rows, err := c.exec().Query(ctx, sql, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
result := make([][]any, 0, 16)
|
||||
|
||||
for rows.Next() {
|
||||
values, err := rows.Values()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
row := make([]any, len(values))
|
||||
copy(row, values)
|
||||
|
||||
result = append(result, row)
|
||||
}
|
||||
|
||||
return result, rows.Err()
|
||||
}
|
||||
@@ -1,34 +0,0 @@
|
||||
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
|
||||
}
|
||||
|
||||
// 创建事务 client
|
||||
txClient := &Client{
|
||||
pool: c.pool,
|
||||
tx: tx,
|
||||
}
|
||||
|
||||
defer func(tx pgx.Tx, ctx context.Context) {
|
||||
err := tx.Rollback(ctx)
|
||||
if err != nil {
|
||||
|
||||
}
|
||||
}(tx, ctx)
|
||||
|
||||
if err := fn(txClient); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return tx.Commit(ctx)
|
||||
}
|
||||
@@ -1,140 +0,0 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func (c *Client) Update(
|
||||
ctx context.Context,
|
||||
table string,
|
||||
keyField string,
|
||||
data map[string]any,
|
||||
) error {
|
||||
keyVal, err := validateData(table, keyField, data)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var (
|
||||
setClauses []string
|
||||
args []any
|
||||
i = 1
|
||||
)
|
||||
|
||||
for col, val := range data {
|
||||
if col == keyField {
|
||||
continue
|
||||
}
|
||||
setClauses = append(setClauses, fmt.Sprintf("%s=$%d", col, i))
|
||||
args = append(args, val)
|
||||
i++
|
||||
}
|
||||
|
||||
// WHERE 条件
|
||||
where := fmt.Sprintf("%s=$%d", keyField, i)
|
||||
args = append(args, keyVal)
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"UPDATE %s SET %s WHERE %s",
|
||||
table,
|
||||
strings.Join(setClauses, ", "),
|
||||
where,
|
||||
)
|
||||
|
||||
_, err = c.exec().Exec(ctx, sql, args...)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *Client) BatchUpdate(
|
||||
ctx context.Context,
|
||||
table string,
|
||||
keyField string,
|
||||
dataList []map[string]any,
|
||||
) error {
|
||||
|
||||
if len(dataList) == 0 {
|
||||
return fmt.Errorf("数据不能为空")
|
||||
}
|
||||
|
||||
// 用第一条数据确定字段
|
||||
first := dataList[0]
|
||||
if _, err := validateData(table, keyField, first); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 提取字段(排除主键)+ 排序(关键)
|
||||
var columns []string
|
||||
for col := range first {
|
||||
if col != keyField {
|
||||
columns = append(columns, col)
|
||||
}
|
||||
}
|
||||
sort.Strings(columns)
|
||||
|
||||
var (
|
||||
args []any
|
||||
argIndex = 1
|
||||
)
|
||||
|
||||
// CASE 语句
|
||||
var setClauses []string
|
||||
|
||||
for _, col := range columns {
|
||||
var caseBuilder strings.Builder
|
||||
|
||||
caseBuilder.WriteString(fmt.Sprintf("%s = CASE %s ", col, keyField))
|
||||
|
||||
for _, data := range dataList {
|
||||
// 校验 key
|
||||
keyVal, err := validateData(table, keyField, data)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
val, ok := data[col]
|
||||
if !ok {
|
||||
return fmt.Errorf("缺少字段: %s", col)
|
||||
}
|
||||
|
||||
caseBuilder.WriteString(fmt.Sprintf(
|
||||
"WHEN $%d THEN $%d ",
|
||||
argIndex,
|
||||
argIndex+1,
|
||||
))
|
||||
|
||||
args = append(args, keyVal, val)
|
||||
argIndex += 2
|
||||
}
|
||||
|
||||
caseBuilder.WriteString("END")
|
||||
setClauses = append(setClauses, caseBuilder.String())
|
||||
}
|
||||
|
||||
// WHERE IN
|
||||
var wherePlaceholders []string
|
||||
|
||||
for _, data := range dataList {
|
||||
keyVal, ok := data[keyField]
|
||||
if !ok || keyVal == nil {
|
||||
return fmt.Errorf("缺少主键字段值: %s", keyField)
|
||||
}
|
||||
|
||||
wherePlaceholders = append(wherePlaceholders, fmt.Sprintf("$%d", argIndex))
|
||||
args = append(args, keyVal)
|
||||
argIndex++
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"UPDATE %s SET %s WHERE %s IN (%s)",
|
||||
table,
|
||||
strings.Join(setClauses, ", "),
|
||||
keyField,
|
||||
strings.Join(wherePlaceholders, ", "),
|
||||
)
|
||||
|
||||
_, err := c.exec().Exec(ctx, sql, args...)
|
||||
return err
|
||||
}
|
||||
Reference in new issue
Block a user