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) + } }) } }