From 1e77bed04fd697da9e2069f457ac39d992d74b68 Mon Sep 17 00:00:00 2001 From: Fabian Holler Date: Tue, 3 Aug 2021 09:06:41 +0200 Subject: [PATCH] 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. --- conn_go19.go | 7 +------ fakedb_test.go | 16 +++++++++++++--- stmt_go19.go | 2 +- stmt_go19_test.go | 47 ++++++++++++++++++++++++++++++++++++----------- 4 files changed, 51 insertions(+), 21 deletions(-) diff --git a/conn_go19.go b/conn_go19.go index 268cabc..4eb10e0 100755 --- a/conn_go19.go +++ b/conn_go19.go @@ -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 } diff --git a/fakedb_test.go b/fakedb_test.go index 97e6e50..fabdb53 100644 --- a/fakedb_test.go +++ b/fakedb_test.go @@ -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 } diff --git a/stmt_go19.go b/stmt_go19.go index 8442101..f54c2c9 100755 --- a/stmt_go19.go +++ b/stmt_go19.go @@ -13,5 +13,5 @@ func (s wrappedStmt) CheckNamedValue(v *driver.NamedValue) error { return checker.CheckNamedValue(v) } - return defaultCheckNamedValue(v) + return driver.ErrSkip } diff --git a/stmt_go19_test.go b/stmt_go19_test.go index f4dd198..cc9213a 100644 --- a/stmt_go19_test.go +++ b/stmt_go19_test.go @@ -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) + } }) } }