141 lines
2.5 KiB
Go
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
|
|
}
|