u
This commit is contained in:
1 parent
89aaef1d6b
commit
6d54a9e407
192 files changed
+971
-51708
No files matched your search
@@ -1,141 +1,93 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"allapp-go/internal/middleware"
|
||||
"context"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
func (c *Client) Update(
|
||||
ctx context.Context,
|
||||
table string,
|
||||
keyField string,
|
||||
data map[string]any,
|
||||
) error {
|
||||
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
|
||||
}
|
||||
|
||||
if !auditExcludeTables[table] {
|
||||
now := time.Now()
|
||||
data = applyMetaFields(ctx, table, data, false)
|
||||
|
||||
delete(data, "update_time")
|
||||
delete(data, "update_by")
|
||||
|
||||
data["update_time"] = now
|
||||
|
||||
userID, flag := middleware.GetUserID(ctx)
|
||||
if flag {
|
||||
data["update_by"] = userID
|
||||
}
|
||||
}
|
||||
|
||||
// 👉 表名加引号(支持 schema)
|
||||
tableParts := strings.Split(table, ".")
|
||||
for i, p := range tableParts {
|
||||
tableParts[i] = `"` + p + `"`
|
||||
}
|
||||
table = strings.Join(tableParts, ".")
|
||||
|
||||
// 👉 keyField 加引号
|
||||
keyField = `"` + keyField + `"`
|
||||
tableSQL := quoteTable(table)
|
||||
keySQL := quoteCol(keyField)
|
||||
|
||||
var (
|
||||
setClauses []string
|
||||
args []any
|
||||
i = 1
|
||||
set []string
|
||||
args []any
|
||||
i = 1
|
||||
)
|
||||
|
||||
for col, val := range data {
|
||||
if col == keyField {
|
||||
for k, v := range data {
|
||||
if k == keyField {
|
||||
continue
|
||||
}
|
||||
|
||||
// 👉 列名加引号
|
||||
colQuoted := `"` + col + `"`
|
||||
|
||||
setClauses = append(setClauses, fmt.Sprintf("%s=$%d", colQuoted, i))
|
||||
args = append(args, val)
|
||||
set = append(set, fmt.Sprintf("%s=$%d", quoteCol(k), i))
|
||||
args = append(args, v)
|
||||
i++
|
||||
}
|
||||
|
||||
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,
|
||||
"UPDATE %s SET %s WHERE %s=$%d",
|
||||
tableSQL,
|
||||
strings.Join(set, ", "),
|
||||
keySQL,
|
||||
i,
|
||||
)
|
||||
|
||||
_, 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("数据不能为空")
|
||||
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 := dataList[0]
|
||||
first := list[0]
|
||||
if _, err := validateData(table, keyField, first); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 👉 表名加引号(支持 schema)
|
||||
tableParts := strings.Split(table, ".")
|
||||
for i, p := range tableParts {
|
||||
tableParts[i] = `"` + p + `"`
|
||||
}
|
||||
table = strings.Join(tableParts, ".")
|
||||
tableSQL := quoteTable(table)
|
||||
keySQL := quoteCol(keyField)
|
||||
|
||||
// 👉 keyField 加引号
|
||||
keyFieldQuoted := `"` + keyField + `"`
|
||||
|
||||
var columns []string
|
||||
for col := range first {
|
||||
if col != keyField {
|
||||
columns = append(columns, col)
|
||||
var cols []string
|
||||
for k := range first {
|
||||
if k != keyField {
|
||||
cols = append(cols, k)
|
||||
}
|
||||
}
|
||||
sort.Strings(columns)
|
||||
sort.Strings(cols)
|
||||
|
||||
var (
|
||||
args []any
|
||||
argIndex = 1
|
||||
sets []string
|
||||
)
|
||||
|
||||
var setClauses []string
|
||||
for _, col := range cols {
|
||||
colSQL := quoteCol(col)
|
||||
|
||||
for _, col := range columns {
|
||||
colQuoted := `"` + col + `"`
|
||||
var caseSQL strings.Builder
|
||||
caseSQL.WriteString(fmt.Sprintf("%s = CASE %s ", colSQL, keySQL))
|
||||
|
||||
var caseBuilder strings.Builder
|
||||
caseBuilder.WriteString(fmt.Sprintf("%s = CASE %s ", colQuoted, keyFieldQuoted))
|
||||
for _, row := range list {
|
||||
row = applyMetaFields(ctx, table, row, false)
|
||||
|
||||
for _, data := range dataList {
|
||||
keyVal, err := validateData(table, keyField, data)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
keyVal, _ := row[keyField]
|
||||
val := row[col]
|
||||
|
||||
val, ok := data[col]
|
||||
if !ok {
|
||||
return fmt.Errorf("缺少字段: %s", col)
|
||||
}
|
||||
|
||||
caseBuilder.WriteString(fmt.Sprintf(
|
||||
caseSQL.WriteString(fmt.Sprintf(
|
||||
"WHEN $%d THEN $%d ",
|
||||
argIndex,
|
||||
argIndex+1,
|
||||
@@ -145,29 +97,23 @@ func (c *Client) BatchUpdate(
|
||||
argIndex += 2
|
||||
}
|
||||
|
||||
caseBuilder.WriteString("END")
|
||||
setClauses = append(setClauses, caseBuilder.String())
|
||||
caseSQL.WriteString("END")
|
||||
sets = append(sets, caseSQL.String())
|
||||
}
|
||||
|
||||
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)
|
||||
var where []string
|
||||
for _, row := range list {
|
||||
where = append(where, fmt.Sprintf("$%d", argIndex))
|
||||
args = append(args, row[keyField])
|
||||
argIndex++
|
||||
}
|
||||
|
||||
sql := fmt.Sprintf(
|
||||
"UPDATE %s SET %s WHERE %s IN (%s)",
|
||||
table,
|
||||
strings.Join(setClauses, ", "),
|
||||
keyFieldQuoted,
|
||||
strings.Join(wherePlaceholders, ", "),
|
||||
tableSQL,
|
||||
strings.Join(sets, ", "),
|
||||
keySQL,
|
||||
strings.Join(where, ", "),
|
||||
)
|
||||
|
||||
_, err := c.exec().Exec(ctx, sql, args...)
|
||||
|
||||
Reference in new issue
Block a user