package db import ( "allapp-go/internal/middleware" "context" "fmt" "sort" "strings" "time" ) // 构建 INSERT SQL(单条) func buildInsertSQL(table string, data map[string]any) (string, []any) { var ( columns []string placeholders []string args []any ) // 处理 schema.table 的情况 tableParts := strings.Split(table, ".") for i, p := range tableParts { tableParts[i] = `"` + p + `"` } table = strings.Join(tableParts, ".") i := 1 for col, val := range data { // 列名加引号 columns = append(columns, `"`+col+`"`) placeholders = append(placeholders, fmt.Sprintf("$%d", i)) args = append(args, val) i++ } sql := fmt.Sprintf( "INSERT INTO %s (%s) VALUES (%s)", table, strings.Join(columns, ", "), strings.Join(placeholders, ", "), ) return sql, args } func (c *Client) Insert( ctx context.Context, table string, keyField string, data map[string]any, ) error { if _, err := validateData(table, keyField, data); err != nil { return err } if !auditExcludeTables[table] { now := time.Now() delete(data, "create_time") delete(data, "update_time") delete(data, "create_by") delete(data, "update_by") data["create_time"] = now data["update_time"] = now userID, flag := middleware.GetUserID(ctx) if flag { data["create_by"] = userID data["update_by"] = userID } } sql, args := buildInsertSQL(table, data) _, err := c.exec().Exec(ctx, sql, args...) return err } func (c *Client) BatchInsert( 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 } // 👉 表名加引号(支持 schema) tableParts := strings.Split(table, ".") for i, p := range tableParts { tableParts[i] = `"` + p + `"` } table = strings.Join(tableParts, ".") var columns []string for col := range first { columns = append(columns, col) } sort.Strings(columns) // 👉 列名加引号 var quotedColumns []string for _, col := range columns { quotedColumns = append(quotedColumns, `"`+col+`"`) } var ( valueStrings []string args []any argIndex = 1 ) for _, data := range dataList { if _, err := validateData(table, keyField, data); err != nil { return err } if len(data) != len(columns) { return fmt.Errorf("批量插入失败:数据字段不一致") } var placeholders []string for _, col := range columns { val, ok := data[col] if !ok { return fmt.Errorf("缺少字段: %s", col) } placeholders = append(placeholders, fmt.Sprintf("$%d", argIndex)) args = append(args, val) argIndex++ } valueStrings = append( valueStrings, fmt.Sprintf("(%s)", strings.Join(placeholders, ", ")), ) } sql := fmt.Sprintf( "INSERT INTO %s (%s) VALUES %s", table, strings.Join(quotedColumns, ", "), strings.Join(valueStrings, ", "), ) _, err := c.exec().Exec(ctx, sql, args...) return err }