diff --git a/CHANGELOG.md b/CHANGELOG.md index e3c6f0d8..c4e7d33f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -16,6 +16,7 @@ The next release will require at least [Go 1.26]. - Support testing of [Go 1.27]. (#650) - Add `WithSpanErrorAttributesGetter` option to set additional attributes (e.g., `db.response.status_code`) on spans when an operation returns an error. (#651) - Add `SpanOptions.RowsChildOfQuery` to create `sql.rows` spans as children of the `sql.conn.query` or `sql.stmt.query` span that produced them, so concurrent queries can be correlated with their result iteration. (#652) +- Support `driver.RowsColumnScanner` on [Go 1.27]. (#649) ### Changed diff --git a/rows.go b/rows.go index 5fc47612..55bca29c 100644 --- a/rows.go +++ b/rows.go @@ -51,7 +51,7 @@ func rowsContext(ctx, queryCtx context.Context, cfg config) context.Context { return ctx } -func newRows(ctx context.Context, rows driver.Rows, cfg config) *otRows { +func newRows(ctx context.Context, rows driver.Rows, cfg config) driver.Rows { var span trace.Span method := MethodRows @@ -61,12 +61,12 @@ func newRows(ctx context.Context, rows driver.Rows, cfg config) *otRows { _, span = createSpan(ctx, cfg, method, false, "", nil) } - return &otRows{ + return wrapRowsColumnScanner(&otRows{ Rows: rows, span: span, cfg: cfg, onClose: onClose, - } + }) } // HasNextResultSet calls the implements the driver.RowsNextResultSet for otRows. @@ -153,15 +153,23 @@ func (r otRows) Close() (err error) { } func (r otRows) Next(dest []driver.Value) (err error) { + r.beforeNext() + + err = r.Rows.Next(dest) + r.afterNext(err) + + return +} + +func (r otRows) beforeNext() { if r.cfg.SpanOptions.RowsNext && r.span != nil { r.span.AddEvent(string(EventRowsNext)) } +} - err = r.Rows.Next(dest) +func (r otRows) afterNext(err error) { // io.EOF is not an error. It is expected to happen during iteration. if err != nil && !errors.Is(err, io.EOF) { recordSpanError(r.span, r.cfg, err) } - - return } diff --git a/rows_go1.27.go b/rows_go1.27.go new file mode 100644 index 00000000..07a39e42 --- /dev/null +++ b/rows_go1.27.go @@ -0,0 +1,57 @@ +// Copyright Sam Xie +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build go1.27 + +package otelsql + +import "database/sql/driver" + +var _ driver.RowsColumnScanner = (*otRowsColumnScanner)(nil) + +type otRowsColumnScanner struct { + *otRows + + scanner driver.RowsColumnScanner +} + +func wrapRowsColumnScanner(rows *otRows) driver.Rows { + scanner, ok := rows.Rows.(driver.RowsColumnScanner) + if !ok { + return rows + } + + return &otRowsColumnScanner{ + otRows: rows, + scanner: scanner, + } +} + +func (r otRowsColumnScanner) NextRow() (err error) { + r.beforeNext() + + err = r.scanner.NextRow() + r.afterNext(err) + + return +} + +func (r otRowsColumnScanner) ScanColumn(scanCtx driver.ScanContext, index int, dest any) (err error) { + err = r.scanner.ScanColumn(scanCtx, index, dest) + if err != nil { + recordSpanError(r.span, r.cfg, err) + } + + return +} diff --git a/rows_go1.27_test.go b/rows_go1.27_test.go new file mode 100644 index 00000000..54e987cb --- /dev/null +++ b/rows_go1.27_test.go @@ -0,0 +1,287 @@ +// Copyright Sam Xie +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build go1.27 + +package otelsql + +import ( + "context" + "database/sql" + "database/sql/driver" + "errors" + "io" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/codes" +) + +type mockRowsColumnScanner struct { + *mockRows + + nextRowErr, scanColumnErr error + nextRowCount, scanColumnCount int + scanColumnIndex int + scanColumnDest any + scanColumnValue driver.Value +} + +var _ driver.RowsColumnScanner = (*mockRowsColumnScanner)(nil) + +func (m *mockRowsColumnScanner) NextRow() error { + m.nextRowCount++ + + return m.nextRowErr +} + +func (m *mockRowsColumnScanner) Columns() []string { + return []string{"value"} +} + +func (m *mockRowsColumnScanner) ScanColumn(scanCtx driver.ScanContext, index int, dest any) error { + m.scanColumnCount++ + m.scanColumnIndex = index + m.scanColumnDest = dest + + if m.scanColumnErr != nil { + return m.scanColumnErr + } + + if m.scanColumnValue != nil { + return sql.ConvertAssign(scanCtx, dest, m.scanColumnValue) + } + + return nil +} + +type mockRowsColumnScannerConn struct { + *mockConn + + rows driver.Rows +} + +var _ driver.QueryerContext = (*mockRowsColumnScannerConn)(nil) + +func (m *mockRowsColumnScannerConn) QueryContext( + context.Context, string, []driver.NamedValue, +) (driver.Rows, error) { + return m.rows, nil +} + +type singleConnConnector struct { + conn driver.Conn +} + +var _ driver.Connector = (*singleConnConnector)(nil) + +func (c *singleConnConnector) Connect(context.Context) (driver.Conn, error) { + return c.conn, nil +} + +func (c *singleConnConnector) Driver() driver.Driver { + return newMockDriver(false) +} + +func TestNewRows_RowsColumnScanner(t *testing.T) { + t.Run("does not expose unsupported interface", func(t *testing.T) { + rows := newRows(t.Context(), newMockRows(false), newConfig()) + + _, ok := rows.(driver.RowsColumnScanner) + assert.False(t, ok) + }) + + t.Run("forwards row scanning", func(t *testing.T) { + ctx, sr, tracer, _ := prepareTraces(false) + dest := new(string) + mr := &mockRowsColumnScanner{mockRows: newMockRows(false)} + + cfg := newConfig() + cfg.Tracer = tracer + cfg.SpanOptions.RowsNext = true + + rows := newRows(ctx, mr, cfg) + scanner, ok := rows.(driver.RowsColumnScanner) + require.True(t, ok) + + require.NoError(t, scanner.NextRow()) + require.NoError(t, scanner.ScanColumn(driver.ScanContext{}, 2, dest)) + + assert.Equal(t, 1, mr.nextRowCount) + assert.Equal(t, 0, mr.nextCount) + assert.Equal(t, 1, mr.scanColumnCount) + assert.Equal(t, 2, mr.scanColumnIndex) + assert.Same(t, dest, mr.scanColumnDest) + + spans := sr.Started() + require.Len(t, spans, 2) + assert.Len(t, spans[1].Events(), 1) + assert.Equal(t, codes.Unset, spans[1].Status().Code) + }) + + t.Run("records NextRow errors", func(t *testing.T) { + ctx, sr, tracer, _ := prepareTraces(false) + nextRowErr := errors.New("next row") + mr := &mockRowsColumnScanner{ + mockRows: newMockRows(false), + nextRowErr: nextRowErr, + } + + cfg := newConfig() + cfg.Tracer = tracer + + rows := newRows(ctx, mr, cfg) + scanner, ok := rows.(driver.RowsColumnScanner) + require.True(t, ok) + require.ErrorIs(t, scanner.NextRow(), nextRowErr) + + spans := sr.Started() + require.Len(t, spans, 2) + assert.Equal(t, codes.Error, spans[1].Status().Code) + assert.Len(t, spans[1].Events(), 1) + }) + + t.Run("records ScanColumn errors", func(t *testing.T) { + ctx, sr, tracer, _ := prepareTraces(false) + scanColumnErr := errors.New("scan column") + mr := &mockRowsColumnScanner{ + mockRows: newMockRows(false), + scanColumnErr: scanColumnErr, + } + + cfg := newConfig() + cfg.Tracer = tracer + + rows := newRows(ctx, mr, cfg) + scanner, ok := rows.(driver.RowsColumnScanner) + require.True(t, ok) + require.ErrorIs(t, scanner.ScanColumn(driver.ScanContext{}, 0, new(string)), scanColumnErr) + + spans := sr.Started() + require.Len(t, spans, 2) + assert.Equal(t, codes.Error, spans[1].Status().Code) + require.Len(t, spans[1].Events(), 1) + assert.Equal(t, "exception", spans[1].Events()[0].Name) + }) + + t.Run("does not record NextRow EOF as an error", func(t *testing.T) { + ctx, sr, tracer, _ := prepareTraces(false) + mr := &mockRowsColumnScanner{ + mockRows: newMockRows(false), + nextRowErr: io.EOF, + } + + cfg := newConfig() + cfg.Tracer = tracer + cfg.SpanOptions.RowsNext = true + + rows := newRows(ctx, mr, cfg) + scanner, ok := rows.(driver.RowsColumnScanner) + require.True(t, ok) + require.ErrorIs(t, scanner.NextRow(), io.EOF) + + spans := sr.Started() + require.Len(t, spans, 2) + assert.Equal(t, codes.Unset, spans[1].Status().Code) + assert.Len(t, spans[1].Events(), 1) + }) +} + +func TestRowsColumnScanner_DatabaseSQL(t *testing.T) { + ctx, sr, tracer, _ := prepareTraces(false) + mr := &mockRowsColumnScanner{ + mockRows: newMockRows(false), + scanColumnValue: "scanned", + } + + cfg := newConfig() + cfg.Tracer = tracer + cfg.SpanOptions.RowsNext = true + + conn := &mockRowsColumnScannerConn{ + mockConn: newMockConn(false), + rows: mr, + } + db := sql.OpenDB(&singleConnConnector{conn: newConn(conn, cfg)}) + + t.Cleanup(func() { + require.NoError(t, db.Close()) + }) + + rows, err := db.QueryContext(ctx, testQueryString) + + require.NoError(t, err) + defer func() { + require.NoError(t, rows.Close()) + }() + + require.True(t, rows.Next()) + + var value string + require.NoError(t, rows.Scan(&value)) + require.NoError(t, rows.Err()) + + assert.Equal(t, "scanned", value) + assert.Equal(t, 1, mr.nextRowCount) + assert.Equal(t, 0, mr.nextCount) + assert.Equal(t, 1, mr.scanColumnCount) + assert.Equal(t, 0, mr.scanColumnIndex) + assert.Same(t, &value, mr.scanColumnDest) + + spans := sr.Started() + require.Len(t, spans, 3) + assert.Len(t, spans[2].Events(), 1) + assert.Equal(t, codes.Unset, spans[2].Status().Code) +} + +func TestRowsColumnScanner_DatabaseSQLScanError(t *testing.T) { + ctx, sr, tracer, _ := prepareTraces(false) + scanColumnErr := errors.New("scan column") + mr := &mockRowsColumnScanner{ + mockRows: newMockRows(false), + scanColumnErr: scanColumnErr, + } + + cfg := newConfig() + cfg.Tracer = tracer + cfg.SpanOptions.RowsNext = true + + conn := &mockRowsColumnScannerConn{ + mockConn: newMockConn(false), + rows: mr, + } + db := sql.OpenDB(&singleConnConnector{conn: newConn(conn, cfg)}) + + t.Cleanup(func() { + require.NoError(t, db.Close()) + }) + + rows, err := db.QueryContext(ctx, testQueryString) + + require.NoError(t, err) + defer func() { + require.NoError(t, rows.Close()) + }() + + require.True(t, rows.Next()) + require.ErrorIs(t, rows.Scan(new(string)), scanColumnErr) + require.NoError(t, rows.Err()) + + spans := sr.Started() + require.Len(t, spans, 3) + assert.Equal(t, codes.Error, spans[2].Status().Code) + require.Len(t, spans[2].Events(), 2) + assert.Equal(t, "exception", spans[2].Events()[1].Name) +} diff --git a/rows_pre_go1.27.go b/rows_pre_go1.27.go new file mode 100644 index 00000000..06a0ecb3 --- /dev/null +++ b/rows_pre_go1.27.go @@ -0,0 +1,23 @@ +// Copyright Sam Xie +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build !go1.27 + +package otelsql + +import "database/sql/driver" + +func wrapRowsColumnScanner(rows *otRows) driver.Rows { + return rows +} diff --git a/rows_test.go b/rows_test.go index 639575de..5c6ca11f 100644 --- a/rows_test.go +++ b/rows_test.go @@ -309,7 +309,9 @@ func TestNewRows(t *testing.T) { attributesGetter: tc.attributesGetter, }) - assert.Equal(t, mr, rows.Rows) + otelRows, ok := rows.(*otRows) + require.True(t, ok) + assert.Equal(t, mr, otelRows.Rows) }) } })