call stmt.ColumnConverter when implemented by parent statement

According to the stdlib driver documentation
(https://github.com/golang/go/blob/bc51e930274a5d5835ac8797978afc0864c9e30c/src/database/sql/driver/driver.go#L385)
value checkers should be called in the following order:

> [..] stopping at the first found match: Stmt.NamedValueChecker,
> Conn.NamedValueChecker, Stmt.ColumnConverter,
> DefaultParameterConverter.

sqlmw was not calling Stmt.ColumnConverter when it was implemented.

This commit changes the behavior to call Stmt.ColumnConverter, if it is
implemented and the NamedValueCheckers are not implemented.
This is done by returning ErrSkip in wrappedStmt.CheckNamedValue() if
neither the parent statement nor the conn implements CheckNamedValue.
The sql package will call wrappedStmt.ColumnConverter() if ErrSkip was
returned.

wrappedStmt.CheckNamedValue() can not check only if the stmt
implements CheckNamedValue and return ErrSkip.
It must also call CheckNamedValue() on the connection if it
was not implemented for the stmt. This is because the stdlib sql package
calls wrappedStmt.CheckNamedValue() if is implemented on the stmt OR on
the connection.

The commit also adds a testcase to verify that ColumnConverter is
called.
This commit is contained in:
Fabian Holler
2021-08-03 11:05:09 +02:00
parent d5c93a81be
commit 1e77bed04f
4 changed files with 51 additions and 21 deletions
+1 -6
View File
@@ -8,15 +8,10 @@ var (
_ driver.NamedValueChecker = wrappedConn{}
)
func defaultCheckNamedValue(nv *driver.NamedValue) (err error) {
nv.Value, err = driver.DefaultParameterConverter.ConvertValue(nv.Value)
return err
}
func (c wrappedConn) CheckNamedValue(v *driver.NamedValue) error {
if checker, ok := c.parent.(driver.NamedValueChecker); ok {
return checker.CheckNamedValue(v)
}
return defaultCheckNamedValue(v)
return driver.ErrSkip
}
+13 -3
View File
@@ -14,7 +14,8 @@ func (d *fakeDriver) Open(_ string) (driver.Conn, error) {
}
type fakeStmt struct {
called bool
checkNamedValueCalled bool
columnConverterCalled bool
}
type fakeStmtWithCheckNamedValue struct {
@@ -25,6 +26,10 @@ type fakeStmtWithoutCheckNamedValue struct {
fakeStmt
}
type fakeStmtWithColumnConverter struct {
fakeStmt
}
func (s fakeStmt) Close() error {
return nil
}
@@ -41,8 +46,13 @@ func (s fakeStmt) Query(_ []driver.Value) (driver.Rows, error) {
return nil, nil
}
func (s *fakeStmtWithColumnConverter) ColumnConverter(_ int) driver.ValueConverter {
s.columnConverterCalled = true
return driver.DefaultParameterConverter
}
func (s *fakeStmtWithCheckNamedValue) CheckNamedValue(_ *driver.NamedValue) (err error) {
s.called = true
s.checkNamedValueCalled = true
return
}
@@ -67,7 +77,7 @@ func (c *fakeConn) PrepareContext(_ context.Context, _ string) (driver.Stmt, err
return c.stmt, nil
}
func (c *fakeConn) Close() error { return nil }
func (c *fakeConn) Close() error { return nil }
func (c *fakeConn) Begin() (driver.Tx, error) { return nil, nil }
+1 -1
View File
@@ -13,5 +13,5 @@ func (s wrappedStmt) CheckNamedValue(v *driver.NamedValue) error {
return checker.CheckNamedValue(v)
}
return defaultCheckNamedValue(v)
return driver.ErrSkip
}
+36 -11
View File
@@ -10,8 +10,10 @@ func TestWrappedStmt_CheckNamedValue(t *testing.T) {
tests := map[string]struct {
fd *fakeDriver
expected struct {
cc bool // Whether the fakeConn's CheckNamedValue was called
sc bool // Whether the fakeStmt's CheckNamedValue was called
cc bool // Whether the fakeConn's CheckNamedValue was called
sc bool // Whether the fakeStmt's CheckNamedValue was called
cci bool // Whether the fakeStmt's ColumnConverter was called
}
}{
"When both conn and stmt implement CheckNamedValue": {
@@ -23,8 +25,9 @@ func TestWrappedStmt_CheckNamedValue(t *testing.T) {
},
},
expected: struct {
cc bool
sc bool
cc bool
sc bool
cci bool
}{cc: false, sc: true},
},
"When only conn implements CheckNamedValue": {
@@ -36,8 +39,9 @@ func TestWrappedStmt_CheckNamedValue(t *testing.T) {
},
},
expected: struct {
cc bool
sc bool
cc bool
sc bool
cci bool
}{cc: true, sc: false},
},
"When only stmt implements CheckNamedValue": {
@@ -49,10 +53,25 @@ func TestWrappedStmt_CheckNamedValue(t *testing.T) {
},
},
expected: struct {
cc bool
sc bool
cc bool
sc bool
cci bool
}{cc: false, sc: true},
},
"When only stmt implements ColumnConverter": {
fd: &fakeDriver{
conn: &fakeConnWithoutCheckNamedValue{
fakeConn: fakeConn{
stmt: &fakeStmtWithColumnConverter{},
},
},
},
expected: struct {
cc bool
sc bool
cci bool
}{cci: true},
},
"When both stmt do not implement CheckNamedValue": {
fd: &fakeDriver{
conn: &fakeConnWithoutCheckNamedValue{
@@ -62,8 +81,9 @@ func TestWrappedStmt_CheckNamedValue(t *testing.T) {
},
},
expected: struct {
cc bool
sc bool
cc bool
sc bool
cci bool
}{cc: false, sc: false},
},
}
@@ -91,8 +111,9 @@ func TestWrappedStmt_CheckNamedValue(t *testing.T) {
}
conn := reflect.ValueOf(test.fd.conn).Elem()
sc := conn.FieldByName("stmt").Elem().Elem().FieldByName("called").Bool()
sc := conn.FieldByName("stmt").Elem().Elem().FieldByName("checkNamedValueCalled").Bool()
cc := conn.FieldByName("called").Bool()
cci := conn.FieldByName("stmt").Elem().Elem().FieldByName("columnConverterCalled").Bool()
if test.expected.sc != sc {
t.Errorf("sc mismatch.\n got: %#v\nwant: %#v", sc, test.expected.sc)
@@ -101,6 +122,10 @@ func TestWrappedStmt_CheckNamedValue(t *testing.T) {
if test.expected.cc != cc {
t.Errorf("cc mismatch.\n got: %#v\nwant: %#v", cc, test.expected.cc)
}
if test.expected.cci != cci {
t.Errorf("columnConverterCalled mismatch.\n got: %#v\nwant: %#v", cci, test.expected.cci)
}
})
}
}