149 lines
2.8 KiB
Go
149 lines
2.8 KiB
Go
package db
|
|
|
|
import (
|
|
"fmt"
|
|
"strings"
|
|
|
|
"base-go-v2/internal/utils/mapx"
|
|
)
|
|
|
|
// ---------------- 内部辅助函数 ----------------
|
|
|
|
// buildInsertSQL 构建插入 SQL 和参数
|
|
func buildInsertSQL(table string, data mapx.M, returning string) (string, []interface{}) {
|
|
columns := make([]string, 0, len(data))
|
|
placeholders := make([]string, 0, len(data))
|
|
values := make([]interface{}, 0, len(data))
|
|
|
|
i := 1
|
|
for k, v := range data {
|
|
columns = append(columns, k)
|
|
placeholders = append(placeholders, fmt.Sprintf("$%d", i))
|
|
values = append(values, v)
|
|
i++
|
|
}
|
|
|
|
sql := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)", table,
|
|
strings.Join(columns, ","), strings.Join(placeholders, ","))
|
|
|
|
if returning != "" {
|
|
sql += " RETURNING " + returning
|
|
}
|
|
|
|
return sql, values
|
|
}
|
|
|
|
// ---------------- 公共方法 ----------------
|
|
|
|
// InsertOne 插入单条记录,返回受影响行数
|
|
func InsertOne(table string, data mapx.M) (int64, error) {
|
|
tx, err := DB.Beginx()
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
|
|
sql, values := buildInsertSQL(table, data, "")
|
|
res, err := tx.Exec(sql, values...)
|
|
if err != nil {
|
|
tx.Rollback()
|
|
return 0, err
|
|
}
|
|
|
|
if err := tx.Commit(); err != nil {
|
|
return 0, err
|
|
}
|
|
|
|
return res.RowsAffected()
|
|
}
|
|
|
|
// InsertBatch 批量插入,返回总受影响行数
|
|
func InsertBatch(table string, dataList []mapx.M) (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 {
|
|
sql, values := buildInsertSQL(table, data, "")
|
|
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
|
|
}
|
|
|
|
// InsertOneReturnPK 插入单条记录,返回主键值
|
|
func InsertOneReturnPK(table string, pkColumn string, data mapx.M) (int64, error) {
|
|
tx, err := DB.Beginx()
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
|
|
sql, values := buildInsertSQL(table, data, pkColumn)
|
|
|
|
var pk int64
|
|
err = tx.Get(&pk, sql, values...)
|
|
if err != nil {
|
|
tx.Rollback()
|
|
return 0, err
|
|
}
|
|
|
|
if err := tx.Commit(); err != nil {
|
|
return 0, err
|
|
}
|
|
|
|
return pk, nil
|
|
}
|
|
|
|
// InsertBatchReturnPK 批量插入多条记录,返回主键数组
|
|
func InsertBatchReturnPK(table string, pkColumn string, dataList []mapx.M) ([]int64, error) {
|
|
if len(dataList) == 0 {
|
|
return nil, nil
|
|
}
|
|
|
|
tx, err := DB.Beginx()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
var pkList []int64
|
|
|
|
for _, data := range dataList {
|
|
sql, values := buildInsertSQL(table, data, pkColumn)
|
|
|
|
var pk int64
|
|
err := tx.Get(&pk, sql, values...)
|
|
if err != nil {
|
|
tx.Rollback()
|
|
return nil, err
|
|
}
|
|
|
|
pkList = append(pkList, pk)
|
|
}
|
|
|
|
if err := tx.Commit(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return pkList, nil
|
|
}
|