Files
workspace/code/allapp/allapp-go-v3/pkg/db/update.go
T
2026-04-24 22:37:52 +08:00

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
}