123 lines
2.2 KiB
Go
123 lines
2.2 KiB
Go
package db
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"sort"
|
|
"strings"
|
|
)
|
|
|
|
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
|
|
}
|
|
|
|
data = applyMetaFields(ctx, table, data, false)
|
|
|
|
var (
|
|
set []string
|
|
args []any
|
|
i = 1
|
|
)
|
|
|
|
for k, v := range data {
|
|
if k == keyField {
|
|
continue
|
|
}
|
|
|
|
set = append(set, fmt.Sprintf("%s=$%d", quoteCol(k), i))
|
|
args = append(args, v)
|
|
i++
|
|
}
|
|
|
|
args = append(args, keyVal)
|
|
|
|
sql := fmt.Sprintf(
|
|
"UPDATE %s SET %s WHERE %s=$%d",
|
|
quoteTable(table),
|
|
strings.Join(set, ", "),
|
|
quoteCol(keyField),
|
|
i,
|
|
)
|
|
|
|
_, err = c.pool.Exec(ctx, sql, args...)
|
|
return err
|
|
}
|
|
|
|
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 := list[0]
|
|
if _, err := validateData(table, keyField, first); err != nil {
|
|
return err
|
|
}
|
|
|
|
tableSQL := quoteTable(table)
|
|
keySQL := quoteCol(keyField)
|
|
|
|
var cols []string
|
|
for k := range first {
|
|
if k != keyField {
|
|
cols = append(cols, k)
|
|
}
|
|
}
|
|
sort.Strings(cols)
|
|
|
|
var (
|
|
args []any
|
|
argIndex = 1
|
|
sets []string
|
|
)
|
|
|
|
// CASE 构建
|
|
for _, col := range cols {
|
|
colSQL := quoteCol(col)
|
|
|
|
var caseSQL strings.Builder
|
|
caseSQL.WriteString(fmt.Sprintf("%s = CASE %s ", colSQL, keySQL))
|
|
|
|
for _, row := range list {
|
|
row = applyMetaFields(ctx, table, row, false)
|
|
|
|
keyVal := row[keyField]
|
|
val := row[col]
|
|
|
|
caseSQL.WriteString(fmt.Sprintf(
|
|
"WHEN $%d THEN $%d ",
|
|
argIndex,
|
|
argIndex+1,
|
|
))
|
|
|
|
args = append(args, keyVal, val)
|
|
argIndex += 2
|
|
}
|
|
|
|
caseSQL.WriteString("END")
|
|
sets = append(sets, caseSQL.String())
|
|
}
|
|
|
|
// ⭐修复 IN 写法(关键)
|
|
var inPlaceholders []string
|
|
for _, row := range list {
|
|
inPlaceholders = append(inPlaceholders, fmt.Sprintf("$%d", argIndex))
|
|
args = append(args, row[keyField])
|
|
argIndex++
|
|
}
|
|
|
|
sql := fmt.Sprintf(
|
|
"UPDATE %s SET %s WHERE %s IN (%s)",
|
|
tableSQL,
|
|
strings.Join(sets, ", "),
|
|
keySQL,
|
|
strings.Join(inPlaceholders, ", "),
|
|
)
|
|
|
|
_, err := c.pool.Exec(ctx, sql, args...) // ⭐修复点
|
|
return err
|
|
}
|