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 }