mirror of
https://github.com/ngrok/sqlmw.git
synced 2026-06-16 16:54:29 +00:00
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:
+1
-6
@@ -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
@@ -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
@@ -13,5 +13,5 @@ func (s wrappedStmt) CheckNamedValue(v *driver.NamedValue) error {
|
||||
return checker.CheckNamedValue(v)
|
||||
}
|
||||
|
||||
return defaultCheckNamedValue(v)
|
||||
return driver.ErrSkip
|
||||
}
|
||||
|
||||
+36
-11
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user