Files
workspace/code/base-project/base-go-v2/internal/db/insert.go
T
2025-11-04 17:32:32 +08:00

147 lines
2.9 KiB
Go

package db
import (
"fmt"
"strings"
)
// ---------------- 内部辅助函数 ----------------
// buildInsertSQL 构建插入 SQL 和参数
func buildInsertSQL(table string, data map[string]interface{}, 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 map[string]interface{}) (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 []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 {
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 map[string]interface{}) (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 []map[string]interface{}) ([]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
}