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 }