package db import ( "fmt" "strings" ) // ---------------- 内部辅助函数 ---------------- // buildUpdateSQL 构建 UPDATE SQL 和参数 // idColumn: 主键列名 func buildUpdateSQL(table string, data map[string]interface{}, idColumn string) (string, []interface{}) { var setParts []string var values []interface{} i := 1 for k, v := range data { setParts = append(setParts, fmt.Sprintf("%s=$%d", k, i)) values = append(values, v) i++ } // WHERE id = $n sql := fmt.Sprintf("UPDATE %s SET %s WHERE %s=$%d", table, strings.Join(setParts, ", "), idColumn, i, ) return sql, values } // ---------------- 公共方法 ---------------- // UpdateOne 更新单条记录,按主键 idColumn func UpdateOne(table string, idColumn string, id interface{}, data map[string]interface{}) (int64, error) { tx, err := DB.Beginx() if err != nil { return 0, err } sql, values := buildUpdateSQL(table, data, idColumn) values = append(values, id) // 最后一个参数是 id res, err := tx.Exec(sql, values...) if err != nil { err := tx.Rollback() if err != nil { return 0, err } return 0, err } if err := tx.Commit(); err != nil { return 0, err } return res.RowsAffected() } // UpdateBatch 批量更新,dataList 中每个 map 必须包含主键 idColumn func UpdateBatch(table string, idColumn string, dataList []map[string]interface{}) (int64, error) { if len(dataList) == 0 { return 0, nil } tx, err := DB.Beginx() if err != nil { return 0, err } total := int64(0) for _, data := range dataList { idValue, ok := data[idColumn] if !ok { tx.Rollback() return 0, fmt.Errorf("缺少主键列 %s", idColumn) } // 移除主键列,否则会重复在 SET 中出现 delete(data, idColumn) sql, values := buildUpdateSQL(table, data, idColumn) values = append(values, idValue) res, err := tx.Exec(sql, values...) if err != nil { tx.Rollback() return 0, err } rows, err := res.RowsAffected() if err != nil { tx.Rollback() return 0, err } total += rows } if err := tx.Commit(); err != nil { return 0, err } return total, nil }