diff --git a/drivers/mssql/internal/cdc.go b/drivers/mssql/internal/cdc.go index 705914a01..1e9a342bb 100644 --- a/drivers/mssql/internal/cdc.go +++ b/drivers/mssql/internal/cdc.go @@ -357,35 +357,25 @@ func (m *MSSQL) fetchTableChangesInLSNRange(ctx context.Context, stream types.St // Query CDC rows for this capture instance between the two LSNs. query := jdbc.MSSQLCDCGetChangesQuery(capture.instanceName) - rows, err := m.client.QueryContext(ctx, query, fromLSNBytes, toLSNBytes) - if err != nil { - return fmt.Errorf("failed to query MSSQL CDC changes: %s", err) - } - defer rows.Close() - - for rows.Next() { - // Use MapScan to properly convert data types including binary types - // TODO: check if we can use MapScanConcurrent for mssql - // rowBytes is the after-image data-column byte sum (excludes __$* metadata columns), attached to the emitted change below. - record := make(map[string]interface{}) - rowBytes, err := jdbc.MapScan(rows, record, m.dataTypeConverter, mssqlCDCColumnSizer) - if err != nil { - return fmt.Errorf("failed to scan MSSQL CDC row: %s", err) - } + setter := jdbc.NewReader(ctx, query, func(ctx context.Context, q string, args ...any) (*sql.Rows, error) { + return m.client.QueryContext(ctx, q, args...) + }, fromLSNBytes, toLSNBytes) + return jdbc.MapScanConcurrent(setter, m.dataTypeConverter, func(ctx context.Context, record map[string]any, rowBytes int64) error { // Determine operation type from SQL Server CDC operation codes. // For updates, CDC emits "before" (3) and "after" (4); we skip "before". var operationType string if val, ok := record["__$operation"]; ok { var opCode = val.(int32) if opCode == 3 { - continue + return nil } operationType = operationTypeFromCDCCode(opCode) } extraColumns := map[string]any{ CDCStartLSN: fmt.Sprintf("%x", record["__$start_lsn"]), - CDCSeqVal: fmt.Sprintf("%x", record["__$seqval"])} + CDCSeqVal: fmt.Sprintf("%x", record["__$seqval"]), + } // Remove metadata columns delete(record, "__$operation") delete(record, "__$start_lsn") @@ -397,12 +387,8 @@ func (m *MSSQL) fetchTableChangesInLSNRange(ctx context.Context, stream types.St extraColumns, rowBytes)); err != nil { return fmt.Errorf("failed to process MSSQL CDC change: %s", err) } - } - - if err := rows.Err(); err != nil { - return err - } - return nil + return nil + }, mssqlCDCColumnSizer) } func (m *MSSQL) currentMaxLSN(ctx context.Context) (string, error) { diff --git a/drivers/mssql/internal/incremental.go b/drivers/mssql/internal/incremental.go index d8e34e0b1..51b386042 100644 --- a/drivers/mssql/internal/incremental.go +++ b/drivers/mssql/internal/incremental.go @@ -22,27 +22,15 @@ func (m *MSSQL) StreamIncrementalChanges(ctx context.Context, stream types.Strea return fmt.Errorf("failed to build incremental condition: %s", err) } - var rows *sql.Rows - rows, err = m.client.QueryContext(ctx, incrementalQuery, queryArgs...) - if err != nil { - return fmt.Errorf("failed to execute incremental query: %s", err) - } - defer rows.Close() - - // Scan rows and process - for rows.Next() { - record := make(types.Record) - rowBytes, err := jdbc.MapScan(rows, record, m.dataTypeConverter, mssqlColumnSizer) - if err != nil { - return fmt.Errorf("failed to scan record: %s", err) - } + setter := jdbc.NewReader(ctx, incrementalQuery, func(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + return m.client.QueryContext(ctx, query, args...) + }, queryArgs...) - if err := processFn(ctx, record, rowBytes); err != nil { - return fmt.Errorf("process error: %s", err) - } + if err := jdbc.MapScanConcurrent(setter, m.dataTypeConverter, processFn, mssqlColumnSizer); err != nil { + return fmt.Errorf("incremental process error: %s", err) } - return rows.Err() + return nil } func (m *MSSQL) FetchMaxCursorValues(ctx context.Context, stream types.StreamInterface) (any, any, error) { diff --git a/drivers/mysql/internal/incremental.go b/drivers/mysql/internal/incremental.go index c4fbc6979..1bb2bb350 100644 --- a/drivers/mysql/internal/incremental.go +++ b/drivers/mysql/internal/incremental.go @@ -22,27 +22,15 @@ func (m *MySQL) StreamIncrementalChanges(ctx context.Context, stream types.Strea return fmt.Errorf("failed to build incremental condition: %s", err) } - var rows *sql.Rows - rows, err = m.client.QueryContext(ctx, incrementalQuery, queryArgs...) - if err != nil { - return fmt.Errorf("failed to execute incremental query: %s", err) - } - defer rows.Close() - - // Scan rows and process - for rows.Next() { - record := make(types.Record) - rowBytes, err := jdbc.MapScan(rows, record, m.dataTypeConverter, mysqlColumnSizer) - if err != nil { - return fmt.Errorf("failed to scan record: %s", err) - } + setter := jdbc.NewReader(ctx, incrementalQuery, func(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + return m.client.QueryContext(ctx, query, args...) + }, queryArgs...) - if err := processFn(ctx, record, rowBytes); err != nil { - return fmt.Errorf("process error: %s", err) - } + if err := jdbc.MapScanConcurrent(setter, m.dataTypeConverter, processFn, mysqlColumnSizer); err != nil { + return fmt.Errorf("incremental process error: %s", err) } - return rows.Err() + return nil } func (m *MySQL) FetchMaxCursorValues(ctx context.Context, stream types.StreamInterface) (any, any, error) { diff --git a/drivers/oracle/internal/incremental.go b/drivers/oracle/internal/incremental.go index 301193378..20ccc1b52 100644 --- a/drivers/oracle/internal/incremental.go +++ b/drivers/oracle/internal/incremental.go @@ -2,6 +2,7 @@ package driver import ( "context" + "database/sql" "fmt" "strings" "time" @@ -26,24 +27,14 @@ func (o *Oracle) StreamIncrementalChanges(ctx context.Context, stream types.Stre return fmt.Errorf("failed to build incremental condition: %s", err) } - rows, err := o.client.QueryContext(ctx, incrementalQuery, queryArgs...) - if err != nil { - return fmt.Errorf("failed to execute incremental query: %s", err) - } - defer rows.Close() - - for rows.Next() { - record := make(types.Record) - rowBytes, err := jdbc.MapScan(rows, record, o.dataTypeConverter, oracleColumnSizer) - if err != nil { - return fmt.Errorf("failed to scan record: %s", err) - } + setter := jdbc.NewReader(ctx, incrementalQuery, func(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + return o.client.QueryContext(ctx, query, args...) + }, queryArgs...) - if err := processFn(ctx, record, rowBytes); err != nil { - return fmt.Errorf("process error: %s", err) - } + if err := jdbc.MapScanConcurrent(setter, o.dataTypeConverter, processFn, oracleColumnSizer); err != nil { + return fmt.Errorf("incremental process error: %s", err) } - return rows.Err() + return nil } func (o *Oracle) FetchMaxCursorValues(ctx context.Context, stream types.StreamInterface) (any, any, error) { diff --git a/drivers/postgres/internal/incremental.go b/drivers/postgres/internal/incremental.go index efc88857a..8742e6230 100644 --- a/drivers/postgres/internal/incremental.go +++ b/drivers/postgres/internal/incremental.go @@ -2,6 +2,7 @@ package driver import ( "context" + "database/sql" "fmt" "github.com/datazip-inc/olake/constants" @@ -21,25 +22,15 @@ func (p *Postgres) StreamIncrementalChanges(ctx context.Context, stream types.St return fmt.Errorf("failed to build incremental condition: %s", err) } - rows, err := p.client.QueryContext(ctx, incrementalQuery, queryArgs...) - if err != nil { - return fmt.Errorf("failed to execute incremental query: %s", err) - } - defer rows.Close() - - for rows.Next() { - record := make(types.Record) - rowBytes, err := jdbc.MapScan(rows, record, p.dataTypeConverter, pgColumnSizer) - if err != nil { - return fmt.Errorf("failed to scan record: %s", err) - } + setter := jdbc.NewReader(ctx, incrementalQuery, func(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + return p.client.QueryContext(ctx, query, args...) + }, queryArgs...) - if err := processFn(ctx, record, rowBytes); err != nil { - return fmt.Errorf("process error: %s", err) - } + if err := jdbc.MapScanConcurrent(setter, p.dataTypeConverter, processFn, pgColumnSizer); err != nil { + return fmt.Errorf("incremental process error: %s", err) } - return rows.Err() + return nil } func (p *Postgres) FetchMaxCursorValues(ctx context.Context, stream types.StreamInterface) (any, any, error) { diff --git a/pkg/jdbc/reader.go b/pkg/jdbc/reader.go index 4a6682dee..20e567e9d 100644 --- a/pkg/jdbc/reader.go +++ b/pkg/jdbc/reader.go @@ -43,6 +43,7 @@ func (o *Reader[T]) Capture(onCapture func(T) error) error { if err != nil { return err } + defer func() { _ = rows.Close() }() for rows.Next() { err := onCapture(rows) @@ -87,48 +88,6 @@ func normalizeDataTypeAndConvert(rawData any, colType *sql.ColumnType, converter return conv, nil } -// TODO: Use MapScanConcurrent instead of MapScan for incremental as well -// -// MapScan scans the current row into dest and returns the row's source-DB byte -// size. columnSizer maps a SQL column type to a function that sizes one raw -// (pre-conversion) value of that column; the size is summed inline in the scan -// loop. NULL values carry no data and are skipped (0 bytes). -func MapScan(rows *sql.Rows, dest map[string]any, converter func(value interface{}, columnType string) (interface{}, error), columnSizer func(colType *sql.ColumnType) func(v any) int64) (int64, error) { - columns, colTypes, err := getColumnMetadata(rows) - if err != nil { - return 0, err - } - - scanValues := make([]any, len(columns)) - for i := range scanValues { - scanValues[i] = new(any) // Allocate pointers for scanning - } - - if err := rows.Scan(scanValues...); err != nil { - return 0, err - } - - var rowBytes int64 - for i, col := range columns { - rawData := *(scanValues[i].(*any)) // Dereference pointer before storing - // If rawData is nil, no byte is added - if rawData != nil { - rowBytes += columnSizer(colTypes[i])(rawData) - } - if converter != nil { - conv, err := normalizeDataTypeAndConvert(rawData, colTypes[i], converter) - if err != nil { - return 0, err - } - dest[col] = conv - } else { - dest[col] = rawData - } - } - - return rowBytes, nil -} - // MapScanConcurrent scans rows concurrently using a producer/consumer pattern. // columnSizer maps a SQL column type to a function that sizes one raw // (pre-conversion) value of that column. Since column types are constant for the @@ -172,8 +131,7 @@ func MapScanConcurrent(setter *Reader[*sql.Rows], converter func(value interface select { case <-ctx.Done(): - // If the processor failed, errgroup cancels the ctx; return nil so the original error wins. - return nil + return ctx.Err() case valuesCh <- vals: return nil } diff --git a/pkg/jdbc/reader_test.go b/pkg/jdbc/reader_test.go new file mode 100644 index 000000000..ac7ee0973 --- /dev/null +++ b/pkg/jdbc/reader_test.go @@ -0,0 +1,383 @@ +package jdbc + +import ( + "context" + "database/sql" + "database/sql/driver" + "errors" + "fmt" + "io" + "sync" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func init() { + sql.Register("olake_jdbc_test", mockDriver{}) +} + +type mockDriver struct{} + +func (mockDriver) Open(name string) (driver.Conn, error) { + return &mockConn{rows: parseMockDSN(name)}, nil +} + +type mockConn struct { + rows *mockDriverRows +} + +func (c *mockConn) Prepare(_ string) (driver.Stmt, error) { + return &mockStmt{rows: c.rows}, nil +} + +func (c *mockConn) Close() error { return nil } + +func (c *mockConn) Begin() (driver.Tx, error) { + return nil, errors.New("transactions not supported in mock driver") +} + +type mockStmt struct { + rows *mockDriverRows +} + +func (s *mockStmt) Close() error { return nil } + +func (s *mockStmt) NumInput() int { return 0 } + +func (s *mockStmt) Exec([]driver.Value) (driver.Result, error) { + return nil, errors.New("exec not supported in mock driver") +} + +func (s *mockStmt) Query([]driver.Value) (driver.Rows, error) { + return s.rows.clone(), nil +} + +type mockDriverRows struct { + mu sync.Mutex + cols []string + data [][]driver.Value + idx int + closed bool + rowsErr error + closeCount *atomic.Int32 +} + +func (r *mockDriverRows) Columns() []string { return r.cols } + +func (r *mockDriverRows) Close() error { + r.mu.Lock() + defer r.mu.Unlock() + r.closed = true + if r.closeCount != nil { + r.closeCount.Add(1) + } + return nil +} + +func (r *mockDriverRows) Next(dest []driver.Value) error { + r.mu.Lock() + defer r.mu.Unlock() + if r.rowsErr != nil && r.idx > 0 { + return r.rowsErr + } + if r.idx >= len(r.data) { + return io.EOF + } + copy(dest, r.data[r.idx]) + r.idx++ + return nil +} + +func (r *mockDriverRows) clone() *mockDriverRows { + r.mu.Lock() + defer r.mu.Unlock() + return &mockDriverRows{ + cols: append([]string(nil), r.cols...), + data: append([][]driver.Value(nil), r.data...), + rowsErr: r.rowsErr, + closeCount: r.closeCount, + } +} + +type mockDSN struct { + cols []string + data [][]driver.Value + rowsErr error + closeCount *atomic.Int32 +} + +var mockDSNRegistry sync.Map + +func registerMockDSN(dsn string, cfg mockDSN) { + mockDSNRegistry.Store(dsn, cfg) +} + +func parseMockDSN(dsn string) *mockDriverRows { + if raw, ok := mockDSNRegistry.Load(dsn); ok { + cfg := raw.(mockDSN) + return &mockDriverRows{cols: cfg.cols, data: cfg.data, rowsErr: cfg.rowsErr, closeCount: cfg.closeCount} + } + return &mockDriverRows{} +} + +func openMockDB(t *testing.T, dsn string, cfg mockDSN) (*sql.DB, *atomic.Int32) { + t.Helper() + closeCount := &atomic.Int32{} + cfg.closeCount = closeCount + registerMockDSN(dsn, cfg) + db, err := sql.Open("olake_jdbc_test", dsn) + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + return db, closeCount +} + +func usersMockDSN(t *testing.T) (string, mockDSN) { + t.Helper() + return t.Name() + "_users", mockDSN{ + cols: []string{"id", "name", "active"}, + data: [][]driver.Value{ + {int64(1), "alice", int64(1)}, + {int64(2), "bob", int64(0)}, + {int64(3), "carol", int64(1)}, + }, + } +} + +func testColumnSizer(colType *sql.ColumnType) func(v any) int64 { + _ = colType + return func(v any) int64 { + if v == nil { + return 0 + } + return int64(len(fmt.Sprint(v))) + } +} + +func identityConverter(value interface{}, _ string) (interface{}, error) { + return value, nil +} + +// fakeIter is an in-memory Iterable that does not auto-close on exhaustion, +// unlike database/sql.Rows. Used to assert Capture always calls Close(). +type fakeIter struct { + remaining int + iterErr error + closed atomic.Int32 +} + +func newFakeIter(count int, iterErr error) *fakeIter { + return &fakeIter{remaining: count, iterErr: iterErr} +} + +func (f *fakeIter) Next() bool { + if f.remaining <= 0 { + return false + } + f.remaining-- + return true +} + +func (f *fakeIter) Err() error { + return f.iterErr +} + +func (f *fakeIter) Close() error { + f.closed.Add(1) + return nil +} + +func newFakeReader(iter *fakeIter) *Reader[*fakeIter] { + return NewReader(context.Background(), "SELECT 1", func(_ context.Context, _ string, _ ...any) (*fakeIter, error) { + return iter, nil + }) +} + +func TestReaderCapture_closesOnSuccess(t *testing.T) { + iter := newFakeIter(3, nil) + setter := newFakeReader(iter) + + var rowCount int + err := setter.Capture(func(_ *fakeIter) error { + rowCount++ + return nil + }) + require.NoError(t, err) + assert.Equal(t, 3, rowCount) + assert.Equal(t, int32(1), iter.closed.Load(), "Capture must close the iterable on success") +} + +func TestReaderCapture_closesOnCaptureError(t *testing.T) { + iter := newFakeIter(3, nil) + setter := newFakeReader(iter) + + captureErr := errors.New("capture failed") + err := setter.Capture(func(_ *fakeIter) error { + return captureErr + }) + require.ErrorIs(t, err, captureErr) + assert.Equal(t, int32(1), iter.closed.Load(), "Capture must close the iterable on early return") +} + +func TestReaderCapture_closesOnIterableErr(t *testing.T) { + iterErr := errors.New("iteration failed") + iter := newFakeIter(1, iterErr) + setter := newFakeReader(iter) + + err := setter.Capture(func(_ *fakeIter) error { return nil }) + require.ErrorIs(t, err, iterErr) + assert.Equal(t, int32(1), iter.closed.Load(), "Capture must close the iterable when Err() fails") +} + +func TestReaderCapture_doesNotLeakIterator(t *testing.T) { + ctx := context.Background() + var open atomic.Int32 + for i := 0; i < 5; i++ { + iter := newFakeIter(1, nil) + setter := NewReader(ctx, "SELECT 1", func(_ context.Context, _ string, _ ...any) (*fakeIter, error) { + if open.Load() > 0 { + return nil, errors.New("previous iterator not closed") + } + open.Store(1) + return iter, nil + }) + err := setter.Capture(func(_ *fakeIter) error { return nil }) + require.NoError(t, err, "iteration %d should not leak the only open iterator", i) + assert.Equal(t, int32(1), iter.closed.Load()) + open.Store(0) + } +} + +func TestReaderCapture_rejectsQueryWithSemicolon(t *testing.T) { + ctx := context.Background() + setter := NewReader(ctx, "SELECT 1;", func(_ context.Context, _ string, _ ...any) (*sql.Rows, error) { + return nil, errors.New("exec must not be called for invalid query") + }) + + err := setter.Capture(func(_ *sql.Rows) error { return nil }) + require.Error(t, err) + assert.Contains(t, err.Error(), "ends with ';'") +} + +func TestMapScanConcurrent_emitsAllRows(t *testing.T) { + dsn, cfg := usersMockDSN(t) + db, _ := openMockDB(t, dsn, cfg) + + ctx := context.Background() + setter := NewReader(ctx, "SELECT id, name, active FROM users", func(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + return db.QueryContext(ctx, query, args...) + }) + + var records []map[string]any + err := MapScanConcurrent(setter, identityConverter, func(_ context.Context, record map[string]any, sourceBytes int64) error { + records = append(records, record) + assert.Positive(t, sourceBytes) + return nil + }, testColumnSizer) + require.NoError(t, err) + + require.Len(t, records, 3) + assert.Equal(t, int64(1), records[0]["id"]) + assert.Equal(t, "alice", records[0]["name"]) + assert.Equal(t, int64(2), records[1]["id"]) + assert.Equal(t, "bob", records[1]["name"]) +} + +func TestMapScanConcurrent_propagatesCallbackError(t *testing.T) { + dsn, cfg := usersMockDSN(t) + db, _ := openMockDB(t, dsn, cfg) + + ctx := context.Background() + setter := NewReader(ctx, "SELECT id FROM users", func(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + return db.QueryContext(ctx, query, args...) + }) + + callbackErr := errors.New("downstream failure") + err := MapScanConcurrent(setter, identityConverter, func(_ context.Context, _ map[string]any, _ int64) error { + return callbackErr + }, testColumnSizer) + require.ErrorIs(t, err, callbackErr) +} + +func TestMapScanConcurrent_skipsRowsViaNilCallbackReturn(t *testing.T) { + dsn, cfg := usersMockDSN(t) + db, _ := openMockDB(t, dsn, cfg) + + ctx := context.Background() + setter := NewReader(ctx, "SELECT id, name FROM users", func(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + return db.QueryContext(ctx, query, args...) + }) + + var emitted []int64 + err := MapScanConcurrent(setter, identityConverter, func(_ context.Context, record map[string]any, _ int64) error { + id := record["id"].(int64) + if id == 2 { + return nil // mirrors MSSQL CDC before-image skip + } + emitted = append(emitted, id) + return nil + }, testColumnSizer) + require.NoError(t, err) + assert.Equal(t, []int64{1, 3}, emitted) +} + +func TestMapScanConcurrent_cdcMetadataColumnSizer(t *testing.T) { + dsn := t.Name() + "_cdc" + cfg := mockDSN{ + cols: []string{"__$operation", "__$start_lsn", "id", "status"}, + data: [][]driver.Value{{int64(4), []byte{1, 2}, int64(1), "shipped"}}, + } + db, _ := openMockDB(t, dsn, cfg) + + cdcColumnSizer := func(colType *sql.ColumnType) func(v any) int64 { + if colType.Name() == "__$operation" || colType.Name() == "__$start_lsn" { + return func(any) int64 { return 0 } + } + return testColumnSizer(colType) + } + + ctx := context.Background() + setter := NewReader(ctx, "SELECT __$operation, __$start_lsn, id, status FROM cdc_ct", func(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + return db.QueryContext(ctx, query, args...) + }) + + var rowBytes int64 + err := MapScanConcurrent(setter, identityConverter, func(_ context.Context, record map[string]any, bytes int64) error { + rowBytes = bytes + assert.Equal(t, int64(4), record["__$operation"]) + assert.Equal(t, int64(1), record["id"]) + return nil + }, cdcColumnSizer) + require.NoError(t, err) + assert.Positive(t, rowBytes, "only data columns should contribute to byte count") +} + +func TestMapScanConcurrent_stopsProducerOnConsumerError(t *testing.T) { + dsn := t.Name() + "_cancel" + rows := make([][]driver.Value, 100) + for i := range rows { + rows[i] = []driver.Value{int64(i + 1)} + } + cfg := mockDSN{ + cols: []string{"id"}, + data: rows, + } + db, closeCount := openMockDB(t, dsn, cfg) + + ctx := context.Background() + setter := NewReader(ctx, "SELECT id FROM users", func(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + return db.QueryContext(ctx, query, args...) + }) + + callbackErr := errors.New("downstream failure") + var emitted int + err := MapScanConcurrent(setter, identityConverter, func(_ context.Context, _ map[string]any, _ int64) error { + emitted++ + return callbackErr + }, testColumnSizer) + require.ErrorIs(t, err, callbackErr) + assert.Equal(t, 1, emitted, "consumer should stop after first row") + assert.Equal(t, int32(1), closeCount.Load(), "rows must be closed when producer is canceled") +} diff --git a/types/interface.go b/types/interface.go index 919d5ef66..e542cb23f 100644 --- a/types/interface.go +++ b/types/interface.go @@ -37,4 +37,5 @@ type StateInterface interface { type Iterable interface { Next() bool Err() error + Close() error }