diff --git a/pkg/cdc/sql_builder.go b/pkg/cdc/sql_builder.go index 4414df8c80e83..034a2dba49eb0 100644 --- a/pkg/cdc/sql_builder.go +++ b/pkg/cdc/sql_builder.go @@ -248,19 +248,20 @@ const ( "table_name = '%s'" CDCCollectTableInfoSqlTemplate = "SELECT " + - " rel_id, " + - " relname, " + - " reldatabase_id, " + - " reldatabase, " + - " rel_createsql, " + - " account_id " + - "FROM `mo_catalog`.`mo_tables` " + + " tbl.rel_id, " + + " tbl.relname, " + + " tbl.reldatabase_id, " + + " tbl.reldatabase, " + + " tbl.rel_createsql, " + + " tbl.account_id, " + + " tbl.`constraint` " + + "FROM `mo_catalog`.`mo_tables` tbl " + "WHERE " + - " account_id IN (%s) " + + " tbl.account_id IN (%s) " + "%s" + "%s" + - " AND relkind = '%s' " + - " AND reldatabase NOT IN (%s)" + " AND tbl.relkind = '%s' " + + " AND tbl.reldatabase NOT IN (%s)" CDCInsertMOISCPLogSqlTemplate = `REPLACE INTO mo_catalog.mo_iscp_log (` + `account_id,` + `table_id,` + @@ -475,6 +476,7 @@ var CDCSQLTemplates = [CDCSqlTemplateCount]struct { "reldatabase", "rel_createsql", "account_id", + "constraint", }, }, CDCGetWatermarkWhereSqlTemplate_Idx: { @@ -1071,13 +1073,13 @@ func (b cdcSQLBuilder) CollectTableInfoSQL(accountIDs string, dbNames string, ta if dbNames == "*" { return "" } - return " AND reldatabase IN (" + dbNames + ") " + return " AND tbl.reldatabase IN (" + dbNames + ") " }(), func() string { if tableNames == "*" { return "" } - return " AND relname IN (" + tableNames + ") " + return " AND tbl.relname IN (" + tableNames + ") " }(), catalog.SystemOrdinaryRel, AddSingleQuotesJoin(catalog.SystemDatabases), diff --git a/pkg/cdc/table_change_stream.go b/pkg/cdc/table_change_stream.go index 58e29e70a831e..61f3803747cc4 100644 --- a/pkg/cdc/table_change_stream.go +++ b/pkg/cdc/table_change_stream.go @@ -587,6 +587,10 @@ func (s *TableChangeStream) cleanup(ctx context.Context) { ) defer s.wg.Done() defer func() { + // Keep ownership until cleanup finishes so a replacement reader cannot + // start while this stream is still closing its sinker/watermark state. + s.runningReaders.CompareAndDelete(s.runningReaderKey, s) + // Decrement table stream state gauge on cleanup if s.progressTracker != nil { state, _ := s.progressTracker.GetState() @@ -603,9 +607,6 @@ func (s *TableChangeStream) cleanup(ctx context.Context) { ) }() - // Remove from running readers - s.runningReaders.Delete(s.runningReaderKey) - // Remove watermark cache removeStart := time.Now() if err := s.watermarkUpdater.RemoveCachedWM(ctx, s.watermarkKey); err != nil { diff --git a/pkg/cdc/table_change_stream_test.go b/pkg/cdc/table_change_stream_test.go index 093fb2e6f17a1..dd80cfde2e9f9 100644 --- a/pkg/cdc/table_change_stream_test.go +++ b/pkg/cdc/table_change_stream_test.go @@ -378,6 +378,90 @@ func TestTableChangeStream_Run_DuplicateReader(t *testing.T) { } } +func TestTableChangeStreamCleanupDoesNotDeleteReplacedReader(t *testing.T) { + runningReaders := &sync.Map{} + + key := "db1.t1" + oldStream := &TableChangeStream{ + accountId: 1, + taskId: "task1", + tableInfo: &DbTableInfo{SourceDbName: "db1", SourceTblName: "t1", SourceTblId: 1}, + sinker: newTableStreamRecordingSinker(), + watermarkUpdater: newWatermarkUpdaterStub(), + watermarkKey: &WatermarkKey{AccountId: 1, TaskId: "task1", DBName: "db1", TableName: "t1"}, + runningReaders: runningReaders, + runningReaderKey: key, + progressTracker: nil, + watermarkStallThreshold: defaultWatermarkStallThreshold, + } + newStream := &TableChangeStream{ + accountId: 1, + taskId: "task1", + tableInfo: &DbTableInfo{SourceDbName: "db1", SourceTblName: "t1", SourceTblId: 1}, + sinker: newTableStreamRecordingSinker(), + watermarkUpdater: newWatermarkUpdaterStub(), + watermarkKey: &WatermarkKey{AccountId: 1, TaskId: "task1", DBName: "db1", TableName: "t1"}, + runningReaders: runningReaders, + runningReaderKey: key, + } + runningReaders.Store(key, oldStream) + runningReaders.Store(key, newStream) + + oldStream.wg.Add(1) + oldStream.cleanup(context.Background()) + + stored, ok := runningReaders.Load(key) + require.True(t, ok, "replaced reader ownership should remain") + require.Same(t, newStream, stored, "old cleanup must not delete a newer reader") +} + +func TestTableChangeStreamCleanupKeepsOwnershipUntilCloseFinishes(t *testing.T) { + runningReaders := &sync.Map{} + + key := "db1.t1" + sinker := newBlockingCloseSinker() + stream := &TableChangeStream{ + accountId: 1, + taskId: "task1", + tableInfo: &DbTableInfo{SourceDbName: "db1", SourceTblName: "t1", SourceTblId: 1}, + sinker: sinker, + watermarkUpdater: newWatermarkUpdaterStub(), + watermarkKey: &WatermarkKey{AccountId: 1, TaskId: "task1", DBName: "db1", TableName: "t1"}, + runningReaders: runningReaders, + runningReaderKey: key, + progressTracker: nil, + watermarkStallThreshold: defaultWatermarkStallThreshold, + } + runningReaders.Store(key, stream) + + stream.wg.Add(1) + cleanupDone := make(chan struct{}) + go func() { + stream.cleanup(context.Background()) + close(cleanupDone) + }() + + select { + case <-sinker.closeStarted: + case <-time.After(time.Second): + t.Fatal("expected cleanup to reach sinker close") + } + + stored, ok := runningReaders.Load(key) + require.True(t, ok, "reader ownership should remain while cleanup is still closing") + require.Same(t, stream, stored, "old stream should keep ownership until cleanup finishes") + + close(sinker.unblockClose) + select { + case <-cleanupDone: + case <-time.After(time.Second): + t.Fatal("cleanup did not finish after unblocking close") + } + + _, ok = runningReaders.Load(key) + require.False(t, ok, "reader ownership should be removed after cleanup finishes") +} + // Integration: commit failure triggers EnsureCleanup rollback, then recovery succeeds func TestTableChangeStream_CommitFailure_EnsureCleanup_ThenRecover(t *testing.T) { updaterStub := newWatermarkUpdaterStub() @@ -2453,6 +2537,28 @@ func newTableStreamRecordingSinker() *tableStreamRecordingSinker { return &tableStreamRecordingSinker{recordingSinker: newRecordingSinker()} } +type blockingCloseSinker struct { + *tableStreamRecordingSinker + closeStarted chan struct{} + unblockClose chan struct{} + closeOnce sync.Once +} + +func newBlockingCloseSinker() *blockingCloseSinker { + return &blockingCloseSinker{ + tableStreamRecordingSinker: newTableStreamRecordingSinker(), + closeStarted: make(chan struct{}), + unblockClose: make(chan struct{}), + } +} + +func (s *blockingCloseSinker) Close() { + s.closeOnce.Do(func() { + close(s.closeStarted) + }) + <-s.unblockClose +} + func (s *tableStreamRecordingSinker) Sink(ctx context.Context, data *DecoderOutput) { s.record("sink") s.mu.Lock() diff --git a/pkg/cdc/table_scanner.go b/pkg/cdc/table_scanner.go index 4794a08693fd3..253e90de6ecd2 100644 --- a/pkg/cdc/table_scanner.go +++ b/pkg/cdc/table_scanner.go @@ -19,7 +19,6 @@ import ( "fmt" "runtime/debug" "slices" - "strings" "sync" "sync/atomic" "time" @@ -37,6 +36,7 @@ import ( "github.com/matrixorigin/matrixone/pkg/defines" "github.com/matrixorigin/matrixone/pkg/logutil" "github.com/matrixorigin/matrixone/pkg/util/executor" + "github.com/matrixorigin/matrixone/pkg/vm/engine" ) const ( @@ -744,6 +744,7 @@ func (s *TableDetector) scanTable() error { } defer result.Close() + var scanErr error result.ReadRows(func(rows int, cols []*vector.Vector) bool { for i := 0; i < rows; i++ { tblId := vector.MustFixedColWithTypeCheck[uint64](cols[0])[i] @@ -752,9 +753,21 @@ func (s *TableDetector) scanTable() error { dbName := cols[3].GetStringAt(i) createSql := cols[4].GetStringAt(i) accountId := vector.MustFixedColWithTypeCheck[uint32](cols[5])[i] + hasForeignKey, decodeErr := tableHasForeignKeyConstraint(cols[6].GetBytesAt(i)) + if decodeErr != nil { + scanErr = decodeErr + logutil.Warn( + "cdc.table_detector.scan_constraint_failed", + zap.Uint32("account-id", accountId), + zap.String("db", dbName), + zap.String("table", tblName), + zap.Error(decodeErr), + ) + return false + } // skip table with foreign key - if strings.Contains(strings.ToLower("createSql"), "foreign key") { + if hasForeignKey { continue } @@ -776,17 +789,21 @@ func (s *TableDetector) scanTable() error { mp[accountId][key] = newInfo } else { idChanged := oldInfo.OnlyDiffinTblId(newInfo) - oldInfo.SourceDbId = dbId - oldInfo.SourceDbName = dbName - oldInfo.SourceTblId = tblId - oldInfo.SourceTblName = tblName - oldInfo.SourceCreateSql = createSql - oldInfo.IdChanged = oldInfo.IdChanged || idChanged - mp[accountId][key] = oldInfo + updatedInfo := oldInfo.Clone() + updatedInfo.SourceDbId = dbId + updatedInfo.SourceDbName = dbName + updatedInfo.SourceTblId = tblId + updatedInfo.SourceTblName = tblName + updatedInfo.SourceCreateSql = createSql + updatedInfo.IdChanged = updatedInfo.IdChanged || idChanged + mp[accountId][key] = updatedInfo } } return true }) + if scanErr != nil { + return scanErr + } // replace the old table map s.mu.Lock() @@ -794,3 +811,26 @@ func (s *TableDetector) scanTable() error { s.mu.Unlock() return nil } + +func tableHasForeignKeyConstraint(data []byte) (hasForeignKey bool, err error) { + if len(data) == 0 { + return false, nil + } + + defer func() { + if r := recover(); r != nil { + err = moerr.NewInternalErrorNoCtxf("unmarshal table constraint failed: %v", r) + } + }() + + constraintDef := &engine.ConstraintDef{} + if err := constraintDef.UnmarshalBinary(data); err != nil { + return false, err + } + for _, constraint := range constraintDef.Cts { + if foreignKeyDef, ok := constraint.(*engine.ForeignKeyDef); ok && len(foreignKeyDef.Fkeys) > 0 { + return true, nil + } + } + return false, nil +} diff --git a/pkg/cdc/table_scanner_test.go b/pkg/cdc/table_scanner_test.go index 8526bad334598..7014204d85761 100644 --- a/pkg/cdc/table_scanner_test.go +++ b/pkg/cdc/table_scanner_test.go @@ -30,9 +30,13 @@ import ( "github.com/matrixorigin/matrixone/pkg/common/moerr" "github.com/matrixorigin/matrixone/pkg/container/batch" "github.com/matrixorigin/matrixone/pkg/container/types" + "github.com/matrixorigin/matrixone/pkg/pb/plan" + "github.com/matrixorigin/matrixone/pkg/sql/parsers" + "github.com/matrixorigin/matrixone/pkg/sql/parsers/dialect" "github.com/matrixorigin/matrixone/pkg/testutil" "github.com/matrixorigin/matrixone/pkg/util/executor" mock_executor "github.com/matrixorigin/matrixone/pkg/util/executor/test" + "github.com/matrixorigin/matrixone/pkg/vm/engine" "github.com/prashantv/gostub" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -64,6 +68,55 @@ func TestApplyTableDetectorOptions(t *testing.T) { assert.Equal(t, DefaultCleanupWarnThreshold, defaultOpts.CleanupWarnThreshold) } +func makeConstraintSQLValue(t *testing.T, constraints ...engine.Constraint) string { + t.Helper() + + if len(constraints) == 0 { + return "" + } + data, err := (&engine.ConstraintDef{Cts: constraints}).MarshalBinary() + require.NoError(t, err) + return string(data) +} + +func makeForeignKeyConstraintSQLValue(t *testing.T) string { + t.Helper() + + return makeConstraintSQLValue(t, &engine.ForeignKeyDef{ + Fkeys: []*plan.ForeignKeyDef{ + { + Name: "fk_child_parent", + Cols: []uint64{2}, + ForeignTbl: 1000, + ForeignCols: []uint64{ + 1, + }, + }, + }, + }) +} + +func TestTableHasForeignKeyConstraint(t *testing.T) { + hasForeignKey, err := tableHasForeignKeyConstraint(nil) + require.NoError(t, err) + assert.False(t, hasForeignKey) + + primaryKeyOnly := makeConstraintSQLValue(t, &engine.PrimaryKeyDef{ + Pkey: &plan.PrimaryKeyDef{PkeyColName: "id"}, + }) + hasForeignKey, err = tableHasForeignKeyConstraint([]byte(primaryKeyOnly)) + require.NoError(t, err) + assert.False(t, hasForeignKey) + + foreignKey := makeForeignKeyConstraintSQLValue(t) + hasForeignKey, err = tableHasForeignKeyConstraint([]byte(foreignKey)) + require.NoError(t, err) + assert.True(t, hasForeignKey) + + _, err = tableHasForeignKeyConstraint([]byte{byte(engine.ForeignKey)}) + require.Error(t, err) +} + func TestTableScanner1(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() @@ -71,13 +124,14 @@ func TestTableScanner1(t *testing.T) { proc := testutil.NewProcess(t) defer proc.Free() - bat := batch.New([]string{"tblId", "tblName", "dbId", "dbName", "createSql", "accountId"}) + bat := batch.New([]string{"tblId", "tblName", "dbId", "dbName", "createSql", "accountId", "constraint"}) bat.Vecs[0] = testutil.MakeUint64Vector([]uint64{1}, nil, proc.Mp()) bat.Vecs[1] = testutil.MakeVarcharVector([]string{"tblName"}, nil, proc.Mp()) bat.Vecs[2] = testutil.MakeUint64Vector([]uint64{1}, nil, proc.Mp()) bat.Vecs[3] = testutil.MakeVarcharVector([]string{"dbName"}, nil, proc.Mp()) bat.Vecs[4] = testutil.MakeVarcharVector([]string{"createSql"}, nil, proc.Mp()) bat.Vecs[5] = testutil.MakeUint32Vector([]uint32{1}, nil, proc.Mp()) + bat.Vecs[6] = testutil.MakeVarcharVector([]string{""}, nil, proc.Mp()) bat.SetRowCount(1) res := executor.Result{ Mp: proc.Mp(), @@ -156,6 +210,251 @@ func TestTableScanner1(t *testing.T) { assert.Equal(t, 0, len(td.Mp)) } +func TestAuditTableScannerSkipsForeignKeyTable(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + proc := testutil.NewProcess(t) + defer proc.Free() + + createSQL := `CREATE TABLE child ( + id BIGINT PRIMARY KEY, + parent_id BIGINT, + FOREIGN KEY (parent_id) REFERENCES parent(id) +)` + + bat := batch.New([]string{"tblId", "tblName", "dbId", "dbName", "createSql", "accountId", "constraint"}) + bat.Vecs[0] = testutil.MakeUint64Vector([]uint64{1001}, nil, proc.Mp()) + bat.Vecs[1] = testutil.MakeVarcharVector([]string{"child"}, nil, proc.Mp()) + bat.Vecs[2] = testutil.MakeUint64Vector([]uint64{10}, nil, proc.Mp()) + bat.Vecs[3] = testutil.MakeVarcharVector([]string{"source_db"}, nil, proc.Mp()) + bat.Vecs[4] = testutil.MakeVarcharVector([]string{createSQL}, nil, proc.Mp()) + bat.Vecs[5] = testutil.MakeUint32Vector([]uint32{1}, nil, proc.Mp()) + bat.Vecs[6] = testutil.MakeVarcharVector([]string{makeForeignKeyConstraintSQLValue(t)}, nil, proc.Mp()) + bat.SetRowCount(1) + res := executor.Result{ + Mp: proc.Mp(), + Batches: []*batch.Batch{bat}, + } + + mockSqlExecutor := mock_executor.NewMockSQLExecutor(ctrl) + mockSqlExecutor.EXPECT().Exec( + gomock.Any(), + CDCSQLBuilder.CollectTableInfoSQL("1", "'source_db'", "'child'"), + gomock.Any(), + ).Return(res, nil) + + td := &TableDetector{ + Mp: make(map[uint32]TblMap), + Callbacks: make(map[string]TableCallback), + CallBackAccountId: make(map[string]uint32), + SubscribedAccountIds: make(map[uint32][]string), + CallBackDbName: make(map[string][]string), + SubscribedDbNames: make(map[string][]string), + CallBackTableName: make(map[string][]string), + SubscribedTableNames: make(map[string][]string), + exec: mockSqlExecutor, + cleanupPeriod: time.Hour, + cleanupWarn: DefaultCleanupWarnThreshold, + } + defer td.Close() + + td.mu.Lock() + td.registerLocked("audit-task", 1, []string{"source_db"}, []string{"child"}, func(mp map[uint32]TblMap) error { + return nil + }) + td.mu.Unlock() + + err := td.scanTable() + require.NoError(t, err) + assert.Empty(t, td.Mp) +} + +func TestTableScannerDoesNotSkipForeignKeyTextLiteral(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + proc := testutil.NewProcess(t) + defer proc.Free() + + createSQL := "CREATE TABLE child (note VARCHAR(32) DEFAULT 'foreign key')" + + bat := batch.New([]string{"tblId", "tblName", "dbId", "dbName", "createSql", "accountId", "constraint"}) + bat.Vecs[0] = testutil.MakeUint64Vector([]uint64{1001}, nil, proc.Mp()) + bat.Vecs[1] = testutil.MakeVarcharVector([]string{"child"}, nil, proc.Mp()) + bat.Vecs[2] = testutil.MakeUint64Vector([]uint64{10}, nil, proc.Mp()) + bat.Vecs[3] = testutil.MakeVarcharVector([]string{"source_db"}, nil, proc.Mp()) + bat.Vecs[4] = testutil.MakeVarcharVector([]string{createSQL}, nil, proc.Mp()) + bat.Vecs[5] = testutil.MakeUint32Vector([]uint32{1}, nil, proc.Mp()) + bat.Vecs[6] = testutil.MakeVarcharVector([]string{""}, nil, proc.Mp()) + bat.SetRowCount(1) + res := executor.Result{ + Mp: proc.Mp(), + Batches: []*batch.Batch{bat}, + } + + mockSqlExecutor := mock_executor.NewMockSQLExecutor(ctrl) + mockSqlExecutor.EXPECT().Exec( + gomock.Any(), + CDCSQLBuilder.CollectTableInfoSQL("1", "'source_db'", "'child'"), + gomock.Any(), + ).Return(res, nil) + + td := &TableDetector{ + Mp: make(map[uint32]TblMap), + Callbacks: make(map[string]TableCallback), + CallBackAccountId: make(map[string]uint32), + SubscribedAccountIds: make(map[uint32][]string), + CallBackDbName: make(map[string][]string), + SubscribedDbNames: make(map[string][]string), + CallBackTableName: make(map[string][]string), + SubscribedTableNames: make(map[string][]string), + exec: mockSqlExecutor, + cleanupPeriod: time.Hour, + cleanupWarn: DefaultCleanupWarnThreshold, + } + defer td.Close() + + td.mu.Lock() + td.registerLocked("audit-task", 1, []string{"source_db"}, []string{"child"}, func(mp map[uint32]TblMap) error { + return nil + }) + td.mu.Unlock() + + err := td.scanTable() + require.NoError(t, err) + require.Contains(t, td.Mp, uint32(1)) + assert.Contains(t, td.Mp[1], "source_db.child") +} + +func TestTableScannerSkipsForeignKeyMetadataWithoutCreateSQLText(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + proc := testutil.NewProcess(t) + defer proc.Free() + + createSQL := "CREATE TABLE child (id BIGINT PRIMARY KEY, parent_id BIGINT)" + + bat := batch.New([]string{"tblId", "tblName", "dbId", "dbName", "createSql", "accountId", "constraint"}) + bat.Vecs[0] = testutil.MakeUint64Vector([]uint64{1001}, nil, proc.Mp()) + bat.Vecs[1] = testutil.MakeVarcharVector([]string{"child"}, nil, proc.Mp()) + bat.Vecs[2] = testutil.MakeUint64Vector([]uint64{10}, nil, proc.Mp()) + bat.Vecs[3] = testutil.MakeVarcharVector([]string{"source_db"}, nil, proc.Mp()) + bat.Vecs[4] = testutil.MakeVarcharVector([]string{createSQL}, nil, proc.Mp()) + bat.Vecs[5] = testutil.MakeUint32Vector([]uint32{1}, nil, proc.Mp()) + bat.Vecs[6] = testutil.MakeVarcharVector([]string{makeForeignKeyConstraintSQLValue(t)}, nil, proc.Mp()) + bat.SetRowCount(1) + res := executor.Result{ + Mp: proc.Mp(), + Batches: []*batch.Batch{bat}, + } + + mockSqlExecutor := mock_executor.NewMockSQLExecutor(ctrl) + mockSqlExecutor.EXPECT().Exec( + gomock.Any(), + CDCSQLBuilder.CollectTableInfoSQL("1", "'source_db'", "'child'"), + gomock.Any(), + ).Return(res, nil) + + td := &TableDetector{ + Mp: make(map[uint32]TblMap), + Callbacks: make(map[string]TableCallback), + CallBackAccountId: make(map[string]uint32), + SubscribedAccountIds: make(map[uint32][]string), + CallBackDbName: make(map[string][]string), + SubscribedDbNames: make(map[string][]string), + CallBackTableName: make(map[string][]string), + SubscribedTableNames: make(map[string][]string), + exec: mockSqlExecutor, + cleanupPeriod: time.Hour, + cleanupWarn: DefaultCleanupWarnThreshold, + } + defer td.Close() + + td.mu.Lock() + td.registerLocked("audit-task", 1, []string{"source_db"}, []string{"child"}, func(mp map[uint32]TblMap) error { + return nil + }) + td.mu.Unlock() + + err := td.scanTable() + require.NoError(t, err) + assert.Empty(t, td.Mp) +} + +func TestTableScannerConstraintDecodeErrorPreservesOldTableMap(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + proc := testutil.NewProcess(t) + defer proc.Free() + + bat := batch.New([]string{"tblId", "tblName", "dbId", "dbName", "createSql", "accountId", "constraint"}) + bat.Vecs[0] = testutil.MakeUint64Vector([]uint64{1001, 1002}, nil, proc.Mp()) + bat.Vecs[1] = testutil.MakeVarcharVector([]string{"child", "broken"}, nil, proc.Mp()) + bat.Vecs[2] = testutil.MakeUint64Vector([]uint64{10, 10}, nil, proc.Mp()) + bat.Vecs[3] = testutil.MakeVarcharVector([]string{"source_db", "source_db"}, nil, proc.Mp()) + bat.Vecs[4] = testutil.MakeVarcharVector([]string{ + "CREATE TABLE child (id BIGINT PRIMARY KEY)", + "CREATE TABLE broken (id BIGINT PRIMARY KEY)", + }, nil, proc.Mp()) + bat.Vecs[5] = testutil.MakeUint32Vector([]uint32{1, 1}, nil, proc.Mp()) + bat.Vecs[6] = testutil.MakeVarcharVector([]string{"", string([]byte{byte(engine.ForeignKey)})}, nil, proc.Mp()) + bat.SetRowCount(2) + res := executor.Result{ + Mp: proc.Mp(), + Batches: []*batch.Batch{bat}, + } + + mockSqlExecutor := mock_executor.NewMockSQLExecutor(ctrl) + mockSqlExecutor.EXPECT().Exec( + gomock.Any(), + CDCSQLBuilder.CollectTableInfoSQL("1", "'source_db'", "*"), + gomock.Any(), + ).Return(res, nil) + + oldInfo := &DbTableInfo{ + SourceDbId: 9, + SourceDbName: "source_db", + SourceTblId: 9001, + SourceTblName: "child", + SourceCreateSql: "CREATE TABLE child (old_id BIGINT PRIMARY KEY)", + IdChanged: false, + } + td := &TableDetector{ + Mp: map[uint32]TblMap{1: {GenDbTblKey("source_db", "child"): oldInfo}}, + Callbacks: make(map[string]TableCallback), + CallBackAccountId: make(map[string]uint32), + SubscribedAccountIds: make(map[uint32][]string), + CallBackDbName: make(map[string][]string), + SubscribedDbNames: make(map[string][]string), + CallBackTableName: make(map[string][]string), + SubscribedTableNames: make(map[string][]string), + exec: mockSqlExecutor, + cleanupPeriod: time.Hour, + cleanupWarn: DefaultCleanupWarnThreshold, + } + defer td.Close() + + td.mu.Lock() + td.registerLocked("audit-task", 1, []string{"source_db"}, []string{"*"}, func(mp map[uint32]TblMap) error { + return nil + }) + td.mu.Unlock() + + err := td.scanTable() + require.Error(t, err) + require.Contains(t, td.Mp, uint32(1)) + gotInfo := td.Mp[1][GenDbTblKey("source_db", "child")] + require.Same(t, oldInfo, gotInfo) + assert.Equal(t, uint64(9), gotInfo.SourceDbId) + assert.Equal(t, uint64(9001), gotInfo.SourceTblId) + assert.Equal(t, "CREATE TABLE child (old_id BIGINT PRIMARY KEY)", gotInfo.SourceCreateSql) + assert.False(t, gotInfo.IdChanged) + assert.NotContains(t, td.Mp[1], GenDbTblKey("source_db", "broken")) +} + func TestTableDetectorRegisterStartsScanAsync(t *testing.T) { td := &TableDetector{ Mp: make(map[uint32]TblMap), @@ -683,17 +982,21 @@ func TestTableDetectorConcurrentRegister(t *testing.T) { func Test_CollectTableInfoSQL(t *testing.T) { var builder cdcSQLBuilder sql := builder.CollectTableInfoSQL("1,2,3", "*", "*") + _, err := parsers.ParseOne(context.Background(), dialect.MYSQL, sql, 1) + require.NoError(t, err) sql = strings.ToUpper(sql) t.Log(sql) - expected := "SELECT REL_ID, RELNAME, RELDATABASE_ID, " + - "RELDATABASE, REL_CREATESQL, ACCOUNT_ID " + - "FROM `MO_CATALOG`.`MO_TABLES` " + - "WHERE ACCOUNT_ID IN (1,2,3) AND RELKIND = 'R' " + - "AND RELDATABASE NOT IN ('INFORMATION_SCHEMA','MO_CATALOG','MO_DEBUG','MO_TASK','MYSQL','SYSTEM','SYSTEM_METRICS')" + expected := "SELECT TBL.REL_ID, TBL.RELNAME, TBL.RELDATABASE_ID, " + + "TBL.RELDATABASE, TBL.REL_CREATESQL, TBL.ACCOUNT_ID, TBL.`CONSTRAINT` " + + "FROM `MO_CATALOG`.`MO_TABLES` TBL " + + "WHERE TBL.ACCOUNT_ID IN (1,2,3) AND TBL.RELKIND = 'R' " + + "AND TBL.RELDATABASE NOT IN ('INFORMATION_SCHEMA','MO_CATALOG','MO_DEBUG','MO_TASK','MYSQL','SYSTEM','SYSTEM_METRICS')" assert.Equal(t, expected, sql) sql = builder.CollectTableInfoSQL("0", "'source_db'", "'orders'") - expected = "SELECT REL_ID, RELNAME, RELDATABASE_ID, RELDATABASE, REL_CREATESQL, ACCOUNT_ID FROM `MO_CATALOG`.`MO_TABLES` WHERE ACCOUNT_ID IN (0) AND RELDATABASE IN ('SOURCE_DB') AND RELNAME IN ('ORDERS') AND RELKIND = 'R' AND RELDATABASE NOT IN ('INFORMATION_SCHEMA','MO_CATALOG','MO_DEBUG','MO_TASK','MYSQL','SYSTEM','SYSTEM_METRICS')" + _, err = parsers.ParseOne(context.Background(), dialect.MYSQL, sql, 1) + require.NoError(t, err) + expected = "SELECT TBL.REL_ID, TBL.RELNAME, TBL.RELDATABASE_ID, TBL.RELDATABASE, TBL.REL_CREATESQL, TBL.ACCOUNT_ID, TBL.`CONSTRAINT` FROM `MO_CATALOG`.`MO_TABLES` TBL WHERE TBL.ACCOUNT_ID IN (0) AND TBL.RELDATABASE IN ('SOURCE_DB') AND TBL.RELNAME IN ('ORDERS') AND TBL.RELKIND = 'R' AND TBL.RELDATABASE NOT IN ('INFORMATION_SCHEMA','MO_CATALOG','MO_DEBUG','MO_TASK','MYSQL','SYSTEM','SYSTEM_METRICS')" assert.Equal(t, strings.ToUpper(expected), strings.ToUpper(sql)) } @@ -856,26 +1159,28 @@ func TestTableScanner_UpdateTableInfo(t *testing.T) { proc := testutil.NewProcess(t) defer proc.Free() - bat1 := batch.New([]string{"tblId", "tblName", "dbId", "dbName", "createSql", "accountId"}) + bat1 := batch.New([]string{"tblId", "tblName", "dbId", "dbName", "createSql", "accountId", "constraint"}) bat1.Vecs[0] = testutil.MakeUint64Vector([]uint64{1001}, nil, proc.Mp()) bat1.Vecs[1] = testutil.MakeVarcharVector([]string{"tbl1"}, nil, proc.Mp()) bat1.Vecs[2] = testutil.MakeUint64Vector([]uint64{1}, nil, proc.Mp()) bat1.Vecs[3] = testutil.MakeVarcharVector([]string{"db1"}, nil, proc.Mp()) bat1.Vecs[4] = testutil.MakeVarcharVector([]string{"create table tbl1 (a int)"}, nil, proc.Mp()) bat1.Vecs[5] = testutil.MakeUint32Vector([]uint32{1}, nil, proc.Mp()) + bat1.Vecs[6] = testutil.MakeVarcharVector([]string{""}, nil, proc.Mp()) bat1.SetRowCount(1) res1 := executor.Result{ Mp: proc.Mp(), Batches: []*batch.Batch{bat1}, } - bat2 := batch.New([]string{"tblId", "tblName", "dbId", "dbName", "createSql", "accountId"}) + bat2 := batch.New([]string{"tblId", "tblName", "dbId", "dbName", "createSql", "accountId", "constraint"}) bat2.Vecs[0] = testutil.MakeUint64Vector([]uint64{1002}, nil, proc.Mp()) bat2.Vecs[1] = testutil.MakeVarcharVector([]string{"tbl1"}, nil, proc.Mp()) bat2.Vecs[2] = testutil.MakeUint64Vector([]uint64{1}, nil, proc.Mp()) bat2.Vecs[3] = testutil.MakeVarcharVector([]string{"db1"}, nil, proc.Mp()) bat2.Vecs[4] = testutil.MakeVarcharVector([]string{"create table tbl1 (a int)"}, nil, proc.Mp()) bat2.Vecs[5] = testutil.MakeUint32Vector([]uint32{1}, nil, proc.Mp()) + bat2.Vecs[6] = testutil.MakeVarcharVector([]string{""}, nil, proc.Mp()) bat2.SetRowCount(1) res2 := executor.Result{ Mp: proc.Mp(), diff --git a/pkg/frontend/cdc_exector.go b/pkg/frontend/cdc_exector.go index b5a5c22973804..18209ac1174fa 100644 --- a/pkg/frontend/cdc_exector.go +++ b/pkg/frontend/cdc_exector.go @@ -131,6 +131,8 @@ type CDCTaskExecutor struct { watermarkUpdater *cdc.CDCWatermarkUpdater // runningReaders store the running execute pipelines, map key pattern: db.table runningReaders *sync.Map + // removedReaderShutdowns stores in-progress shutdowns for readers that disappeared from scan results. + removedReaderShutdowns sync.Map // stateMachine manages executor state transitions stateMachine *ExecutorStateMachine @@ -836,6 +838,98 @@ func (exec *CDCTaskExecutor) stopAllReaders() { ) } +type removedReaderShutdown struct { + reader cdc.ChangeReader + done chan struct{} +} + +func (exec *CDCTaskExecutor) stopReadersMissingFromScan(accountTbls cdc.TblMap) { + if exec.runningReaders == nil { + return + } + + exec.runningReaders.Range(func(key, value interface{}) bool { + tableKey, ok := key.(string) + if !ok { + return true + } + if _, ok = accountTbls[tableKey]; ok { + return true + } + + reader, ok := value.(cdc.ChangeReader) + if !ok { + exec.runningReaders.Delete(key) + return true + } + + if !exec.matchesAnySourcePattern(tableKey) { + return true + } + + exec.stopRemovedReader(tableKey, key, reader) + return true + }) +} + +func (exec *CDCTaskExecutor) stopRemovedReader(tableKey string, mapKey interface{}, reader cdc.ChangeReader) { + shutdown := &removedReaderShutdown{ + reader: reader, + done: make(chan struct{}), + } + + actual, loaded := exec.removedReaderShutdowns.LoadOrStore(tableKey, shutdown) + if loaded { + existing, ok := actual.(*removedReaderShutdown) + if ok && existing.reader == reader { + select { + case <-existing.done: + exec.removedReaderShutdowns.CompareAndDelete(tableKey, existing) + default: + return + } + } else { + exec.removedReaderShutdowns.CompareAndDelete(tableKey, actual) + } + _, loaded = exec.removedReaderShutdowns.LoadOrStore(tableKey, shutdown) + if loaded { + return + } + } + + logutil.Info( + "cdc.frontend.task.stop_reader_removed_from_scan", + zap.String("task-id", exec.spec.TaskId), + zap.String("task-name", exec.spec.TaskName), + zap.String("table", tableKey), + ) + + go func() { + reader.Close() + reader.Wait() + exec.runningReaders.CompareAndDelete(mapKey, reader) + close(shutdown.done) + exec.removedReaderShutdowns.CompareAndDelete(tableKey, shutdown) + }() +} + +func (exec *CDCTaskExecutor) removedReaderShutdownInProgress(tableKey string, reader cdc.ChangeReader) bool { + actual, ok := exec.removedReaderShutdowns.Load(tableKey) + if !ok { + return false + } + shutdown, ok := actual.(*removedReaderShutdown) + if !ok || shutdown.reader != reader { + return false + } + select { + case <-shutdown.done: + return false + default: + return true + } +} + func (exec *CDCTaskExecutor) initAesKeyByInternalExecutor(ctx context.Context, accountId uint32) (err error) { if len(cdc.AesKey) > 0 { return nil @@ -1130,14 +1224,25 @@ func (exec *CDCTaskExecutor) handleNewTablesForGeneration( // Track failed tables for better error reporting failedTables := make(map[string]error) successCount := 0 + accountTbls := allAccountTbls[accountId] + exec.stopReadersMissingFromScan(accountTbls) - for key, info := range allAccountTbls[accountId] { + for key, info := range accountTbls { // already running if val, ok := exec.runningReaders.Load(key); ok { if reader, ok := val.(cdc.ChangeReader); ok { readerInfo := reader.GetTableInfo() // wait the old reader to stop if info.OnlyDiffinTblId(readerInfo) { + if exec.removedReaderShutdownInProgress(key, reader) { + logutil.Info( + "cdc.frontend.task.skip_wait_removed_reader_shutdown", + zap.String("table", key), + zap.Uint64("old-table-id", readerInfo.SourceTblId), + zap.Uint64("new-table-id", info.SourceTblId), + ) + continue + } logutil.Info( "cdc.frontend.task.wait_old_reader", zap.String("table", key), @@ -1433,6 +1538,23 @@ func (exec *CDCTaskExecutor) matchAnyPattern(key string, info *cdc.DbTableInfo) return false } +func (exec *CDCTaskExecutor) matchesAnySourcePattern(key string) bool { + match := func(s, p string) bool { + if p == cdc.CDCPitrGranularity_All { + return true + } + return s == p + } + + db, table := cdc.SplitDbTblKey(key) + for _, pt := range exec.tables.Pts { + if match(db, pt.Source.Database) && match(table, pt.Source.Table) { + return true + } + } + return false +} + // reader ----> sinker ----> remote db func (exec *CDCTaskExecutor) addExecPipelineForTable( ctx context.Context, diff --git a/pkg/frontend/cdc_test.go b/pkg/frontend/cdc_test.go index 9e469b7715997..c0b892dcb69d5 100644 --- a/pkg/frontend/cdc_test.go +++ b/pkg/frontend/cdc_test.go @@ -22,6 +22,7 @@ import ( "reflect" "regexp" "sync" + "sync/atomic" "testing" "time" @@ -4270,6 +4271,413 @@ func TestCdcTask_handleNewTables_existingReaderWithDifferentTableID(t *testing.T cdcTask.handleNewTables(mp) } +func TestCdcTask_handleNewTablesStopsReaderRemovedFromScan(t *testing.T) { + stub1 := gostub.Stub(&cdc.GetTxnOp, func(context.Context, engine.Engine, client.TxnClient, string) (client.TxnOperator, error) { + return nil, nil + }) + defer stub1.Reset() + + stub2 := gostub.Stub(&cdc.FinishTxnOp, func(context.Context, error, client.TxnOperator, engine.Engine) {}) + defer stub2.Reset() + + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + eng := mock_frontend.NewMockEngine(ctrl) + eng.EXPECT().New(gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + + closeCh := make(chan struct{}) + readerInfo := &cdc.DbTableInfo{ + SourceDbName: "db1", + SourceTblName: "child", + SourceTblId: 1001, + SinkDbName: "live_sink_db", + SinkTblName: "live_sink_table", + } + oldReader := &mockChangeReader{ + info: readerInfo, + closeCh: closeCh, + } + + cdcTask := &CDCTaskExecutor{ + spec: &task.CreateCdcDetails{ + TaskId: "task-removed-from-scan", + TaskName: "task-removed-from-scan", + Accounts: []*task.Account{ + {Id: 0}, + }, + }, + tables: cdc.PatternTuples{ + Pts: []*cdc.PatternTuple{ + { + Source: cdc.PatternTable{ + Database: "db1", + Table: cdc.CDCPitrGranularity_All, + }, + Sink: cdc.PatternTable{ + Database: cdc.CDCPitrGranularity_All, + Table: cdc.CDCPitrGranularity_All, + }, + }, + }, + }, + cnEngine: eng, + runningReaders: &sync.Map{}, + } + cdcTask.runningReaders.Store("db1.child", oldReader) + + err := cdcTask.handleNewTables(map[uint32]cdc.TblMap{0: {}}) + require.NoError(t, err) + + select { + case <-closeCh: + case <-time.After(time.Second): + t.Fatal("expected removed reader to be closed") + } + require.Eventually(t, func() bool { + _, ok := cdcTask.runningReaders.Load("db1.child") + return !ok + }, time.Second, time.Millisecond) + require.Equal(t, "live_sink_db", readerInfo.SinkDbName) + require.Equal(t, "live_sink_table", readerInfo.SinkTblName) +} + +func TestCdcTask_handleNewTablesKeepsBlockedRemovedReaderOwnership(t *testing.T) { + stub1 := gostub.Stub(&cdc.GetTxnOp, func(context.Context, engine.Engine, client.TxnClient, string) (client.TxnOperator, error) { + return nil, nil + }) + defer stub1.Reset() + + stub2 := gostub.Stub(&cdc.FinishTxnOp, func(context.Context, error, client.TxnOperator, engine.Engine) {}) + defer stub2.Reset() + + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + eng := mock_frontend.NewMockEngine(ctrl) + eng.EXPECT().New(gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + + closeCh := make(chan struct{}) + waitCh := make(chan struct{}) + var closeCalls atomic.Int32 + var waitCalls atomic.Int32 + oldReader := &mockChangeReader{ + info: &cdc.DbTableInfo{SourceDbName: "db1", SourceTblName: "child", SourceTblId: 1001}, + closeCh: closeCh, + waitCh: waitCh, + closeCalls: &closeCalls, + waitCalls: &waitCalls, + } + t.Cleanup(func() { + close(waitCh) + }) + + cdcTask := &CDCTaskExecutor{ + spec: &task.CreateCdcDetails{ + TaskId: "task-removed-from-scan-timeout", + TaskName: "task-removed-from-scan-timeout", + Accounts: []*task.Account{ + {Id: 0}, + }, + }, + tables: cdc.PatternTuples{ + Pts: []*cdc.PatternTuple{ + { + Source: cdc.PatternTable{ + Database: "db1", + Table: cdc.CDCPitrGranularity_All, + }, + Sink: cdc.PatternTable{ + Database: cdc.CDCPitrGranularity_All, + Table: cdc.CDCPitrGranularity_All, + }, + }, + }, + }, + cnEngine: eng, + runningReaders: &sync.Map{}, + } + cdcTask.runningReaders.Store("db1.child", oldReader) + + err := cdcTask.handleNewTables(map[uint32]cdc.TblMap{0: {}}) + require.NoError(t, err) + + select { + case <-closeCh: + case <-time.After(time.Second): + t.Fatal("expected removed reader to be closed") + } + val, ok := cdcTask.runningReaders.Load("db1.child") + require.True(t, ok) + require.Same(t, oldReader, val) + require.Equal(t, int32(1), closeCalls.Load()) + require.Eventually(t, func() bool { + return waitCalls.Load() == 1 + }, time.Second, time.Millisecond) + + start := time.Now() + err = cdcTask.handleNewTables(map[uint32]cdc.TblMap{0: {}}) + require.NoError(t, err) + require.Less(t, time.Since(start), 200*time.Millisecond) + require.Equal(t, int32(1), closeCalls.Load()) + require.Equal(t, int32(1), waitCalls.Load()) + val, ok = cdcTask.runningReaders.Load("db1.child") + require.True(t, ok) + require.Same(t, oldReader, val) + + start = time.Now() + err = cdcTask.handleNewTables(map[uint32]cdc.TblMap{ + 0: { + "db1.child": &cdc.DbTableInfo{SourceDbName: "db1", SourceTblName: "child", SourceTblId: 2002}, + }, + }) + require.NoError(t, err) + require.Less(t, time.Since(start), 200*time.Millisecond) + require.Equal(t, int32(1), closeCalls.Load()) + require.Equal(t, int32(1), waitCalls.Load()) + val, ok = cdcTask.runningReaders.Load("db1.child") + require.True(t, ok) + require.Same(t, oldReader, val) + + err = cdcTask.handleNewTables(map[uint32]cdc.TblMap{ + 0: { + "db1.child": &cdc.DbTableInfo{SourceDbName: "db1", SourceTblName: "child", SourceTblId: 1001}, + }, + }) + require.NoError(t, err) + val, ok = cdcTask.runningReaders.Load("db1.child") + require.True(t, ok) + require.Same(t, oldReader, val) +} + +func TestCdcTask_handleNewTablesDoesNotWaitForBlockedRemovedReaderClose(t *testing.T) { + stub1 := gostub.Stub(&cdc.GetTxnOp, func(context.Context, engine.Engine, client.TxnClient, string) (client.TxnOperator, error) { + return nil, nil + }) + defer stub1.Reset() + + stub2 := gostub.Stub(&cdc.FinishTxnOp, func(context.Context, error, client.TxnOperator, engine.Engine) {}) + defer stub2.Reset() + + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + eng := mock_frontend.NewMockEngine(ctrl) + eng.EXPECT().New(gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + + closeCh := make(chan struct{}) + closeBlockCh := make(chan struct{}) + var closeCalls atomic.Int32 + var waitCalls atomic.Int32 + oldReader := &mockChangeReader{ + info: &cdc.DbTableInfo{SourceDbName: "db1", SourceTblName: "child", SourceTblId: 1001}, + closeCh: closeCh, + closeBlockCh: closeBlockCh, + closeCalls: &closeCalls, + waitCalls: &waitCalls, + } + t.Cleanup(func() { + close(closeBlockCh) + }) + + cdcTask := &CDCTaskExecutor{ + spec: &task.CreateCdcDetails{ + TaskId: "task-removed-from-scan-close-timeout", + TaskName: "task-removed-from-scan-close-timeout", + Accounts: []*task.Account{ + {Id: 0}, + }, + }, + tables: cdc.PatternTuples{ + Pts: []*cdc.PatternTuple{ + { + Source: cdc.PatternTable{ + Database: "db1", + Table: cdc.CDCPitrGranularity_All, + }, + Sink: cdc.PatternTable{ + Database: cdc.CDCPitrGranularity_All, + Table: cdc.CDCPitrGranularity_All, + }, + }, + }, + }, + cnEngine: eng, + runningReaders: &sync.Map{}, + } + cdcTask.runningReaders.Store("db1.child", oldReader) + + start := time.Now() + err := cdcTask.handleNewTables(map[uint32]cdc.TblMap{0: {}}) + require.NoError(t, err) + require.Less(t, time.Since(start), 200*time.Millisecond) + + select { + case <-closeCh: + case <-time.After(time.Second): + t.Fatal("expected removed reader close to start") + } + require.Equal(t, int32(1), closeCalls.Load()) + require.Equal(t, int32(0), waitCalls.Load()) + val, ok := cdcTask.runningReaders.Load("db1.child") + require.True(t, ok) + require.Same(t, oldReader, val) +} + +func TestCdcTask_handleNewTablesStartsRemovedReaderShutdownsWithoutWaiting(t *testing.T) { + stub1 := gostub.Stub(&cdc.GetTxnOp, func(context.Context, engine.Engine, client.TxnClient, string) (client.TxnOperator, error) { + return nil, nil + }) + defer stub1.Reset() + + stub2 := gostub.Stub(&cdc.FinishTxnOp, func(context.Context, error, client.TxnOperator, engine.Engine) {}) + defer stub2.Reset() + + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + eng := mock_frontend.NewMockEngine(ctrl) + eng.EXPECT().New(gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + + closeBlockCh := make(chan struct{}) + var closeCalls atomic.Int32 + var waitCalls atomic.Int32 + t.Cleanup(func() { + close(closeBlockCh) + }) + + cdcTask := &CDCTaskExecutor{ + spec: &task.CreateCdcDetails{ + TaskId: "task-removed-from-scan-batch-timeout", + TaskName: "task-removed-from-scan-batch-timeout", + Accounts: []*task.Account{ + {Id: 0}, + }, + }, + tables: cdc.PatternTuples{ + Pts: []*cdc.PatternTuple{ + { + Source: cdc.PatternTable{ + Database: "db1", + Table: cdc.CDCPitrGranularity_All, + }, + Sink: cdc.PatternTable{ + Database: cdc.CDCPitrGranularity_All, + Table: cdc.CDCPitrGranularity_All, + }, + }, + }, + }, + cnEngine: eng, + runningReaders: &sync.Map{}, + } + + const readerCount = 3 + for i := 0; i < readerCount; i++ { + tableName := fmt.Sprintf("child_%d", i) + reader := &mockChangeReader{ + info: &cdc.DbTableInfo{ + SourceDbName: "db1", + SourceTblName: tableName, + SourceTblId: uint64(1000 + i), + }, + closeBlockCh: closeBlockCh, + closeCalls: &closeCalls, + waitCalls: &waitCalls, + } + cdcTask.runningReaders.Store("db1."+tableName, reader) + } + + start := time.Now() + err := cdcTask.handleNewTables(map[uint32]cdc.TblMap{0: {}}) + elapsed := time.Since(start) + require.NoError(t, err) + require.Less(t, elapsed, 200*time.Millisecond) + require.Eventually(t, func() bool { + return closeCalls.Load() == readerCount + }, time.Second, time.Millisecond) + require.Equal(t, int32(0), waitCalls.Load()) + + for i := 0; i < readerCount; i++ { + val, ok := cdcTask.runningReaders.Load(fmt.Sprintf("db1.child_%d", i)) + require.True(t, ok) + require.NotNil(t, val) + } +} + +func TestCdcTask_handleNewTablesDoesNotWaitAcrossTaskCallbacks(t *testing.T) { + stub1 := gostub.Stub(&cdc.GetTxnOp, func(context.Context, engine.Engine, client.TxnClient, string) (client.TxnOperator, error) { + return nil, nil + }) + defer stub1.Reset() + + stub2 := gostub.Stub(&cdc.FinishTxnOp, func(context.Context, error, client.TxnOperator, engine.Engine) {}) + defer stub2.Reset() + + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + eng := mock_frontend.NewMockEngine(ctrl) + eng.EXPECT().New(gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + + closeBlockCh := make(chan struct{}) + var closeCalls atomic.Int32 + t.Cleanup(func() { + close(closeBlockCh) + }) + + const taskCount = 3 + tasks := make([]*CDCTaskExecutor, 0, taskCount) + for i := 0; i < taskCount; i++ { + cdcTask := &CDCTaskExecutor{ + spec: &task.CreateCdcDetails{ + TaskId: fmt.Sprintf("task-removed-from-scan-callback-%d", i), + TaskName: fmt.Sprintf("task-removed-from-scan-callback-%d", i), + Accounts: []*task.Account{ + {Id: 0}, + }, + }, + tables: cdc.PatternTuples{ + Pts: []*cdc.PatternTuple{ + { + Source: cdc.PatternTable{ + Database: "db1", + Table: cdc.CDCPitrGranularity_All, + }, + Sink: cdc.PatternTable{ + Database: cdc.CDCPitrGranularity_All, + Table: cdc.CDCPitrGranularity_All, + }, + }, + }, + }, + cnEngine: eng, + runningReaders: &sync.Map{}, + } + reader := &mockChangeReader{ + info: &cdc.DbTableInfo{ + SourceDbName: "db1", + SourceTblName: fmt.Sprintf("child_%d", i), + SourceTblId: uint64(1000 + i), + }, + closeBlockCh: closeBlockCh, + closeCalls: &closeCalls, + } + cdcTask.runningReaders.Store(fmt.Sprintf("db1.child_%d", i), reader) + tasks = append(tasks, cdcTask) + } + + start := time.Now() + for _, cdcTask := range tasks { + err := cdcTask.handleNewTables(map[uint32]cdc.TblMap{0: {}}) + require.NoError(t, err) + } + require.Less(t, time.Since(start), 200*time.Millisecond) + require.Eventually(t, func() bool { + return closeCalls.Load() == taskCount + }, time.Second, time.Millisecond) +} + // setupCDCTestStubs sets up all necessary stubs for CDC tests that create TableChangeStream // This prevents nil pointer panics when TableChangeStream.Run() is called in goroutines func setupCDCTestStubs(t *testing.T) []*gostub.Stubs { @@ -4312,21 +4720,46 @@ func setupCDCTestStubs(t *testing.T) []*gostub.Stubs { } type mockChangeReader struct { - info *cdc.DbTableInfo - wg *sync.WaitGroup + info *cdc.DbTableInfo + wg *sync.WaitGroup + closeCh chan struct{} + closeOnce sync.Once + closeBlockCh chan struct{} + waitCh chan struct{} + closeCalls *atomic.Int32 + waitCalls *atomic.Int32 } -func (m mockChangeReader) Run(ctx context.Context, ar *cdc.ActiveRoutine) {} +func (m *mockChangeReader) Run(ctx context.Context, ar *cdc.ActiveRoutine) {} -func (m mockChangeReader) Close() {} +func (m *mockChangeReader) Close() { + if m.closeCalls != nil { + m.closeCalls.Add(1) + } + if m.closeCh != nil { + m.closeOnce.Do(func() { + close(m.closeCh) + }) + } + if m.closeBlockCh != nil { + <-m.closeBlockCh + } +} -func (m mockChangeReader) Wait() { +func (m *mockChangeReader) Wait() { + if m.waitCalls != nil { + m.waitCalls.Add(1) + } + if m.waitCh != nil { + <-m.waitCh + return + } if m.wg != nil { m.wg.Wait() } } -func (m mockChangeReader) GetTableInfo() *cdc.DbTableInfo { +func (m *mockChangeReader) GetTableInfo() *cdc.DbTableInfo { return m.info }