142 lines
2.5 KiB
Go
142 lines
2.5 KiB
Go
package db
|
|
|
|
import (
|
|
"allapp-go/internal/middleware"
|
|
"context"
|
|
"fmt"
|
|
"sort"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
// 构建 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
|
|
}
|
|
|
|
if !auditExcludeTables[table] {
|
|
now := time.Now()
|
|
|
|
delete(data, "create_time")
|
|
delete(data, "update_time")
|
|
delete(data, "create_by")
|
|
delete(data, "update_by")
|
|
|
|
data["create_time"] = now
|
|
data["update_time"] = now
|
|
|
|
userID, flag := middleware.GetUserID(ctx)
|
|
|
|
if flag {
|
|
data["create_by"] = userID
|
|
data["update_by"] = userID
|
|
}
|
|
}
|
|
|
|
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
|
|
}
|