This commit is contained in:
oneao committed 2026-04-20 17:20:47 +08:00
1 parent 89aaef1d6b
commit 6d54a9e407
192 files changed
+971 -51708

No files matched your search

+47 -101
View File
@@ -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...)