104 lines
2.1 KiB
Go
104 lines
2.1 KiB
Go
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
|
|
}
|