11package main
22
33import (
4+ "context"
45 sql "database/sql"
56 "encoding/json"
67 "fmt"
7- pq "github.com/lib/pq"
88 "github.com/jackc/pgx"
9+ pq "github.com/lib/pq"
910 "gopkg.in/alecthomas/kingpin.v2"
1011 "io/ioutil"
1112 "log"
1213 "math"
1314 "os"
15+ "reflect"
1416 "regexp"
1517 "runtime"
1618 "strings"
@@ -23,20 +25,17 @@ type WorkerFunc func(time.Time, time.Duration, time.Duration,
2325 string , []interface {},
2426 * sync.WaitGroup , ReportFunc )
2527
26-
2728type TypeCaster func (interface {}) interface {}
2829
29-
3030type CopyInfo struct {
3131 TableName string
32- Columns []string
33- Types []TypeCaster
34- Rows [][]interface {}
32+ Columns []string
33+ Types []TypeCaster
34+ Rows [][]interface {}
3535}
3636
37-
3837func get_type_casters () map [string ]TypeCaster {
39- cast_map := map [string ] TypeCaster {
38+ cast_map := map [string ]TypeCaster {
4039 "int4" : func (v interface {}) interface {} {
4140 i := int (v .(float64 ))
4241 return i
@@ -54,7 +53,6 @@ func get_type_casters() map[string]TypeCaster {
5453 return cast_map
5554}
5655
57-
5856func get_copy_info (db * sql.DB , query string , args []interface {}) CopyInfo {
5957 re := regexp .MustCompile (`COPY (\w+)\s*\(\s*((?:\w+)(?:,\s*\w+)*)\s*\)` )
6058 match := re .FindStringSubmatch (query )
@@ -67,7 +65,7 @@ func get_copy_info(db *sql.DB, query string, args []interface{}) CopyInfo {
6765
6866 col_parts := strings .Split (match [2 ], "," )
6967 cols := make ([]string , len (col_parts ))
70- for i , cp := range ( col_parts ) {
68+ for i , cp := range col_parts {
7169 cols [i ] = strings .Trim (cp , " " )
7270 }
7371
@@ -144,13 +142,12 @@ func get_copy_info(db *sql.DB, query string, args []interface{}) CopyInfo {
144142
145143 return CopyInfo {
146144 TableName : table ,
147- Columns : cols ,
148- Types : casters ,
149- Rows : copyrows ,
145+ Columns : cols ,
146+ Types : casters ,
147+ Rows : copyrows ,
150148 }
151149}
152150
153-
154151func lib_pq_worker (
155152 start time.Time , duration time.Duration , timeout time.Duration ,
156153 query string , query_args []interface {}, wg * sync.WaitGroup ,
@@ -221,6 +218,56 @@ func lib_pq_worker(
221218 break
222219 }
223220 }
221+ } else if len (query_args ) > 0 && reflect .ValueOf (query_args [0 ]).Kind () == reflect .Map {
222+ args := query_args [0 ].(map [string ]interface {})
223+ row := args ["row" ].([]interface {})
224+ count := int (args ["count" ].(float64 ))
225+
226+
227+ for time .Since (start ) < duration || duration == 0 {
228+ req_start := time .Now ()
229+
230+ txn , err := db .Begin ()
231+ if err != nil {
232+ log .Fatal (err )
233+ }
234+
235+ stmt , err := txn .Prepare (query )
236+ if err != nil {
237+ log .Fatal (err )
238+ }
239+
240+ for i := 0 ; i < count ; i ++ {
241+ _ , err := stmt .Exec (row ... )
242+ if err != nil {
243+ log .Fatal (err )
244+ }
245+ nrows += 1
246+ }
247+
248+ err = stmt .Close ()
249+ if err != nil {
250+ log .Fatal (err )
251+ }
252+
253+ err = txn .Commit ()
254+ if err != nil {
255+ log .Fatal (err )
256+ }
257+
258+ req_time := time .Since (req_start ).Nanoseconds () / 1000 / 10
259+ latency_stats [req_time ] += 1
260+ if req_time > max_latency {
261+ max_latency = req_time
262+ }
263+ if req_time < min_latency {
264+ min_latency = req_time
265+ }
266+ queries += 1
267+ if duration == 0 {
268+ break
269+ }
270+ }
224271 } else {
225272
226273 var record []interface {}
@@ -310,16 +357,11 @@ func pgx_worker(
310357 }
311358 defer adminConn .Close ()
312359
313- db , err := pgx .Connect (pgx.ConnConfig {
314- Host : * pghost ,
315- Port : uint16 (* pgport ),
316- Database : * pgdatabase ,
317- User : * pguser ,
318- })
360+ db , err := pgx .Connect (context .Background (), conninfo )
319361 if err != nil {
320362 log .Fatal (err )
321363 }
322- defer db .Close ()
364+ defer db .Close (context . Background () )
323365
324366 latency_stats := make ([]int64 , timeout / 1000 / 10 )
325367 min_latency := int64 (math .MaxInt64 )
@@ -334,9 +376,10 @@ func pgx_worker(
334376 req_start := time .Now ()
335377
336378 copy_count , err := db .CopyFrom (
337- pgx.Identifier {copy .TableName },
338- copy .Columns ,
339- pgx .CopyFromRows (copy .Rows ),
379+ context .Background (),
380+ pgx.Identifier {copy .TableName },
381+ copy .Columns ,
382+ pgx .CopyFromRows (copy .Rows ),
340383 )
341384
342385 if err != nil {
@@ -357,20 +400,68 @@ func pgx_worker(
357400 break
358401 }
359402 }
403+ } else if len (query_args ) > 0 && reflect .ValueOf (query_args [0 ]).Kind () == reflect .Map {
404+ args := query_args [0 ].(map [string ]interface {})
405+ row := args ["row" ].([]interface {})
406+ count := int (args ["count" ].(float64 ))
407+
408+ _ , err = db .Prepare (context .Background (), "testquery" , query )
409+ if err != nil {
410+ log .Fatal (err )
411+ }
412+
413+ for time .Since (start ) < duration || duration == 0 {
414+ req_start := time .Now ()
415+
416+ batch := & pgx.Batch {}
417+ for i := 0 ; i < count ; i ++ {
418+ batch .Queue ("testquery" , row ... )
419+ }
420+
421+ br := db .SendBatch (context .Background (), batch )
422+ for i := 0 ; i < count ; i ++ {
423+ rows , err := br .Query ()
424+ if err != nil {
425+ log .Fatal (err )
426+ }
427+ if rows .Err () != nil {
428+ log .Fatal (rows .Err ())
429+ }
430+ nrows += 1
431+ }
432+
433+ err = br .Close ()
434+ if err != nil {
435+ log .Fatal (err )
436+ }
437+
438+ req_time := time .Since (req_start ).Nanoseconds () / 1000 / 10
439+ latency_stats [req_time ] += 1
440+ if req_time > max_latency {
441+ max_latency = req_time
442+ }
443+ if req_time < min_latency {
444+ min_latency = req_time
445+ }
446+ queries += 1
447+ if duration == 0 {
448+ break
449+ }
450+ }
360451 } else {
361452 fixed_query_args := make ([]interface {}, len (query_args ))
362453 for i , arg := range query_args {
363454 fixed_query_args [i ] = fmt .Sprintf ("%v" , arg )
364455 }
365456
366- _ , err = db .Prepare ("testquery" , query )
457+ _ , err = db .Prepare (context . Background (), "testquery" , query )
367458 if err != nil {
368459 log .Fatal (err )
369460 }
370461
371462 for time .Since (start ) < duration || duration == 0 {
372463 req_start := time .Now ()
373- rows , err := db .Query ("testquery" , fixed_query_args ... )
464+ rows , err := db .Query (context . Background (), "testquery" , fixed_query_args ... )
374465 if err != nil {
375466 log .Fatal (err )
376467 }
0 commit comments