Files
workspace/code/app/app-go/pkg/db/update.go
T
2026-04-07 17:19:52 +08:00

141 lines
2.5 KiB
Go

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
}