rows: fix PR feedback

This commit is contained in:
Tristan Colgate
2021-12-14 12:38:31 +00:00
parent f4f50f46dc
commit 7aff84f564
4 changed files with 27 additions and 15 deletions
-5
View File
@@ -8,11 +8,6 @@ import (
//go:generate go run ./tools/rows_picker_gen.go -o rows_picker.go
// Compile time validation that our types implement the expected interfaces
var (
_ driver.Rows = wrappedRows{}
)
// RowsUnwrapper must be used by any middleware that provides its own wrapping
// for driver.Rows. Unwrap should return the original driver.Rows the
// middleware received. You may wish to wrap the driver.Rows returned by the
+16 -4
View File
@@ -4,11 +4,23 @@ import (
"context"
"database/sql"
"database/sql/driver"
"log"
"fmt"
"reflect"
"sync/atomic"
"testing"
)
var driverCount = int32(0)
func driverName(t *testing.T) string {
c := atomic.LoadInt32(&driverCount)
name := fmt.Sprintf("driver-%s-%d", t.Name(), c)
c++
atomic.StoreInt32(&driverCount, c)
return name
}
type rowsCloseInterceptor struct {
NullInterceptor
@@ -24,7 +36,7 @@ func (r *rowsCloseInterceptor) RowsClose(ctx context.Context, rows driver.Rows)
}
func TestRowsClose(t *testing.T) {
driverName := t.Name()
driverName := driverName(t)
interceptor := rowsCloseInterceptor{}
con := fakeConn{}
@@ -98,7 +110,7 @@ func TestRowsNext(t *testing.T) {
rows: rows,
}
con.stmt = stmt
driverName := t.Name()
driverName := driverName(t)
interceptor := rowsNextInterceptor{}
sql.Register(
@@ -188,7 +200,7 @@ func TestRows_LikePGX(t *testing.T) {
rows: rows,
}
con.stmt = stmt
driverName := t.Name()
driverName := driverName(t)
interceptor := rowsNextInterceptor{}
sql.Register(
+6 -5
View File
@@ -11,7 +11,7 @@ import (
// driver.DefaultParameterConverter is used when neither stmt nor con
// implements any value converters.
func TestDefaultParameterConversion(t *testing.T) {
driverNameWithSQLmw := t.Name() + "sqlmw"
driverName := driverName(t)
expectVal := int64(1)
con := &fakeConn{}
@@ -24,11 +24,11 @@ func TestDefaultParameterConversion(t *testing.T) {
con.stmt = fakeStmt
sql.Register(
driverNameWithSQLmw,
driverName,
Driver(&fakeDriver{conn: con}, &NullInterceptor{}),
)
db, err := sql.Open(driverNameWithSQLmw, "")
db, err := sql.Open(driverName, "")
if err != nil {
t.Fatalf("Failed to open: %v", err)
}
@@ -134,8 +134,9 @@ func TestWrappedStmt_CheckNamedValue(t *testing.T) {
for name, test := range tests {
t.Run(name, func(t *testing.T) {
sql.Register("fake-driver:"+name, Driver(test.fd, &fakeInterceptor{}))
db, err := sql.Open("fake-driver:"+name, "dummy")
driverName := driverName(t)
sql.Register(driverName, Driver(test.fd, &fakeInterceptor{}))
db, err := sql.Open(driverName, "dummy")
if err != nil {
t.Errorf("Failed to open: %v", err)
}
+5 -1
View File
@@ -23,7 +23,6 @@ func main() {
log.Fatalf("could not create file %q, %v", *fn, err)
}
}
defer out.Close()
intfs := []string{
"NextResultSet",
@@ -51,6 +50,11 @@ func main() {
fmt.Fprintln(out, "")
genWrapRows(out, intfs)
err = out.Close()
if err != nil {
log.Fatalf("could close file, %v", err)
}
}
func genComment(w io.Writer) {