package db import ( "allapp-go/internal/middleware" "context" "fmt" "sort" "strings" "time" ) func (c *Client) Update( ctx context.Context, table string, keyField string, data map[string]any, ) error { keyVal, err := validateData(table, keyField, data) if err != nil { return err } if !auditExcludeTables[table] { now := time.Now() delete(data, "update_time") delete(data, "update_by") data["update_time"] = now userID, flag := middleware.GetUserID(ctx) if flag { data["update_by"] = userID } } // 👉 表名加引号(支持 schema) tableParts := strings.Split(table, ".") for i, p := range tableParts { tableParts[i] = `"` + p + `"` } table = strings.Join(tableParts, ".") // 👉 keyField 加引号 keyField = `"` + keyField + `"` var ( setClauses []string args []any i = 1 ) for col, val := range data { if col == keyField { continue } // 👉 列名加引号 colQuoted := `"` + col + `"` setClauses = append(setClauses, fmt.Sprintf("%s=$%d", colQuoted, i)) args = append(args, val) i++ } where := fmt.Sprintf("%s=$%d", keyField, i) args = append(args, keyVal) sql := fmt.Sprintf( "UPDATE %s SET %s WHERE %s", table, strings.Join(setClauses, ", "), where, ) _, err = c.exec().Exec(ctx, sql, args...) return err } func (c *Client) BatchUpdate( 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, ".") // 👉 keyField 加引号 keyFieldQuoted := `"` + keyField + `"` var columns []string for col := range first { if col != keyField { columns = append(columns, col) } } sort.Strings(columns) var ( args []any argIndex = 1 ) var setClauses []string for _, col := range columns { colQuoted := `"` + col + `"` var caseBuilder strings.Builder caseBuilder.WriteString(fmt.Sprintf("%s = CASE %s ", colQuoted, keyFieldQuoted)) for _, data := range dataList { keyVal, err := validateData(table, keyField, data) if err != nil { return err } val, ok := data[col] if !ok { return fmt.Errorf("缺少字段: %s", col) } caseBuilder.WriteString(fmt.Sprintf( "WHEN $%d THEN $%d ", argIndex, argIndex+1, )) args = append(args, keyVal, val) argIndex += 2 } caseBuilder.WriteString("END") setClauses = append(setClauses, caseBuilder.String()) } var wherePlaceholders []string for _, data := range dataList { keyVal, ok := data[keyField] if !ok || keyVal == nil { return fmt.Errorf("缺少主键字段值: %s", keyField) } wherePlaceholders = append(wherePlaceholders, fmt.Sprintf("$%d", argIndex)) args = append(args, keyVal) argIndex++ } sql := fmt.Sprintf( "UPDATE %s SET %s WHERE %s IN (%s)", table, strings.Join(setClauses, ", "), keyFieldQuoted, strings.Join(wherePlaceholders, ", "), ) _, err := c.exec().Exec(ctx, sql, args...) return err }