diff --git a/cmd/kafka-consumer/writer.go b/cmd/kafka-consumer/writer.go index aea783c091..5f9991e915 100644 --- a/cmd/kafka-consumer/writer.go +++ b/cmd/kafka-consumer/writer.go @@ -26,13 +26,13 @@ import ( "github.com/pingcap/ticdc/downstreamadapter/sink" "github.com/pingcap/ticdc/downstreamadapter/sink/eventrouter" commonType "github.com/pingcap/ticdc/pkg/common" - commonEvent "github.com/pingcap/ticdc/pkg/common/event" + "github.com/pingcap/ticdc/pkg/common/event" "github.com/pingcap/ticdc/pkg/config" "github.com/pingcap/ticdc/pkg/errors" "github.com/pingcap/ticdc/pkg/sink/codec" "github.com/pingcap/ticdc/pkg/sink/codec/common" "github.com/pingcap/ticdc/pkg/sink/codec/simple" - timodel "github.com/pingcap/tidb/pkg/meta/model" + "github.com/pingcap/tidb/pkg/meta/model" "github.com/pingcap/tidb/pkg/parser" "github.com/pingcap/tidb/pkg/parser/ast" "go.uber.org/atomic" @@ -64,10 +64,8 @@ func (p *partitionProgress) updateWatermark(newWatermark uint64, offset kafka.Of zap.Uint64("watermark", newWatermark)) return } - readOldOffset := true - if offset > p.watermarkOffset { - readOldOffset = false - } + readOldOffset := offset <= p.watermarkOffset + log.Warn("partition resolved ts fall back, ignore it", zap.Bool("readOldOffset", readOldOffset), zap.Int32("partition", p.partition), @@ -77,7 +75,7 @@ func (p *partitionProgress) updateWatermark(newWatermark uint64, offset kafka.Of type writer struct { progresses []*partitionProgress - ddlList []*commonEvent.DDLEvent + ddlList []*event.DDLEvent ddlWithMaxCommitTs map[int64]uint64 // this should be used by the canal-json, avro and open protocol @@ -98,7 +96,7 @@ func newWriter(ctx context.Context, o *option) *writer { maxBatchSize: o.maxBatchSize, progresses: make([]*partitionProgress, o.partitionNum), partitionTableAccessor: common.NewPartitionTableAccessor(), - ddlList: make([]*commonEvent.DDLEvent, 0), + ddlList: make([]*event.DDLEvent, 0), ddlWithMaxCommitTs: make(map[int64]uint64), enableTableAcrossNodes: o.enableTableAcrossNodes, } @@ -149,47 +147,32 @@ func (w *writer) run(ctx context.Context) error { return w.mysqlSink.Run(ctx) } -func (w *writer) flushDDLEvent(ctx context.Context, ddl *commonEvent.DDLEvent) error { +func (w *writer) flushDDLEvent(ctx context.Context, ddl *event.DDLEvent) error { var ( done = make(chan struct{}, 1) - total int flushed atomic.Int64 ) tableIDs := w.getBlockTableIDs(ddl) commitTs := ddl.GetCommitTs() - resolvedEvents := make([]*commonEvent.DMLEvent, 0) - // resolvedGroups records which EventsGroup has flushed events so we can - // advance its AppliedWatermark after the flush is fully finished. - resolvedGroups := make([]struct { - group *util.EventsGroup - maxCommitTs uint64 - }, 0) + resolvedEvents := make([]*event.DMLEvent, 0) for tableID := range tableIDs { for _, progress := range w.progresses { g, ok := progress.eventsGroup[tableID] if !ok { continue } - before := len(resolvedEvents) - resolvedEvents = g.ResolveInto(commitTs, resolvedEvents) - resolvedCount := len(resolvedEvents) - before - if resolvedCount == 0 { - continue + messages := g.ResolveInto(commitTs, nil) + events := make([]*event.DMLEvent, 0, len(messages)) + for _, message := range messages { + events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) } - - resolvedGroups = append(resolvedGroups, struct { - group *util.EventsGroup - maxCommitTs uint64 - }{ - group: g, - maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), - }) - total += resolvedCount + resolvedEvents = append(resolvedEvents, events...) } } + total := len(resolvedEvents) if total == 0 { return w.mysqlSink.WriteBlockEvent(ddl) } @@ -214,11 +197,6 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *commonEvent.DDLEvent) e log.Info("flush DML events before DDL done", zap.Uint64("DDLCommitTs", commitTs), zap.Int("total", total), zap.Duration("duration", time.Since(start)), zap.Any("tables", tableIDs)) - for _, item := range resolvedGroups { - if item.maxCommitTs > item.group.AppliedWatermark { - item.group.AppliedWatermark = item.maxCommitTs - } - } return w.mysqlSink.WriteBlockEvent(ddl) case <-ticker.C: log.Warn("DML events cannot be flushed in time", @@ -228,20 +206,20 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *commonEvent.DDLEvent) e } } -func (w *writer) getBlockTableIDs(ddl *commonEvent.DDLEvent) map[int64]struct{} { +func (w *writer) getBlockTableIDs(ddl *event.DDLEvent) map[int64]struct{} { // The DDL event is delivered after all messages belongs to the tables which are blocked by the DDL event // so we can make assumption that the all DMLs received before the DDL event. // since one table's events may be produced to the different partitions, so we have to flush all partitions. // if block the whole database, flush all tables, otherwise flush the blocked tables. tableIDs := make(map[int64]struct{}) switch ddl.GetBlockedTables().InfluenceType { - case commonEvent.InfluenceTypeDB, commonEvent.InfluenceTypeAll: + case event.InfluenceTypeDB, event.InfluenceTypeAll: for _, progress := range w.progresses { for tableID := range progress.eventsGroup { tableIDs[tableID] = struct{}{} } } - case commonEvent.InfluenceTypeNormal: + case event.InfluenceTypeNormal: for _, item := range ddl.GetBlockedTables().TableIDs { tableIDs[item] = struct{}{} } @@ -256,7 +234,7 @@ func (w *writer) getBlockTableIDs(ddl *commonEvent.DDLEvent) map[int64]struct{} // DDLs may be received out of commit-ts order (e.g. due to MQ delivery or buffering), so Write() sorts // ddlList by commit-ts before executing. ddlWithMaxCommitTs is a guard against per-table commit-ts // regressions: executing an older DDL after a newer one may corrupt downstream schema/DML ordering. -func (w *writer) appendDDL(ddl *commonEvent.DDLEvent) { +func (w *writer) appendDDL(ddl *event.DDLEvent) { // If commitTs goes backwards for a blocked table, ignore this DDL instead of applying it out of order. tableIDs := w.getBlockTableIDs(ddl) for tableID := range tableIDs { @@ -290,37 +268,22 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { var ( done = make(chan struct{}, 1) - total int flushed atomic.Int64 ) watermark := w.globalWatermark() - resolvedEvents := make([]*commonEvent.DMLEvent, 0) - // resolvedGroups records which EventsGroup has flushed events so we can - // advance its AppliedWatermark after the flush is fully finished. - resolvedGroups := make([]struct { - group *util.EventsGroup - maxCommitTs uint64 - }, 0) + resolvedEvents := make([]*event.DMLEvent, 0) for _, p := range w.progresses { for _, group := range p.eventsGroup { - before := len(resolvedEvents) - resolvedEvents = group.ResolveInto(watermark, resolvedEvents) - resolvedCount := len(resolvedEvents) - before - if resolvedCount == 0 { - continue + messages := group.ResolveInto(watermark, nil) + events := make([]*event.DMLEvent, 0, len(messages)) + for _, message := range messages { + events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) } - - resolvedGroups = append(resolvedGroups, struct { - group *util.EventsGroup - maxCommitTs uint64 - }{ - group: group, - maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), - }) - total += resolvedCount + resolvedEvents = append(resolvedEvents, events...) } } + total := len(resolvedEvents) if total == 0 { return nil } @@ -346,11 +309,6 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { case <-done: log.Info("flush DML events done", zap.Uint64("watermark", watermark), zap.Int("total", total), zap.Duration("duration", time.Since(start))) - for _, item := range resolvedGroups { - if item.maxCommitTs > item.group.AppliedWatermark { - item.group.AppliedWatermark = item.maxCommitTs - } - } return nil case <-ticker.C: log.Warn("DML events cannot be flushed in time", zap.Uint64("watermark", watermark), @@ -392,12 +350,12 @@ func (w *writer) WriteMessage(ctx context.Context, message *kafka.Message) bool ddl := progress.decoder.NextDDLEvent() if dec, ok := progress.decoder.(*simple.Decoder); ok { - cachedEvents := dec.GetCachedEvents() - for _, row := range cachedEvents { + cachedMessages := dec.GetCachedMessages() + for _, dmlMessage := range cachedMessages { log.Info("simple protocol cached event resolved, append to the group", - zap.Int64("tableID", row.GetTableID()), zap.Uint64("commitTs", row.CommitTs), + zap.Int64("tableID", dmlMessage.TableID), zap.Uint64("commitTs", dmlMessage.GetCommitTs()), zap.Int32("partition", partition), zap.Any("offset", offset)) - w.appendRow2Group(row, progress, offset) + w.appendMessage2Group(dmlMessage, progress, offset) } } @@ -421,25 +379,33 @@ func (w *writer) WriteMessage(ctx context.Context, message *kafka.Message) bool needFlush = true case common.MessageTypeRow: var counter int - row := progress.decoder.NextDMLEvent() - if row == nil { + dmlMessage := progress.decoder.NextDMLMessage() + if dmlMessage == nil { if w.protocol != config.ProtocolSimple { - log.Panic("DML event is nil, it's not expected", + log.Panic("DML message is nil, it's not expected", zap.Int32("partition", partition), zap.Any("offset", offset)) } - log.Debug("DML event is nil, it's cached", zap.Int32("partition", partition), zap.Any("offset", offset)) + log.Debug("DML message is nil, it's cached", zap.Int32("partition", partition), zap.Any("offset", offset)) break } - w.appendRow2Group(row, progress, offset) + w.appendMessage2Group(dmlMessage, progress, offset) counter++ for { _, hasNext = progress.decoder.HasNext() if !hasNext { break } - row = progress.decoder.NextDMLEvent() - w.appendRow2Group(row, progress, offset) + dmlMessage = progress.decoder.NextDMLMessage() + if dmlMessage == nil { + if w.protocol != config.ProtocolSimple { + log.Panic("DML message is nil, it's not expected", + zap.Int32("partition", partition), zap.Any("offset", offset)) + } + log.Debug("DML message is nil, it's cached", zap.Int32("partition", partition), zap.Any("offset", offset)) + break + } + w.appendMessage2Group(dmlMessage, progress, offset) counter++ } // If the message containing only one event exceeds the length limit, CDC will allow it and issue a warning. @@ -478,7 +444,7 @@ func (w *writer) Write(ctx context.Context, messageType common.MessageType) bool } watermark := w.globalWatermark() - ddlList := make([]*commonEvent.DDLEvent, 0) + ddlList := make([]*event.DDLEvent, 0) for i, todoDDL := range w.ddlList { // DDL ordering must follow commitTs (see appendDDL). Traditionally we wait until the global // resolved-ts (watermark) has reached the DDL commitTs, which guarantees all partitions have @@ -495,15 +461,15 @@ func (w *writer) Write(ctx context.Context, messageType common.MessageType) bool // populating BlockedTableNames and/or adding referenced table IDs (or partition IDs) into // BlockedTables.TableIDs. We only bypass watermark for CREATE TABLE when the DDL only blocks the // special DDL span and has no referenced blocked table names. - action := timodel.ActionType(todoDDL.Type) + action := model.ActionType(todoDDL.Type) bypassWatermark := false switch action { - case timodel.ActionCreateSchema: + case model.ActionCreateSchema: bypassWatermark = true - case timodel.ActionCreateTable: + case model.ActionCreateTable: blockedTables := todoDDL.GetBlockedTables() bypassWatermark = blockedTables != nil && - blockedTables.InfluenceType == commonEvent.InfluenceTypeNormal && + blockedTables.InfluenceType == event.InfluenceTypeNormal && len(blockedTables.TableIDs) == 1 && blockedTables.TableIDs[0] == commonType.DDLSpanTableID && len(todoDDL.GetBlockedTableNames()) == 0 @@ -537,7 +503,10 @@ func (w *writer) Write(ctx context.Context, messageType common.MessageType) bool return true } -func (w *writer) onDDL(ddl *commonEvent.DDLEvent) { +func (w *writer) onDDL(ddl *event.DDLEvent) { + if ddl.Query == "" { + return + } switch w.protocol { case config.ProtocolCanalJSON, config.ProtocolOpen, config.ProtocolAvro, config.ProtocolSimple, config.ProtocolDebezium, config.ProtocolDebeziumAvro: @@ -546,23 +515,57 @@ func (w *writer) onDDL(ddl *commonEvent.DDLEvent) { } // TODO: support more corner cases // e.g. create partition table + drop table(rename table) + create normal table: the partitionTableAccessor should drop the table when the table become normal. - switch timodel.ActionType(ddl.Type) { - case timodel.ActionCreateTable: + switch model.ActionType(ddl.Type) { + case model.ActionCreateTable: + if w.markPartitionTableFromDDL(ddl) { + return + } stmt, err := parser.New().ParseOneStmt(ddl.Query, "", "") if err != nil { log.Panic("parse ddl query failed", zap.String("query", ddl.Query), zap.Error(err)) } - if v, ok := stmt.(*ast.CreateTableStmt); ok && v.Partition != nil { - w.partitionTableAccessor.Add(ddl.GetSchemaName(), ddl.GetTableName()) + if v, ok := stmt.(*ast.CreateTableStmt); ok { + if v.Partition != nil { + w.addPartitionTable(ddl.GetSchemaName(), ddl.GetTableName()) + return + } + if v.ReferTable != nil { + referSchema := v.ReferTable.Schema.O + if referSchema == "" { + referSchema = ddl.GetSchemaName() + } + if w.partitionTableAccessor.IsPartitionTable(referSchema, v.ReferTable.Name.O) { + w.addPartitionTable(ddl.GetSchemaName(), ddl.GetTableName()) + } + } } - case timodel.ActionRenameTable: + case model.ActionRenameTable: if w.partitionTableAccessor.IsPartitionTable(ddl.ExtraSchemaName, ddl.ExtraTableName) { - w.partitionTableAccessor.Add(ddl.GetSchemaName(), ddl.GetTableName()) + w.addPartitionTable(ddl.GetSchemaName(), ddl.GetTableName()) } + w.markPartitionTableFromDDL(ddl) } } -func (w *writer) checkPartition(row *commonEvent.DMLEvent, partition int32, offset kafka.Offset) { +func (w *writer) markPartitionTableFromDDL(ddl *event.DDLEvent) bool { + if ddl.TableInfo == nil || !ddl.TableInfo.IsPartitionTable() { + return false + } + + w.addPartitionTable(ddl.GetSchemaName(), ddl.GetTableName()) + w.addPartitionTable(ddl.TableInfo.GetSchemaName(), ddl.TableInfo.GetTableName()) + w.addPartitionTable(ddl.TableInfo.GetSchemaName(), ddl.TableInfo.GetTableName()) + return true +} + +func (w *writer) addPartitionTable(schema, table string) { + if schema == "" || table == "" { + return + } + w.partitionTableAccessor.Add(schema, table) +} + +func (w *writer) checkPartition(row *event.DMLEvent, partition int32, offset kafka.Offset) { var ( partitioner = w.eventRouter.GetPartitionGenerator(row.TableInfo.GetSchemaName(), row.TableInfo.GetTableName()) partitionNum = int32(len(w.progresses)) @@ -589,54 +592,94 @@ func (w *writer) checkPartition(row *commonEvent.DMLEvent, partition int32, offs } } -func (w *writer) appendRow2Group(dml *commonEvent.DMLEvent, progress *partitionProgress, offset kafka.Offset) { - w.checkPartition(dml, progress.partition, offset) +func (w *writer) messageWithPartitionCheck(message *common.DMLMessage, partition int32, offset kafka.Offset) *common.DMLMessage { + return common.NewDMLMessage(message.TableID, message.Schema, message.Table, message.GetCommitTs(), message.RowType, func() *event.DMLEvent { + row := message.ToDMLEvent() + w.checkPartition(row, partition, offset) + return row + }) +} + +func (w *writer) appendMessage2Group(message *common.DMLMessage, progress *partitionProgress, offset kafka.Offset) { // if the kafka cluster is normal, this should not hit. // else if the cluster is abnormal, the consumer may consume old message, then cause the watermark fallback. var ( - tableID = dml.GetTableID() - schema = dml.TableInfo.GetSchemaName() - table = dml.TableInfo.GetTableName() - commitTs = dml.GetCommitTs() + tableID = message.TableID + schema = message.Schema + table = message.Table + commitTs = message.GetCommitTs() ) group := progress.eventsGroup[tableID] if group == nil { group = util.NewEventsGroup(progress.partition, tableID) progress.eventsGroup[tableID] = group } - // IMPORTANT: Kafka offsets are append-only, but CommitTs can go backwards after - // a TiCDC restart/retry (at-least-once replay). We must not drop such events - // solely based on a "seen" watermark (e.g. HighWatermark). The only safe - // ignore condition is "already flushed to downstream". - if commitTs <= group.AppliedWatermark { - log.Warn("DML event replayed after applied, ignore it", + if commitTs < progress.watermark { + log.Warn("DML Event fallback row, since less than the partition watermark, ignore it", zap.Int64("tableID", tableID), zap.Int32("partition", group.Partition), zap.Uint64("commitTs", commitTs), zap.Any("offset", offset), - zap.Uint64("appliedWatermark", group.AppliedWatermark), zap.Uint64("highWatermark", group.HighWatermark), - zap.Uint64("partitionWatermark", progress.watermark), zap.Any("watermarkOffset", progress.watermarkOffset), - zap.String("schema", schema), zap.String("table", table), zap.Any("protocol", w.protocol)) + zap.Uint64("watermark", progress.watermark), zap.Any("watermarkOffset", progress.watermarkOffset), + zap.String("schema", schema), zap.String("table", table)) return } - forceInsert := commitTs < group.HighWatermark || commitTs < progress.watermark || w.enableTableAcrossNodes - if forceInsert { - log.Warn("DML event commit ts fallback, append with forceInsert", + if commitTs >= group.HighWatermark { + message = w.messageWithPartitionCheck(message, progress.partition, offset) + group.AppendMessage(message, false) + log.Debug("DML event append to the group", zap.Int32("partition", group.Partition), zap.Any("offset", offset), - zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), - zap.Uint64("appliedWatermark", group.AppliedWatermark), - zap.Uint64("partitionWatermark", progress.watermark), zap.Any("watermarkOffset", progress.watermarkOffset), + zap.Uint64("commitTs", commitTs), zap.Uint64("HighWatermark", group.HighWatermark), zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), - zap.Stringer("eventType", dml.RowTypes[0]), zap.Any("protocol", w.protocol), - zap.Bool("IsPartition", dml.TableInfo.TableName.IsPartition)) - group.Append(dml, true) + zap.Stringer("eventType", message.RowType)) return } - group.Append(dml, false) - log.Info("DML event append to the group", - zap.Int32("partition", group.Partition), zap.Any("offset", offset), - zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), - zap.Uint64("appliedWatermark", group.AppliedWatermark), - zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), - zap.Stringer("eventType", dml.RowTypes[0])) + if w.enableTableAcrossNodes { + log.Warn("DML events fallback, but enableTableAcrossNodes is true, still append it", + zap.Int32("partition", group.Partition), zap.Any("offset", offset), + zap.Uint64("commitTs", commitTs), zap.Uint64("HighWatermark", group.HighWatermark), + zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), + zap.Stringer("eventType", message.RowType)) + group.AppendMessage(w.messageWithPartitionCheck(message, progress.partition, offset), true) + return + } + switch w.protocol { + case config.ProtocolSimple: + // simple protocol set the table id for all row message, it can be known which table the row message belongs to, + // also consider the table partition. + // open protocol set the partition table id if the table is partitioned. + // for normal table, the table id is generated by the fake table id generator by using schema and table name. + // so one event group for one normal table or one table partition, replayed messages can be ignored. + log.Warn("DML event fallback row, since less than the group high watermark, ignore it", + zap.Int32("partition", progress.partition), zap.Any("offset", offset), + zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), + zap.Any("partitionWatermark", progress.watermark), zap.Any("watermarkOffset", progress.watermarkOffset), + zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), + zap.Stringer("eventType", message.RowType), + // zap.Any("columns", row.Columns), zap.Any("preColumns", row.PreColumns), + zap.Any("protocol", w.protocol)) + case config.ProtocolCanalJSON, config.ProtocolOpen, config.ProtocolAvro, + config.ProtocolDebezium, config.ProtocolDebeziumAvro: + // for partition table, these protocols cannot assign physical table id to each dml message, + // we cannot distinguish whether it's a real fallback event or not, still append it. + if w.partitionTableAccessor.IsPartitionTable(schema, table) { + log.Warn("DML events fallback, but the table is a partition table, still append it", + zap.Int32("partition", group.Partition), zap.Any("offset", offset), + zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), + zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), + zap.Stringer("eventType", message.RowType), zap.Any("protocol", w.protocol)) + group.AppendMessage(w.messageWithPartitionCheck(message, progress.partition, offset), true) + return + } + log.Warn("DML event fallback row, since less than the group high watermark, ignore it", + zap.Int32("partition", progress.partition), zap.Any("offset", offset), + zap.Uint64("commitTs", commitTs), zap.Uint64("HighWatermark", group.HighWatermark), + zap.Any("partitionWatermark", progress.watermark), zap.Any("watermarkOffset", progress.watermarkOffset), + zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), + zap.Stringer("eventType", message.RowType), + // zap.Any("columns", row.Columns), zap.Any("preColumns", row.PreColumns), + zap.Any("protocol", w.protocol)) + default: + log.Panic("unknown protocol", zap.Any("protocol", w.protocol)) + } } func openDB(ctx context.Context, dsn string) (*sql.DB, error) { diff --git a/cmd/kafka-consumer/writer_test.go b/cmd/kafka-consumer/writer_test.go index c88cdea827..7396b2ec32 100644 --- a/cmd/kafka-consumer/writer_test.go +++ b/cmd/kafka-consumer/writer_test.go @@ -17,11 +17,16 @@ import ( "context" "testing" + "github.com/confluentinc/confluent-kafka-go/v2/kafka" + "github.com/pingcap/ticdc/cmd/util" "github.com/pingcap/ticdc/downstreamadapter/sink" + "github.com/pingcap/ticdc/downstreamadapter/sink/eventrouter" "github.com/pingcap/ticdc/pkg/common" commonEvent "github.com/pingcap/ticdc/pkg/common/event" + "github.com/pingcap/ticdc/pkg/config" codecCommon "github.com/pingcap/ticdc/pkg/sink/codec/common" timodel "github.com/pingcap/tidb/pkg/meta/model" + "github.com/pingcap/tidb/pkg/util/chunk" "github.com/stretchr/testify/require" ) @@ -278,3 +283,53 @@ func TestWriterWrite_handlesOutOfOrderDDLsByCommitTs(t *testing.T) { require.Len(t, w.ddlList, 1) require.Equal(t, "CREATE TABLE `common_1`.`a` (`a` BIGINT PRIMARY KEY,`b` INT)", w.ddlList[0].Query) } + +func TestAppendRow2GroupKeepsDebeziumPartitionTableFallback(t *testing.T) { + for _, protocol := range []config.Protocol{ + config.ProtocolDebezium, + config.ProtocolDebeziumAvro, + } { + t.Run(protocol.String(), func(t *testing.T) { + replicaCfg := config.GetDefaultReplicaConfig() + eventRouter, err := eventrouter.NewEventRouter(replicaCfg.Sink, "test-topic", false, false) + require.NoError(t, err) + + w := &writer{ + progresses: []*partitionProgress{{partition: 0, eventsGroup: make(map[int64]*util.EventsGroup)}}, + eventRouter: eventRouter, + protocol: protocol, + partitionTableAccessor: codecCommon.NewPartitionTableAccessor(), + } + + w.partitionTableAccessor.Add("target", "src") + ddl := &commonEvent.DDLEvent{ + Query: "CREATE TABLE `target`.`dst` LIKE `target`.`src`", + SchemaName: "target", + TableName: "dst", + Type: byte(timodel.ActionCreateTable), + } + w.onDDL(ddl) + require.True(t, w.partitionTableAccessor.IsPartitionTable("target", "dst")) + + newDMLEvent := func(commitTs uint64) *commonEvent.DMLEvent { + return &commonEvent.DMLEvent{ + PhysicalTableID: 1, + CommitTs: commitTs, + RowTypes: []common.RowType{common.RowTypeUpdate}, + Rows: chunk.NewChunkWithCapacity(nil, 0), + TableInfo: &common.TableInfo{ + TableName: common.TableName{Schema: "target", Table: "dst"}, + }, + } + } + + progress := w.progresses[0] + w.appendMessage2Group(codecCommon.NewDMLMessageFromEvent(newDMLEvent(200)), progress, kafka.Offset(10)) + w.appendMessage2Group(codecCommon.NewDMLMessageFromEvent(newDMLEvent(100)), progress, kafka.Offset(11)) + + resolved := progress.eventsGroup[1].ResolveInto(150, nil) + require.Len(t, resolved, 1) + require.Equal(t, uint64(100), resolved[0].GetCommitTs()) + }) + } +} diff --git a/cmd/pulsar-consumer/consumer.go b/cmd/pulsar-consumer/consumer.go index d23ed61675..2ae884ae5d 100644 --- a/cmd/pulsar-consumer/consumer.go +++ b/cmd/pulsar-consumer/consumer.go @@ -115,7 +115,7 @@ func (c *consumer) readMessage(ctx context.Context) error { if !needCommit { continue } - err := c.pulsarConsumer.AckID(consumerMsg.Message.ID()) + err := c.pulsarConsumer.AckIDCumulative(consumerMsg.ID()) if err != nil { log.Panic("Error ack message", zap.Error(err)) } diff --git a/cmd/pulsar-consumer/writer.go b/cmd/pulsar-consumer/writer.go index 5db4d3decd..418fc5fc35 100644 --- a/cmd/pulsar-consumer/writer.go +++ b/cmd/pulsar-consumer/writer.go @@ -143,43 +143,28 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *commonEvent.DDLEvent) e var ( done = make(chan struct{}, 1) - total int flushed atomic.Int64 ) tableIDs := w.getBlockTableIDs(ddl) commitTs := ddl.GetCommitTs() resolvedEvents := make([]*commonEvent.DMLEvent, 0) - // resolvedGroups records which EventsGroup has flushed events so we can - // advance its AppliedWatermark after the flush is fully finished. - resolvedGroups := make([]struct { - group *util.EventsGroup - maxCommitTs uint64 - }, 0) for tableID := range tableIDs { for _, progress := range w.progresses { g, ok := progress.eventsGroup[tableID] if !ok { continue } - before := len(resolvedEvents) - resolvedEvents = g.ResolveInto(commitTs, resolvedEvents) - resolvedCount := len(resolvedEvents) - before - if resolvedCount == 0 { - continue + messages := g.ResolveInto(commitTs, nil) + events := make([]*commonEvent.DMLEvent, 0, len(messages)) + for _, message := range messages { + events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) } - - resolvedGroups = append(resolvedGroups, struct { - group *util.EventsGroup - maxCommitTs uint64 - }{ - group: g, - maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), - }) - total += resolvedCount + resolvedEvents = append(resolvedEvents, events...) } } + total := len(resolvedEvents) if total == 0 { return w.mysqlSink.WriteBlockEvent(ddl) } @@ -204,11 +189,6 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *commonEvent.DDLEvent) e log.Info("flush DML events before DDL done", zap.Uint64("DDLCommitTs", commitTs), zap.Int("total", total), zap.Duration("duration", time.Since(start)), zap.Any("tables", tableIDs)) - for _, item := range resolvedGroups { - if item.maxCommitTs > item.group.AppliedWatermark { - item.group.AppliedWatermark = item.maxCommitTs - } - } return w.mysqlSink.WriteBlockEvent(ddl) case <-ticker.C: log.Warn("DML events cannot be flushed in time", @@ -280,37 +260,22 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { var ( done = make(chan struct{}, 1) - total int flushed atomic.Int64 ) watermark := w.globalWatermark() resolvedEvents := make([]*commonEvent.DMLEvent, 0) - // resolvedGroups records which EventsGroup has flushed events so we can - // advance its AppliedWatermark after the flush is fully finished. - resolvedGroups := make([]struct { - group *util.EventsGroup - maxCommitTs uint64 - }, 0) for _, p := range w.progresses { for _, group := range p.eventsGroup { - before := len(resolvedEvents) - resolvedEvents = group.ResolveInto(watermark, resolvedEvents) - resolvedCount := len(resolvedEvents) - before - if resolvedCount == 0 { - continue + messages := group.ResolveInto(watermark, nil) + events := make([]*commonEvent.DMLEvent, 0, len(messages)) + for _, message := range messages { + events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) } - - resolvedGroups = append(resolvedGroups, struct { - group *util.EventsGroup - maxCommitTs uint64 - }{ - group: group, - maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), - }) - total += resolvedCount + resolvedEvents = append(resolvedEvents, events...) } } + total := len(resolvedEvents) if total == 0 { return nil } @@ -334,11 +299,6 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { case <-done: log.Info("flush DML events done", zap.Uint64("watermark", watermark), zap.Int("total", total), zap.Duration("duration", time.Since(start))) - for _, item := range resolvedGroups { - if item.maxCommitTs > item.group.AppliedWatermark { - item.group.AppliedWatermark = item.maxCommitTs - } - } return nil case <-ticker.C: log.Warn("DML events cannot be flushed in time", zap.Uint64("watermark", watermark), @@ -387,12 +347,11 @@ func (w *writer) WriteMessage(ctx context.Context, message pulsar.Message) bool zap.Any("blockedTables", ddl.GetBlockedTables())) needFlush = true case common.MessageTypeRow: - row := progress.decoder.NextDMLEvent() - if row == nil { - log.Panic("DML event is nil, it's not expected") + dmlMessage := progress.decoder.NextDMLMessage() + if dmlMessage == nil { + log.Panic("DML message is nil, it's not expected") } - - w.appendRow2Group(row, progress) + w.appendMessage2Group(dmlMessage, progress) default: log.Panic("unknown message type", zap.Any("messageType", messageType)) } @@ -484,59 +443,110 @@ func (w *writer) onDDL(ddl *commonEvent.DDLEvent) { // e.g. create partition table + drop table(rename table) + create normal table: the partitionTableAccessor should drop the table when the table become normal. switch timodel.ActionType(ddl.Type) { case timodel.ActionCreateTable: + if w.markPartitionTableFromDDL(ddl) { + return + } stmt, err := parser.New().ParseOneStmt(ddl.Query, "", "") if err != nil { log.Panic("parse ddl query failed", zap.String("query", ddl.Query), zap.Error(err)) } - if v, ok := stmt.(*ast.CreateTableStmt); ok && v.Partition != nil { - w.partitionTableAccessor.Add(ddl.GetSchemaName(), ddl.GetTableName()) + if v, ok := stmt.(*ast.CreateTableStmt); ok { + if v.Partition != nil { + w.addPartitionTable(ddl.GetSchemaName(), ddl.GetTableName()) + return + } + if v.ReferTable != nil { + referSchema := v.ReferTable.Schema.O + if referSchema == "" { + referSchema = ddl.GetSchemaName() + } + if w.partitionTableAccessor.IsPartitionTable(referSchema, v.ReferTable.Name.O) { + w.addPartitionTable(ddl.GetSchemaName(), ddl.GetTableName()) + } + } } case timodel.ActionRenameTable: if w.partitionTableAccessor.IsPartitionTable(ddl.ExtraSchemaName, ddl.ExtraTableName) { - w.partitionTableAccessor.Add(ddl.GetSchemaName(), ddl.GetTableName()) + w.addPartitionTable(ddl.GetSchemaName(), ddl.GetTableName()) } + w.markPartitionTableFromDDL(ddl) } } -func (w *writer) appendRow2Group(dml *commonEvent.DMLEvent, progress *partitionProgress) { +func (w *writer) markPartitionTableFromDDL(ddl *commonEvent.DDLEvent) bool { + if ddl.TableInfo == nil || !ddl.TableInfo.IsPartitionTable() { + return false + } + + w.addPartitionTable(ddl.GetSchemaName(), ddl.GetTableName()) + w.addPartitionTable(ddl.TableInfo.GetSchemaName(), ddl.TableInfo.GetTableName()) + w.addPartitionTable(ddl.TableInfo.GetSchemaName(), ddl.TableInfo.GetTableName()) + return true +} + +func (w *writer) addPartitionTable(schema, table string) { + if schema == "" || table == "" { + return + } + w.partitionTableAccessor.Add(schema, table) +} + +func (w *writer) appendMessage2Group(message *common.DMLMessage, progress *partitionProgress) { var ( - tableID = dml.GetTableID() - schema = dml.TableInfo.GetSchemaName() - table = dml.TableInfo.GetTableName() - commitTs = dml.GetCommitTs() + tableID = message.TableID + schema = message.Schema + table = message.Table + commitTs = message.GetCommitTs() ) group := progress.eventsGroup[tableID] if group == nil { group = util.NewEventsGroup(progress.partition, tableID) progress.eventsGroup[tableID] = group } - if commitTs <= group.AppliedWatermark { - log.Warn("DML event replayed after applied, ignore it", + if commitTs < progress.watermark { + log.Warn("DML Event fallback row, since less than the partition watermark, ignore it", zap.Int64("tableID", tableID), zap.Int32("partition", group.Partition), - zap.Uint64("commitTs", commitTs), - zap.Uint64("appliedWatermark", group.AppliedWatermark), zap.Uint64("highWatermark", group.HighWatermark), - zap.Uint64("partitionWatermark", progress.watermark), - zap.String("schema", schema), zap.String("table", table), zap.Any("protocol", w.protocol)) + zap.Uint64("commitTs", commitTs), zap.Uint64("watermark", progress.watermark), + zap.String("schema", schema), zap.String("table", table)) + return + } + if commitTs >= group.HighWatermark { + group.AppendMessage(message, false) + log.Debug("DML event append to the group", + zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), + zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), + zap.Stringer("eventType", message.RowType)) return } - forceInsert := commitTs < group.HighWatermark || commitTs < progress.watermark || w.enableTableAcrossNodes - if forceInsert { - log.Warn("DML event commit ts fallback, append with forceInsert", - zap.Int32("partition", group.Partition), + if w.enableTableAcrossNodes { + log.Warn("DML events fallback, but enableTableAcrossNodes is true, still append it", zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), - zap.Uint64("appliedWatermark", group.AppliedWatermark), - zap.Uint64("partitionWatermark", progress.watermark), zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), - zap.Stringer("eventType", dml.RowTypes[0]), zap.Any("protocol", w.protocol), - zap.Bool("IsPartition", dml.TableInfo.TableName.IsPartition)) - group.Append(dml, true) + zap.Stringer("eventType", message.RowType)) + group.AppendMessage(message, true) return } - group.Append(dml, false) - log.Info("DML event append to the group", - zap.Int32("partition", group.Partition), - zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), - zap.Uint64("appliedWatermark", group.AppliedWatermark), - zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), - zap.Stringer("eventType", dml.RowTypes[0])) + switch w.protocol { + case config.ProtocolCanalJSON: + // for partition table, the canal-json message cannot assign physical table id to each dml message, + // we cannot distinguish whether it's a real fallback event or not, still append it. + isPartitionTable := w.partitionTableAccessor != nil && + w.partitionTableAccessor.IsPartitionTable(schema, table) + if isPartitionTable { + log.Warn("DML events fallback, but it's canal-json and partition table, still append it", + zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), + zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), + zap.Stringer("eventType", message.RowType)) + group.AppendMessage(message, true) + return + } + log.Warn("DML event fallback row, since less than the group high watermark, ignore it", + zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), + zap.Any("partitionWatermark", progress.watermark), zap.Any("watermark", progress.watermark), + zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), + zap.Stringer("eventType", message.RowType), + zap.Any("protocol", w.protocol), zap.Bool("IsPartition", isPartitionTable)) + default: + log.Panic("unknown protocol", zap.Any("protocol", w.protocol)) + } } diff --git a/cmd/pulsar-consumer/writer_test.go b/cmd/pulsar-consumer/writer_test.go index 8ed543763f..f071376b97 100644 --- a/cmd/pulsar-consumer/writer_test.go +++ b/cmd/pulsar-consumer/writer_test.go @@ -16,15 +16,18 @@ package main import ( "context" "testing" + "time" + "github.com/apache/pulsar-client-go/pulsar" + "github.com/golang/mock/gomock" "github.com/pingcap/ticdc/cmd/util" "github.com/pingcap/ticdc/downstreamadapter/sink" + sinkmock "github.com/pingcap/ticdc/downstreamadapter/sink/mock" "github.com/pingcap/ticdc/pkg/common" commonEvent "github.com/pingcap/ticdc/pkg/common/event" "github.com/pingcap/ticdc/pkg/config" codeccommon "github.com/pingcap/ticdc/pkg/sink/codec/common" timodel "github.com/pingcap/tidb/pkg/meta/model" - "github.com/pingcap/tidb/pkg/util/chunk" "github.com/stretchr/testify/require" ) @@ -282,57 +285,158 @@ func TestWriterWrite_handlesOutOfOrderDDLsByCommitTs(t *testing.T) { require.Equal(t, "CREATE TABLE `common_1`.`a` (`a` BIGINT PRIMARY KEY,`b` INT)", w.ddlList[0].Query) } -func TestAppendRow2Group_DoesNotDropCommitTsFallbackBeforeApplied(t *testing.T) { - // Scenario: - // 1) TiCDC writes DML messages to Pulsar in commitTs order. - // 2) Under network partition / changefeed restart, TiCDC may replay older commitTs - // at a later time (commitTs appears to go backwards). - // - // The pulsar-consumer must not drop these "fallback commitTs" events unless they - // have already been flushed to downstream (AppliedWatermark), otherwise replayed - // messages cannot heal missing windows. - w := &writer{ - progresses: []*partitionProgress{ - { - partition: 0, - eventsGroup: make(map[int64]*util.EventsGroup), - }, - }, - protocol: config.ProtocolCanalJSON, - partitionTableAccessor: codeccommon.NewPartitionTableAccessor(), - } +func TestWriteMessageDefersDMLAssemblyUntilFlush(t *testing.T) { + ctx := context.Background() + ctrl := gomock.NewController(t) + s := sinkmock.NewMockSink(ctrl) + s.EXPECT().AddDMLEvent(gomock.Any()).Do(func(event *commonEvent.DMLEvent) { + event.PostFlush() + }).Times(1) - newDMLEvent := func(tableID int64, commitTs uint64) *commonEvent.DMLEvent { - return &commonEvent.DMLEvent{ - PhysicalTableID: tableID, - CommitTs: commitTs, - RowTypes: []common.RowType{common.RowTypeUpdate}, - Rows: chunk.NewChunkWithCapacity(nil, 0), + decoder := &deferredDMLDecoder{ + row: &commonEvent.DMLEvent{ + PhysicalTableID: 1, + CommitTs: 100, + RowTypes: []common.RowType{common.RowTypeInsert}, TableInfo: &common.TableInfo{ - TableName: common.TableName{Schema: "test", Table: "t"}, + TableName: common.TableName{Schema: "test", Table: "t", TableID: 1}, }, - } + }, + } + progress := &partitionProgress{ + partition: 0, + eventsGroup: make(map[int64]*util.EventsGroup), + decoder: decoder, } + w := &writer{ + progresses: []*partitionProgress{progress}, + mysqlSink: s, + protocol: config.ProtocolCanalJSON, + } + + needCommit := w.WriteMessage(ctx, fakePulsarMessage{key: "k", payload: []byte(`{"fake":"row"}`)}) + require.False(t, needCommit) + require.Equal(t, 1, decoder.addKeyValueCount) + require.Equal(t, 1, decoder.hasNextCount) + require.Equal(t, 1, decoder.nextDMLMessageCount) + require.Zero(t, decoder.toDMLEventCount) + require.Len(t, progress.eventsGroup[1].ResolveInto(99, nil), 0) - progress := w.progresses[0] + progress.watermark = 100 + require.True(t, w.Write(ctx, codeccommon.MessageTypeResolved)) + require.Equal(t, 1, decoder.addKeyValueCount) + require.Equal(t, 1, decoder.hasNextCount) + require.Equal(t, 1, decoder.nextDMLMessageCount) + require.Equal(t, 1, decoder.toDMLEventCount) + require.Empty(t, progress.eventsGroup[1].ResolveInto(100, nil)) + require.Equal(t, []byte(`{"fake":"row"}`), decoder.lastValue) +} - // Step 1: observe a larger commitTs first (e.g. produced before restart). - w.appendRow2Group(newDMLEvent(1, 200), progress) +type deferredDMLDecoder struct { + row *commonEvent.DMLEvent - // Step 2: observe a smaller commitTs later (e.g. replayed after restart). - w.appendRow2Group(newDMLEvent(1, 100), progress) + addKeyValueCount int + hasNextCount int + nextDMLMessageCount int + toDMLEventCount int + lastValue []byte +} - group := progress.eventsGroup[1] - require.NotNil(t, group) +func (d *deferredDMLDecoder) AddKeyValue(_, value []byte) { + d.addKeyValueCount++ + d.lastValue = append(d.lastValue[:0], value...) +} - // Expect: commitTs=100 is still kept and can be resolved. - resolved := group.ResolveInto(150, nil) - require.Len(t, resolved, 1) - require.Equal(t, uint64(100), resolved[0].CommitTs) +func (d *deferredDMLDecoder) HasNext() (codeccommon.MessageType, bool) { + d.hasNextCount++ + return codeccommon.MessageTypeRow, true +} + +func (d *deferredDMLDecoder) NextResolvedEvent() uint64 { + return 0 +} + +func (d *deferredDMLDecoder) NextDMLMessage() *codeccommon.DMLMessage { + d.nextDMLMessageCount++ + return codeccommon.NewDMLMessage(1, "test", "t", 100, common.RowTypeInsert, func() *commonEvent.DMLEvent { + d.toDMLEventCount++ + return d.row + }) +} + +func (d *deferredDMLDecoder) NextDDLEvent() *commonEvent.DDLEvent { + return nil +} + +type fakePulsarMessage struct { + key string + payload []byte +} + +func (m fakePulsarMessage) Topic() string { + return "" +} + +func (m fakePulsarMessage) ProducerName() string { + return "" +} + +func (m fakePulsarMessage) Properties() map[string]string { + return nil +} - // Step 3: once downstream has flushed beyond commitTs=100, replay is safe to ignore. - group.AppliedWatermark = 200 - w.appendRow2Group(newDMLEvent(1, 100), progress) - resolved = group.ResolveInto(150, nil) - require.Empty(t, resolved) +func (m fakePulsarMessage) Payload() []byte { + return m.payload +} + +func (m fakePulsarMessage) ID() pulsar.MessageID { + return nil +} + +func (m fakePulsarMessage) PublishTime() time.Time { + return time.Time{} +} + +func (m fakePulsarMessage) EventTime() time.Time { + return time.Time{} +} + +func (m fakePulsarMessage) Key() string { + return m.key +} + +func (m fakePulsarMessage) OrderingKey() string { + return "" +} + +func (m fakePulsarMessage) RedeliveryCount() uint32 { + return 0 +} + +func (m fakePulsarMessage) IsReplicated() bool { + return false +} + +func (m fakePulsarMessage) GetReplicatedFrom() string { + return "" +} + +func (m fakePulsarMessage) GetSchemaValue(any) error { + return nil +} + +func (m fakePulsarMessage) SchemaVersion() []byte { + return nil +} + +func (m fakePulsarMessage) GetEncryptionContext() *pulsar.EncryptionContext { + return nil +} + +func (m fakePulsarMessage) Index() *uint64 { + return nil +} + +func (m fakePulsarMessage) BrokerPublishTime() *time.Time { + return nil } diff --git a/cmd/storage-consumer/consumer.go b/cmd/storage-consumer/consumer.go index 3fd3814c1b..1c7aafbe90 100644 --- a/cmd/storage-consumer/consumer.go +++ b/cmd/storage-consumer/consumer.go @@ -244,12 +244,12 @@ func (c *consumer) getNewFiles( return tableDMLMap, err } -func (c *consumer) appendRow2Group(dml *event.DMLEvent, enableTableAcrossNodes bool) { +func (c *consumer) appendMessage2Group(message *common.DMLMessage, enableTableAcrossNodes bool) { var ( - tableID = dml.GetTableID() - schema = dml.TableInfo.GetSchemaName() - table = dml.TableInfo.GetTableName() - commitTs = dml.GetCommitTs() + tableID = message.TableID + schema = message.Schema + table = message.Table + commitTs = message.GetCommitTs() ) group := c.eventsGroup[tableID] if group == nil { @@ -257,25 +257,26 @@ func (c *consumer) appendRow2Group(dml *event.DMLEvent, enableTableAcrossNodes b c.eventsGroup[tableID] = group } if commitTs >= group.HighWatermark { - group.Append(dml, false) - log.Info("DML event append to the group", + group.AppendMessage(message, false) + log.Debug("DML event append to the group", zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), - zap.Stringer("eventType", dml.RowTypes[0])) + zap.Stringer("eventType", message.RowType)) return } if enableTableAcrossNodes { log.Warn("DML events fallback, but enableTableAcrossNodes is true, still append it", zap.Uint64("commitTs", commitTs), zap.Uint64("highWatermark", group.HighWatermark), zap.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), - zap.Stringer("eventType", dml.RowTypes[0])) - group.Append(dml, true) + zap.Stringer("eventType", message.RowType)) + group.AppendMessage(message, true) return } log.Warn("dml event commit ts fallback, ignore", - zap.Uint64("commitTs", dml.CommitTs), + zap.Uint64("commitTs", commitTs), zap.Any("highWatermark", group.HighWatermark), - zap.Stringer("row", dml), + zap.String("schema", schema), + zap.String("table", table), ) } @@ -330,9 +331,8 @@ func (c *consumer) appendDMLEvents( if tp == common.MessageTypeRow { c.dmlCount.Add(1) - row := decoder.NextDMLEvent() - row.PhysicalTableID = tableID - c.appendRow2Group(row, fileIdx.EnableTableAcrossNodes) + message := decoder.NextDMLMessage() + c.appendMessage2Group(messageWithPhysicalTableID(message, tableID), fileIdx.EnableTableAcrossNodes) filteredCnt++ } } @@ -345,12 +345,27 @@ func (c *consumer) appendDMLEvents( return err } +func messageWithPhysicalTableID(message *common.DMLMessage, tableID int64) *common.DMLMessage { + return common.NewDMLMessage(tableID, message.Schema, message.Table, message.GetCommitTs(), message.RowType, func() *event.DMLEvent { + row := message.ToDMLEvent() + row.PhysicalTableID = tableID + return row + }) +} + func (c *consumer) flushDMLEvents(ctx context.Context, tableID int64) error { group := c.eventsGroup[tableID] if group == nil { return nil } - events := group.GetAllEvents() + messages := group.GetAllMessages() + if len(messages) == 0 { + return nil + } + events := make([]*event.DMLEvent, 0, len(messages)) + for _, message := range messages { + events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) + } total := len(events) if total == 0 { return nil diff --git a/cmd/util/event_group.go b/cmd/util/event_group.go index 95ff621510..b6eebe2039 100644 --- a/cmd/util/event_group.go +++ b/cmd/util/event_group.go @@ -14,26 +14,20 @@ package util import ( - "slices" "sort" "github.com/pingcap/log" - "github.com/pingcap/ticdc/pkg/common" commonEvent "github.com/pingcap/ticdc/pkg/common/event" + codeccommon "github.com/pingcap/ticdc/pkg/sink/codec/common" "go.uber.org/zap" ) -// EventsGroup could store change event message. +// EventsGroup stores change event messages. type EventsGroup struct { Partition int32 tableID int64 - events []*commonEvent.DMLEvent - // HighWatermark is the maximum CommitTs ever observed in this group. - // - // It is a "seen" watermark, not a "flushed/applied" watermark. When the - // consumer reads faster than it flushes to downstream, HighWatermark can be - // larger than AppliedWatermark. + messages []*codeccommon.DMLMessage HighWatermark uint64 // AppliedWatermark is the maximum CommitTs that has been successfully flushed // to the downstream for this group. @@ -49,108 +43,101 @@ func NewEventsGroup(partition int32, tableID int64) *EventsGroup { return &EventsGroup{ Partition: partition, tableID: tableID, - events: make([]*commonEvent.DMLEvent, 0, 1024), + messages: make([]*codeccommon.DMLMessage, 0, 1024), } } -// Append will append an event to event groups. -func (g *EventsGroup) Append(row *commonEvent.DMLEvent, force bool) { - if row.CommitTs > g.HighWatermark { - g.HighWatermark = row.CommitTs - } - - var lastDMLEvent *commonEvent.DMLEvent - if len(g.events) > 0 { - lastDMLEvent = g.events[len(g.events)-1] +// AppendMessage appends a message to event groups. +func (g *EventsGroup) AppendMessage(message *codeccommon.DMLMessage, force bool) { + commitTs := message.GetCommitTs() + if commitTs > g.HighWatermark { + g.HighWatermark = commitTs } - mergeDMLEvent := func(dst, src *commonEvent.DMLEvent) { - dst.Rows.Append(src.Rows, 0, src.Rows.NumRows()) - dst.RowTypes = append(dst.RowTypes, src.RowTypes...) - dst.Length += src.Length - dst.PostTxnFlushed = append(dst.PostTxnFlushed, src.PostTxnFlushed...) - } - - if lastDMLEvent == nil || lastDMLEvent.GetCommitTs() < row.GetCommitTs() { - g.events = append(g.events, row) - return + var lastMessage *codeccommon.DMLMessage + if len(g.messages) > 0 { + lastMessage = g.messages[len(g.messages)-1] } - if lastDMLEvent.GetCommitTs() == row.GetCommitTs() { - mergeDMLEvent(lastDMLEvent, row) + if lastMessage == nil || lastMessage.GetCommitTs() <= commitTs { + g.messages = append(g.messages, message) return } if force { - // A smaller CommitTs can appear at a larger Kafka offset after a TiCDC - // restart/retry (at-least-once replay). In this case we need to insert the - // event by CommitTs order. If the CommitTs already exists, merge it so one - // upstream transaction isn't split into multiple downstream transactions. - i := sort.Search(len(g.events), func(i int) bool { - return g.events[i].CommitTs > row.CommitTs + i := sort.Search(len(g.messages), func(i int) bool { + return g.messages[i].GetCommitTs() > commitTs }) - if i > 0 && g.events[i-1].CommitTs == row.CommitTs { - previous := g.events[i-1] - // If the table info version is incompatible, we cannot merge the events, - // and the event may be replayed again after the table info is updated. - // So we just skip this event to avoid potential panic in the downstream. - if !compareTableInfo(previous, row) { - log.Warn("skip replayed DML event due to incompatible table info, the event may be replayed again after the table info is updated", - zap.Int32("partition", g.Partition), zap.Int64("tableID", g.tableID), - zap.Uint64("commitTs", row.CommitTs), - zap.Any("previous", previous), - zap.Any("now", row)) - return - } - mergeDMLEvent(previous, row) - return - } - g.events = slices.Insert(g.events, i, row) + g.messages = append(g.messages, nil) + copy(g.messages[i+1:], g.messages[i:]) + g.messages[i] = message return } log.Panic("append event with smaller commit ts", zap.Int32("partition", g.Partition), zap.Int64("tableID", g.tableID), - zap.Uint64("lastCommitTs", lastDMLEvent.GetCommitTs()), zap.Uint64("commitTs", row.GetCommitTs())) + zap.Uint64("lastCommitTs", lastMessage.GetCommitTs()), zap.Uint64("commitTs", commitTs)) } -func compareTableInfo(previous, now *commonEvent.DMLEvent) bool { - previousInfo := previous.TableInfo.ToTiDBTableInfo() - nowInfo := now.TableInfo.ToTiDBTableInfo() - if previousInfo.UpdateTS > nowInfo.UpdateTS { - log.Panic("previous dml event has bigger table info version", - zap.Any("previous", previous), - zap.Any("now", now)) - } - return common.NewColumnSchema4Decoder(previousInfo).SameWithTableInfo(nowInfo) -} - -// ResolveInto appends all events with CommitTs <= resolve into dst and removes them from the group. -// ResolveInto copies pointers into dst first, then clears the -// resolved prefix so Go GC can reclaim resolved events once downstream is done with them. -func (g *EventsGroup) ResolveInto(resolve uint64, dst []*commonEvent.DMLEvent) []*commonEvent.DMLEvent { - i := sort.Search(len(g.events), func(i int) bool { - return g.events[i].CommitTs > resolve +// ResolveInto appends all messages with CommitTs <= resolve into dst and removes them from the group. +// ResolveInto copies pointers into dst first, then clears the resolved prefix so Go GC can reclaim +// resolved messages once downstream is done with them. +func (g *EventsGroup) ResolveInto(resolve uint64, dst []*codeccommon.DMLMessage) []*codeccommon.DMLMessage { + i := sort.Search(len(g.messages), func(i int) bool { + return g.messages[i].GetCommitTs() > resolve }) if i == 0 { return dst } // Copy pointers out first so we can safely clear the group's slice without affecting callers. - dst = append(dst, g.events[:i]...) - clear(g.events[:i]) - g.events = g.events[i:] - if len(g.events) != 0 { + dst = append(dst, g.messages[:i]...) + clear(g.messages[:i]) + g.messages = g.messages[i:] + if len(g.messages) != 0 { log.Debug("not all events resolved", zap.Int32("partition", g.Partition), zap.Int64("tableID", g.tableID), - zap.Int("resolved", i), zap.Int("remained", len(g.events)), - zap.Uint64("resolveTs", resolve), zap.Uint64("firstCommitTs", g.events[0].CommitTs)) + zap.Int("resolved", i), zap.Int("remained", len(g.messages)), + zap.Uint64("resolveTs", resolve), zap.Uint64("firstCommitTs", g.messages[0].GetCommitTs())) } return dst } -// GetAllEvents will get all events. -func (g *EventsGroup) GetAllEvents() []*commonEvent.DMLEvent { - result := g.events - g.events = nil +// GetAllMessages gets all messages. +func (g *EventsGroup) GetAllMessages() []*codeccommon.DMLMessage { + result := g.messages + g.messages = nil return result } + +// AppendOrMergeDMLEvent appends a DML event, or merges it into the previous event +// when both events belong to the same table group and have the same commit-ts. +func AppendOrMergeDMLEvent(events []*commonEvent.DMLEvent, row *commonEvent.DMLEvent) []*commonEvent.DMLEvent { + var lastDMLEvent *commonEvent.DMLEvent + if len(events) > 0 { + lastDMLEvent = events[len(events)-1] + } + + // mergeDMLEvent := func(dst, src *commonEvent.DMLEvent) { + // dst.Rows.Append(src.Rows, 0, src.Rows.NumRows()) + // dst.RowTypes = append(dst.RowTypes, src.RowTypes...) + // dst.Length += src.Length + // dst.PostTxnFlushed = append(dst.PostTxnFlushed, src.PostTxnFlushed...) + // } + + if lastDMLEvent == nil || lastDMLEvent.GetCommitTs() < row.GetCommitTs() { + return append(events, row) + } + + if lastDMLEvent.GetCommitTs() == row.GetCommitTs() { + lastDMLEvent.Rows.Append(row.Rows, 0, row.Rows.NumRows()) + lastDMLEvent.RowTypes = append(lastDMLEvent.RowTypes, row.RowTypes...) + lastDMLEvent.Length += row.Length + lastDMLEvent.PostTxnFlushed = append(lastDMLEvent.PostTxnFlushed, row.PostTxnFlushed...) + return events + } + + log.Panic("append event with smaller commit ts", + zap.Int64("tableID", row.GetTableID()), + zap.Uint64("lastCommitTs", lastDMLEvent.GetCommitTs()), zap.Uint64("commitTs", row.GetCommitTs())) + return events +} diff --git a/cmd/util/event_group_test.go b/cmd/util/event_group_test.go index f989cbaff9..0bb4f82fe9 100644 --- a/cmd/util/event_group_test.go +++ b/cmd/util/event_group_test.go @@ -18,55 +18,27 @@ import ( "github.com/pingcap/ticdc/pkg/common" commonEvent "github.com/pingcap/ticdc/pkg/common/event" - timodel "github.com/pingcap/tidb/pkg/meta/model" - "github.com/pingcap/tidb/pkg/parser/ast" + codeccommon "github.com/pingcap/ticdc/pkg/sink/codec/common" "github.com/pingcap/tidb/pkg/util/chunk" "github.com/stretchr/testify/require" ) -func TestEventsGroupAppendForceMergesExistingCommitTs(t *testing.T) { - // Scenario: - // 1) An upstream transaction (commitTs=100) is split into multiple messages. - // 2) Due to sink retry/restart, a later transaction (commitTs=200) is observed first. - // 3) A "late" fragment of the commitTs=100 transaction arrives afterwards. - // - // The EventsGroup must merge the late fragment into the existing commitTs=100 event, - // instead of turning it into a second commitTs=100 item (which would split one upstream - // transaction into multiple downstream transactions). - group := NewEventsGroup(0, 1) +func newTestDMLMessage(commitTs uint64) *codeccommon.DMLMessage { + return codeccommon.NewDMLMessage(1, "test", "t", commitTs, common.RowTypeInsert, nil) +} - newDMLEvent := func(commitTs uint64) *commonEvent.DMLEvent { - return &commonEvent.DMLEvent{ - CommitTs: commitTs, - RowTypes: []common.RowType{common.RowTypeUpdate}, - Rows: chunk.NewChunkWithCapacity(nil, 0), - Length: 0, - TableInfo: common.NewTableInfo4Decoder("test", &timodel.TableInfo{ - ID: 100, - Name: ast.NewCIStr("t"), - Columns: []*timodel.ColumnInfo{ - {Name: ast.NewCIStr("a")}, - }, - }), - } +func newTestDMLEvent(commitTs uint64, rowTypes ...common.RowType) *commonEvent.DMLEvent { + return &commonEvent.DMLEvent{ + PhysicalTableID: 1, + CommitTs: commitTs, + Length: int32(len(rowTypes)), + RowTypes: rowTypes, + Rows: chunk.NewChunkWithCapacity(nil, 0), } - - group.Append(newDMLEvent(100), false) - group.Append(newDMLEvent(200), false) - group.Append(newDMLEvent(100), true) - - require.Equal(t, uint64(200), group.HighWatermark) - - var dst []*commonEvent.DMLEvent - dst = group.ResolveInto(150, dst) - require.Len(t, dst, 1) - require.Equal(t, uint64(100), dst[0].CommitTs) - require.Len(t, dst[0].RowTypes, 2) } func TestEventsGroupResolveIntoAppendsAndClearsResolvedPrefix(t *testing.T) { // Scenario: A consumer resolves a prefix of events by watermark/commit-ts and appends them - // into a downstream batch slice. We must clear the resolved prefix in the group's backing // array to avoid retaining already-flushed events and causing unbounded memory growth. // // Steps: @@ -75,75 +47,106 @@ func TestEventsGroupResolveIntoAppendsAndClearsResolvedPrefix(t *testing.T) { // 3. Verify (a) returned events are correct, (b) group keeps only the remaining event, // (c) the resolved prefix in the original backing slice is cleared (nil'd). group := NewEventsGroup(0, 1) - e1 := &commonEvent.DMLEvent{CommitTs: 1} - e2 := &commonEvent.DMLEvent{CommitTs: 2} - e3 := &commonEvent.DMLEvent{CommitTs: 3} - group.Append(e1, false) - group.Append(e2, false) - group.Append(e3, false) + m1 := newTestDMLMessage(1) + m2 := newTestDMLMessage(2) + m3 := newTestDMLMessage(3) + group.AppendMessage(m1, false) + group.AppendMessage(m2, false) + group.AppendMessage(m3, false) // Keep a reference to the original slice header so we can validate that ResolveInto clears // the resolved prefix in-place (this is what prevents GC retention of flushed events). - original := group.events + original := group.messages - var dst []*commonEvent.DMLEvent + var dst []*codeccommon.DMLMessage dst = group.ResolveInto(2, dst) require.Len(t, dst, 2) - require.Same(t, e1, dst[0]) - require.Same(t, e2, dst[1]) + require.Same(t, m1, dst[0]) + require.Same(t, m2, dst[1]) - require.Len(t, group.events, 1) - require.Same(t, e3, group.events[0]) + require.Len(t, group.messages, 1) + require.Same(t, m3, group.messages[0]) // The resolved prefix must be nil so the group doesn't keep flushed events alive via its // backing array (classic Go slice memory retention pitfall). require.Nil(t, original[0]) require.Nil(t, original[1]) - require.Same(t, e3, original[2]) + require.Same(t, m3, original[2]) } func TestEventsGroupResolveIntoNoopWhenNothingResolved(t *testing.T) { // Scenario: resolveTs is behind all buffered events. // Expectation: ResolveInto should be a no-op (dst unchanged, group unchanged). group := NewEventsGroup(0, 1) - e1 := &commonEvent.DMLEvent{CommitTs: 10} - e2 := &commonEvent.DMLEvent{CommitTs: 20} - group.Append(e1, false) - group.Append(e2, false) + m1 := newTestDMLMessage(10) + m2 := newTestDMLMessage(20) + group.AppendMessage(m1, false) + group.AppendMessage(m2, false) - original := group.events - dst := make([]*commonEvent.DMLEvent, 0, 1) + original := group.messages + dst := make([]*codeccommon.DMLMessage, 0, 1) dst = group.ResolveInto(5, dst) require.Len(t, dst, 0) - require.Len(t, group.events, 2) - require.Same(t, e1, group.events[0]) - require.Same(t, e2, group.events[1]) + require.Len(t, group.messages, 2) + require.Same(t, m1, group.messages[0]) + require.Same(t, m2, group.messages[1]) // No prefix should be cleared because nothing was resolved. - require.Same(t, e1, original[0]) - require.Same(t, e2, original[1]) + require.Same(t, m1, original[0]) + require.Same(t, m2, original[1]) } func TestEventsGroupResolveIntoClearsAllWhenFullyResolved(t *testing.T) { // Scenario: resolveTs advances beyond all buffered events. // Expectation: group is emptied and all backing-array pointers for resolved events are cleared. group := NewEventsGroup(0, 1) - e1 := &commonEvent.DMLEvent{CommitTs: 1} - e2 := &commonEvent.DMLEvent{CommitTs: 2} - group.Append(e1, false) - group.Append(e2, false) + m1 := newTestDMLMessage(1) + m2 := newTestDMLMessage(2) + group.AppendMessage(m1, false) + group.AppendMessage(m2, false) - original := group.events - var dst []*commonEvent.DMLEvent + original := group.messages + var dst []*codeccommon.DMLMessage dst = group.ResolveInto(100, dst) require.Len(t, dst, 2) - require.Same(t, e1, dst[0]) - require.Same(t, e2, dst[1]) + require.Same(t, m1, dst[0]) + require.Same(t, m2, dst[1]) - require.Len(t, group.events, 0) + require.Len(t, group.messages, 0) require.Nil(t, original[0]) require.Nil(t, original[1]) } + +func TestAppendOrMergeDMLEventMergesSameCommitTs(t *testing.T) { + var flushed []int + e1 := newTestDMLEvent(10, common.RowTypeInsert) + e1.AddPostFlushFunc(func() { flushed = append(flushed, 1) }) + e2 := newTestDMLEvent(10, common.RowTypeDelete) + e2.AddPostFlushFunc(func() { flushed = append(flushed, 2) }) + + events := AppendOrMergeDMLEvent(nil, e1) + events = AppendOrMergeDMLEvent(events, e2) + + require.Len(t, events, 1) + require.Same(t, e1, events[0]) + require.Equal(t, int32(2), events[0].Length) + require.Equal(t, []common.RowType{common.RowTypeInsert, common.RowTypeDelete}, events[0].RowTypes) + + events[0].PostFlush() + require.Equal(t, []int{1, 2}, flushed) +} + +func TestAppendOrMergeDMLEventAppendsDifferentCommitTs(t *testing.T) { + e1 := newTestDMLEvent(10, common.RowTypeInsert) + e2 := newTestDMLEvent(20, common.RowTypeDelete) + + events := AppendOrMergeDMLEvent(nil, e1) + events = AppendOrMergeDMLEvent(events, e2) + + require.Len(t, events, 2) + require.Same(t, e1, events[0]) + require.Same(t, e2, events[1]) +} diff --git a/pkg/sink/codec/avro/arvo.go b/pkg/sink/codec/avro/arvo.go index 89c15736ea..a537495c1e 100644 --- a/pkg/sink/codec/avro/arvo.go +++ b/pkg/sink/codec/avro/arvo.go @@ -86,8 +86,12 @@ func (a *BatchEncoder) encodeKey(ctx context.Context, topic string, e *commonEve if len(index) == 0 { return nil, nil } + row := e.GetRows() + if e.IsDelete() { + row = e.GetPreRows() + } keyColumns := &avroEncodeInput{ - row: e.GetRows(), + row: row, index: index, colInfos: colInfos, columnselector: e.ColumnSelector, @@ -121,20 +125,31 @@ func (a *BatchEncoder) encodeKey(ctx context.Context, topic string, e *commonEve } func (a *BatchEncoder) encodeValue(ctx context.Context, topic string, e *commonEvent.RowEvent) ([]byte, error) { + row := e.GetRows() + colInfos := e.TableInfo.GetColumns() + var index []int if e.IsDelete() { - return nil, nil - } - length := e.GetRows().Len() - if length == 0 { - return nil, nil - } - index := make([]int, length) - for i := 0; i < length; i++ { - index[i] = i + if !a.config.EnableTiDBExtension || !a.config.AvroEnableWatermark { + return nil, nil + } + index, colInfos = e.PrimaryKeyColumn() + if len(index) == 0 { + return nil, nil + } + row = e.GetPreRows() + } else { + length := row.Len() + if length == 0 { + return nil, nil + } + index = make([]int, length) + for i := 0; i < length; i++ { + index[i] = i + } } input := &avroEncodeInput{ - row: e.GetRows(), - colInfos: e.TableInfo.GetColumns(), + row: row, + colInfos: colInfos, index: index, columnselector: e.ColumnSelector, } diff --git a/pkg/sink/codec/avro/avro_test.go b/pkg/sink/codec/avro/avro_test.go index 528bd92f97..778d453818 100644 --- a/pkg/sink/codec/avro/avro_test.go +++ b/pkg/sink/codec/avro/avro_test.go @@ -18,6 +18,7 @@ import ( "testing" "github.com/linkedin/goavro/v2" + commonType "github.com/pingcap/ticdc/pkg/common" "github.com/pingcap/ticdc/pkg/config" "github.com/pingcap/ticdc/pkg/sink/codec/common" "github.com/pingcap/ticdc/pkg/uuid" @@ -123,6 +124,78 @@ func TestAvroEncode(t *testing.T) { } } +func TestAvroEncodeDeleteEventUsesPreRowForKey(t *testing.T) { + codecConfig := common.NewConfig(config.ProtocolAvro) + codecConfig.EnableTiDBExtension = true + + ctx := t.Context() + + encoder, err := SetupEncoderAndSchemaRegistry4Testing(ctx, codecConfig) + defer TeardownEncoderAndSchemaRegistry4Testing() + require.NoError(t, err) + require.NotNil(t, encoder) + + _, _, _, event := common.NewLargeEvent4Test(t) + topic := "avro-delete-test-topic" + require.NoError(t, encoder.AppendRowChangedEvent(ctx, topic, event)) + + messages := encoder.Build() + require.Len(t, messages, 1) + require.NotEmpty(t, messages[0].Key) + require.Nil(t, messages[0].Value) + + cid, data, err := extractConfluentSchemaIDAndBinaryData(messages[0].Key) + require.NoError(t, err) + + avroKeyCodec, err := encoder.schemaM.Lookup(ctx, + topicName2SchemaSubjects(topic, keySchemaSuffix), + schemaID{confluentSchemaID: cid}) + require.NoError(t, err) + + res, _, err := avroKeyCodec.NativeFromBinary(data) + require.NoError(t, err) + require.NotNil(t, res) + require.Equal(t, int32(127), res.(map[string]any)["tu1"]) +} + +func TestAvroEncodeDeleteEventWithWatermarkCarriesCommitTs(t *testing.T) { + codecConfig := common.NewConfig(config.ProtocolAvro) + codecConfig.EnableTiDBExtension = true + codecConfig.AvroEnableWatermark = true + + ctx := t.Context() + + encoder, err := SetupEncoderAndSchemaRegistry4Testing(ctx, codecConfig) + defer TeardownEncoderAndSchemaRegistry4Testing() + require.NoError(t, err) + require.NotNil(t, encoder) + + _, _, _, event := common.NewLargeEvent4Test(t) + topic := "avro-delete-with-watermark-test-topic" + require.NoError(t, encoder.AppendRowChangedEvent(ctx, topic, event)) + + messages := encoder.Build() + require.Len(t, messages, 1) + require.NotEmpty(t, messages[0].Key) + require.NotEmpty(t, messages[0].Value) + + decoder := NewDecoder(codecConfig, 0, encoder.schemaM, topic, nil) + decoder.AddKeyValue(messages[0].Key, messages[0].Value) + + messageType, exists := decoder.HasNext() + require.True(t, exists) + require.Equal(t, common.MessageTypeRow, messageType) + + decoded := decoder.NextDMLMessage().ToDMLEvent() + require.NotNil(t, decoded) + require.Equal(t, event.CommitTs, decoded.GetCommitTs()) + require.Len(t, decoded.RowTypes, 1) + require.Equal(t, commonType.RowTypeDelete, decoded.RowTypes[0]) + + _, exists = decoder.HasNext() + require.False(t, exists) +} + func TestAvroEnvelope(t *testing.T) { t.Parallel() cManager := &confluentSchemaManager{} diff --git a/pkg/sink/codec/avro/decoder.go b/pkg/sink/codec/avro/decoder.go index ce239cc79b..0a187c8913 100644 --- a/pkg/sink/codec/avro/decoder.go +++ b/pkg/sink/codec/avro/decoder.go @@ -113,23 +113,45 @@ func (d *decoder) NextResolvedEvent() uint64 { return ts } -// NextDMLEvent returns the next row changed event if exists -func (d *decoder) NextDMLEvent() *commonEvent.DMLEvent { +// NextDMLMessage returns the next row changed message if exists +func (d *decoder) NextDMLMessage() *common.DMLMessage { + keyMap, valueMap, valueSchema, isDelete, deleteCommitTs := d.decodeDMLPayload() + schemaName, tableName := schemaAndTableName(valueSchema) + commitTs := deleteCommitTs + if commitTs == 0 && !isDelete { + commitTs = uint64(valueMap[tidbCommitTs].(int64)) + } + rowType := commonType.RowTypeInsert + if isDelete { + rowType = commonType.RowTypeDelete + } + tableID := tableIDAllocator.Allocate(schemaName, tableName) + return common.NewDMLMessage(tableID, schemaName, tableName, commitTs, rowType, func() *commonEvent.DMLEvent { + return d.assembleDMLEventFromDecoded(keyMap, valueMap, valueSchema, isDelete, deleteCommitTs) + }) +} + +func (d *decoder) decodeDMLPayload() ( + keyMap map[string]any, + valueMap map[string]any, + valueSchema map[string]any, + isDelete bool, + deleteCommitTs uint64, +) { var ( - valueMap map[string]interface{} - valueSchema map[string]interface{} - err error + keySchema map[string]any + err error ) ctx := context.Background() - keyMap, keySchema, err := d.decodeKey(ctx) + keyMap, keySchema, err = d.decodeKey(ctx) if err != nil { log.Panic("decode key failed", zap.Error(err)) } // for the delete event, only have key part, it holds primary key or the unique key columns. // for the insert / update, extract the value part, it holds all columns. - isDelete := len(d.value) == 0 + isDelete = len(d.value) == 0 if isDelete { // delete event only have key part, treat it as the value part also. valueMap = keyMap @@ -139,8 +161,24 @@ func (d *decoder) NextDMLEvent() *commonEvent.DMLEvent { if err != nil { log.Panic("decode value failed", zap.Error(err)) } + if op, ok := valueMap[tidbOp].(string); ok && op == deleteOperation { + isDelete = true + if commitTs, ok := valueMap[tidbCommitTs].(int64); ok { + deleteCommitTs = uint64(commitTs) + } + } } + return keyMap, valueMap, valueSchema, isDelete, deleteCommitTs +} + +func (d *decoder) assembleDMLEventFromDecoded( + keyMap map[string]any, + valueMap map[string]any, + valueSchema map[string]any, + isDelete bool, + deleteCommitTs uint64, +) *commonEvent.DMLEvent { event, err := assembleEvent(keyMap, valueMap, valueSchema, isDelete) if err != nil { log.Panic("assemble event failed", zap.Error(err)) @@ -149,6 +187,10 @@ func (d *decoder) NextDMLEvent() *commonEvent.DMLEvent { // Delete event only has Primary Key Columns, but the checksum is calculated based on the whole row columns, // checksum verification cannot be done here, so skip it. if isDelete { + if deleteCommitTs != 0 { + event.StartTs = deleteCommitTs + event.CommitTs = deleteCommitTs + } return event } @@ -188,18 +230,18 @@ func (d *decoder) NextDMLEvent() *commonEvent.DMLEvent { // valueMap hold all columns information // schema is corresponding to the valueMap, it can be used to decode the valueMap to construct columns. func assembleEvent( - keyMap, valueMap, schema map[string]interface{}, isDelete bool, + keyMap, valueMap, schema map[string]any, isDelete bool, ) (*commonEvent.DMLEvent, error) { - fields, ok := schema["fields"].([]interface{}) + fields, ok := schema["fields"].([]any) if !ok { return nil, errors.New("schema fields should be a map") } columns := make([]*timodel.ColumnInfo, 0, len(valueMap)) - data := make(map[string]interface{}, 0) + data := make(map[string]any, 0) // fields is ordered by the column id, so iterate over it to build columns // it's also the order to calculate the checksum. for idx, item := range fields { - field, ok := item.(map[string]interface{}) + field, ok := item.(map[string]any) if !ok { return nil, errors.New("schema field should be a map") } @@ -210,18 +252,18 @@ func assembleEvent( break } // query the field to get `tidbType`, and get the mysql type from it. - var holder map[string]interface{} + var holder map[string]any switch ty := field["type"].(type) { - case []interface{}: - if m, ok := ty[0].(map[string]interface{}); ok { - holder = m["connect.parameters"].(map[string]interface{}) - } else if m, ok := ty[1].(map[string]interface{}); ok { - holder = m["connect.parameters"].(map[string]interface{}) + case []any: + if m, ok := ty[0].(map[string]any); ok { + holder = m["connect.parameters"].(map[string]any) + } else if m, ok := ty[1].(map[string]any); ok { + holder = m["connect.parameters"].(map[string]any) } else { log.Panic("type info is anything else", zap.Any("typeInfo", field["type"])) } - case map[string]interface{}: - holder = ty["connect.parameters"].(map[string]interface{}) + case map[string]any: + holder = ty["connect.parameters"].(map[string]any) default: log.Panic("type info is anything else", zap.Any("typeInfo", field["type"])) } @@ -248,10 +290,7 @@ func assembleEvent( columns = append(columns, tiCol) } - // "namespace.schema" - namespace := schema["namespace"].(string) - schemaName := strings.Split(namespace, ".")[1] - tableName := schema["name"].(string) + schemaName, tableName := schemaAndTableName(schema) var commitTs int64 if !isDelete { @@ -282,12 +321,17 @@ func assembleEvent( return event, nil } -func queryTableInfo(schemaName, tableName string, columns []*timodel.ColumnInfo, keyMap map[string]interface{}) *commonType.TableInfo { +func schemaAndTableName(schema map[string]any) (string, string) { + namespace := schema["namespace"].(string) + return strings.Split(namespace, ".")[1], schema["name"].(string) +} + +func queryTableInfo(schemaName, tableName string, columns []*timodel.ColumnInfo, keyMap map[string]any) *commonType.TableInfo { tableInfo := newTableInfo(schemaName, tableName, columns, keyMap) return tableInfo } -func newTableInfo(schemaName, tableName string, columns []*timodel.ColumnInfo, keyMap map[string]interface{}) *commonType.TableInfo { +func newTableInfo(schemaName, tableName string, columns []*timodel.ColumnInfo, keyMap map[string]any) *commonType.TableInfo { tidbTableInfo := new(timodel.TableInfo) tidbTableInfo.ID = tableIDAllocator.Allocate(schemaName, tableName) tableIDAllocator.AddBlockTableID(schemaName, tableName, tidbTableInfo.ID) @@ -310,7 +354,7 @@ func newTableInfo(schemaName, tableName string, columns []*timodel.ColumnInfo, k return commonType.NewTableInfo4Decoder(schemaName, tidbTableInfo) } -func isCorrupted(valueMap map[string]interface{}) bool { +func isCorrupted(valueMap map[string]any) bool { o, ok := valueMap[tidbCorrupted] if !ok { return false @@ -322,7 +366,7 @@ func isCorrupted(valueMap map[string]interface{}) bool { // extract the checksum from the received value map // return true if the checksum found, and return error if the checksum is not valid -func extractExpectedChecksum(valueMap map[string]interface{}) (uint64, bool) { +func extractExpectedChecksum(valueMap map[string]any) (uint64, bool) { o, ok := valueMap[tidbRowLevelChecksum] if !ok { return 0, false @@ -341,12 +385,12 @@ func extractExpectedChecksum(valueMap map[string]interface{}) (uint64, bool) { // value is an interface, need to convert it to the real value with the help of type info. // holder has the value's column info. func getColumnValue( - value interface{}, holder map[string]interface{}, mysqlType byte, flag uint, -) (interface{}, error) { + value any, holder map[string]any, mysqlType byte, flag uint, +) (any, error) { switch t := value.(type) { // for nullable columns, the value is encoded as a map with one pair. // key is the encoded type, value is the encoded value, only care about the value here. - case map[string]interface{}: + case map[string]any: for _, v := range t { value = v } @@ -511,7 +555,7 @@ func extractGlueSchemaIDAndBinaryData(data []byte) (string, []byte, error) { func decodeRawBytes( ctx context.Context, schemaM SchemaManager, data []byte, topic string, -) (map[string]interface{}, map[string]interface{}, error) { +) (map[string]any, map[string]any, error) { var schemaID schemaID var binary []byte var err error @@ -545,12 +589,12 @@ func decodeRawBytes( return nil, nil, err } - result, ok := native.(map[string]interface{}) + result, ok := native.(map[string]any) if !ok { return nil, nil, errors.ErrCodecDecode.GenWithStack("raw avro message is not a map") } - schema := make(map[string]interface{}) + schema := make(map[string]any) if err := json.Unmarshal([]byte(codec.Schema()), &schema); err != nil { return nil, nil, errors.Trace(err) } @@ -558,13 +602,13 @@ func decodeRawBytes( return result, schema, nil } -func (d *decoder) decodeKey(ctx context.Context) (map[string]interface{}, map[string]interface{}, error) { +func (d *decoder) decodeKey(ctx context.Context) (map[string]any, map[string]any, error) { data := d.key d.key = nil return decodeRawBytes(ctx, d.schemaM, data, d.topic) } -func (d *decoder) decodeValue(ctx context.Context) (map[string]interface{}, map[string]interface{}, error) { +func (d *decoder) decodeValue(ctx context.Context) (map[string]any, map[string]any, error) { data := d.value d.value = nil return decodeRawBytes(ctx, d.schemaM, data, d.topic) diff --git a/pkg/sink/codec/avro/encoder_test.go b/pkg/sink/codec/avro/encoder_test.go index 01f2c44f23..33774a8acc 100644 --- a/pkg/sink/codec/avro/encoder_test.go +++ b/pkg/sink/codec/avro/encoder_test.go @@ -69,7 +69,7 @@ func TestDMLEventE2E(t *testing.T) { require.True(t, exist) require.Equal(t, common.MessageTypeRow, messageType) - decodedEvent := decoder.NextDMLEvent() + decodedEvent := decoder.NextDMLMessage().ToDMLEvent() require.NotNil(t, decodedEvent) require.NotZero(t, decodedEvent.GetTableID()) diff --git a/pkg/sink/codec/avro/helper.go b/pkg/sink/codec/avro/helper.go index fcd59aa943..52638795e1 100644 --- a/pkg/sink/codec/avro/helper.go +++ b/pkg/sink/codec/avro/helper.go @@ -43,6 +43,7 @@ const ( const ( insertOperation = "c" updateOperation = "u" + deleteOperation = "d" ) const ( @@ -148,6 +149,8 @@ func getOperation(e *commonEvent.RowEvent) string { return insertOperation } else if e.IsUpdate() { return updateOperation + } else if e.IsDelete() { + return deleteOperation } return "" } diff --git a/pkg/sink/codec/canal/canal_json_decoder.go b/pkg/sink/codec/canal/canal_json_decoder.go index f8b3ac9a84..fe750831b8 100644 --- a/pkg/sink/codec/canal/canal_json_decoder.go +++ b/pkg/sink/codec/canal/canal_json_decoder.go @@ -22,8 +22,10 @@ import ( "path/filepath" "reflect" "slices" + "sort" "strconv" "strings" + "sync" "github.com/pingcap/log" commonType "github.com/pingcap/ticdc/pkg/common" @@ -44,6 +46,12 @@ import ( ) type tableKey struct { + schema string + table string + ddlCommitTs uint64 +} + +type tableNameKey struct { schema string table string } @@ -91,7 +99,9 @@ type decoder struct { storage storeapi.Storage upstreamTiDB *sql.DB + tableInfoMu sync.RWMutex tableInfoCache map[tableKey]*commonType.TableInfo + ddlCommitTs map[tableNameKey][]uint64 } var tableIDAllocator = common.NewTableIDAllocator() @@ -125,6 +135,7 @@ func NewDecoder( storage: externalStorage, upstreamTiDB: db, tableInfoCache: make(map[tableKey]*commonType.TableInfo), + ddlCommitTs: make(map[tableNameKey][]uint64), }, nil } @@ -194,8 +205,7 @@ func (d *decoder) assembleClaimCheckDMLEvent( log.Panic("unmarshal claim check message failed", zap.Any("value", util.RedactAny(value)), zap.Error(err)) } - d.msg = message - return d.NextDMLEvent() + return d.decodeDMLMessage(message) } func buildData(holder *common.ColumnsHolder) (map[string]interface{}, map[string]string) { @@ -292,19 +302,46 @@ func (d *decoder) assembleHandleKeyOnlyDMLEvent( result.Data = []map[string]interface{}{data} } - d.msg = result - return d.NextDMLEvent() + return d.decodeDMLMessage(result) } -// NextDMLEvent implements the Decoder interface +// NextDMLMessage implements the Decoder interface // `HasNext` should be called before this. -func (d *decoder) NextDMLEvent() *commonEvent.DMLEvent { +func (d *decoder) NextDMLMessage() *common.DMLMessage { if d.msg == nil || d.msg.messageType() != common.MessageTypeRow { + messageType := common.MessageTypeUnknown + if d.msg != nil { + messageType = d.msg.messageType() + } log.Panic("message type is not row changed", - zap.Any("messageType", d.msg.messageType()), zap.Any("msg", d.msg)) + zap.Any("messageType", messageType), zap.Any("msg", d.msg)) + } + + msg := d.msg + schemaName := *msg.getSchema() + tableName := *msg.getTable() + tableID := tableIDAllocator.Allocate(schemaName, tableName) + tableIDAllocator.AddBlockTableID(schemaName, tableName, tableID) + + var rowType commonType.RowType + switch msg.eventType() { + case canal.EventType_DELETE: + rowType = commonType.RowTypeDelete + case canal.EventType_INSERT: + rowType = commonType.RowTypeInsert + case canal.EventType_UPDATE: + rowType = commonType.RowTypeUpdate + default: + log.Panic("unknown event type for the DML event", zap.Any("eventType", msg.eventType())) } - message, withExtension := d.msg.(*canalJSONMessageWithTiDBExtension) + return common.NewDMLMessage(tableID, schemaName, tableName, msg.getCommitTs(), rowType, func() *commonEvent.DMLEvent { + return d.decodeDMLMessage(msg) + }) +} + +func (d *decoder) decodeDMLMessage(msg canalJSONMessageInterface) *commonEvent.DMLEvent { + message, withExtension := msg.(*canalJSONMessageWithTiDBExtension) if withExtension { ctx := context.Background() if message.Extensions.OnlyHandleKey && d.upstreamTiDB != nil { @@ -314,11 +351,10 @@ func (d *decoder) NextDMLEvent() *commonEvent.DMLEvent { return d.assembleClaimCheckDMLEvent(ctx, message.Extensions.ClaimCheckLocation) } } - return d.canalJSONMessage2DMLEvent() + return d.canalJSONMessage2DMLEvent(msg) } -func (d *decoder) canalJSONMessage2DMLEvent() *commonEvent.DMLEvent { - msg := d.msg +func (d *decoder) canalJSONMessage2DMLEvent(msg canalJSONMessageInterface) *commonEvent.DMLEvent { tableInfo := d.queryTableInfo(msg) result := new(commonEvent.DMLEvent) @@ -378,9 +414,8 @@ func (d *decoder) NextDDLEvent() *commonEvent.DDLEvent { tableIDAllocator.AddBlockTableID(result.SchemaName, result.TableName, tableIDAllocator.Allocate(result.SchemaName, result.TableName)) result.BlockedTables = common.GetBlockedTables(tableIDAllocator, result) - // if receive a table level DDL, just remove the table info to trigger create a new one. - delete(d.tableInfoCache, tableKey{schema: result.SchemaName, table: result.TableName}) - delete(d.tableInfoCache, tableKey{schema: result.SchemaName, table: result.TableName}) + d.addDDLCommitTs(result.SchemaName, result.TableName, result.GetCommitTs()) + d.addDDLCommitTs(result.ExtraSchemaName, result.ExtraTableName, result.GetCommitTs()) return result } @@ -400,14 +435,15 @@ func (d *decoder) NextResolvedEvent() uint64 { } func formatAllColumnsValue(data map[string]any, columns []*timodel.ColumnInfo) map[string]any { + result := make(map[string]any, len(data)) for _, col := range columns { raw, ok := data[col.Name.O] if !ok { continue } - data[col.Name.O] = formatValue(raw, col.FieldType) + result[col.Name.O] = formatValue(raw, col.FieldType) } - return data + return result } func formatValue(value any, ft types.FieldType) any { @@ -536,9 +572,13 @@ func (d *decoder) queryTableInfo(msg canalJSONMessageInterface) *commonType.Tabl schemaName := *msg.getSchema() tableName := *msg.getTable() + d.tableInfoMu.Lock() + defer d.tableInfoMu.Unlock() + cacheKey := tableKey{ - schema: schemaName, - table: tableName, + schema: schemaName, + table: tableName, + ddlCommitTs: d.getDDLCommitTsLocked(schemaName, tableName, msg.getCommitTs()), } tableInfo, ok := d.tableInfoCache[cacheKey] if !ok { @@ -557,6 +597,41 @@ func (d *decoder) queryTableInfo(msg canalJSONMessageInterface) *commonType.Tabl return tableInfo } +func (d *decoder) addDDLCommitTs(schema, table string, commitTs uint64) { + if schema == "" || table == "" || commitTs == 0 { + return + } + + d.tableInfoMu.Lock() + defer d.tableInfoMu.Unlock() + + key := tableNameKey{schema: schema, table: table} + commitTsList := d.ddlCommitTs[key] + i := sort.Search(len(commitTsList), func(i int) bool { + return commitTsList[i] >= commitTs + }) + if i < len(commitTsList) && commitTsList[i] == commitTs { + return + } + d.ddlCommitTs[key] = slices.Insert(commitTsList, i, commitTs) +} + +func (d *decoder) getDDLCommitTsLocked(schema, table string, commitTs uint64) uint64 { + if commitTs == 0 { + return 0 + } + + commitTsList := d.ddlCommitTs[tableNameKey{schema: schema, table: table}] + i := sort.Search(len(commitTsList), func(i int) bool { + // DMLs with the same commit-ts as a DDL are flushed before that DDL. + return commitTsList[i] >= commitTs + }) + if i == 0 { + return 0 + } + return commitTsList[i-1] +} + func newTiColumns(msg canalJSONMessageInterface) []*timodel.ColumnInfo { type columnPair struct { mysqlType string diff --git a/pkg/sink/codec/canal/canal_json_encoder_test.go b/pkg/sink/codec/canal/canal_json_encoder_test.go index e84ebd109d..f9970974a2 100644 --- a/pkg/sink/codec/canal/canal_json_encoder_test.go +++ b/pkg/sink/codec/canal/canal_json_encoder_test.go @@ -82,7 +82,7 @@ func TestDMLE2E(t *testing.T) { require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - decodedEvent := dml2rowEvent(t, decoder.NextDMLEvent()) + decodedEvent := dml2rowEvent(t, decoder.NextDMLMessage().ToDMLEvent()) require.True(t, decodedEvent.IsInsert()) if enableTiDBExtension { require.Equal(t, insertEvent.CommitTs, decodedEvent.CommitTs) @@ -103,7 +103,7 @@ func TestDMLE2E(t *testing.T) { require.True(t, hasNext) require.EqualValues(t, messageType, common.MessageTypeRow) - decodedEvent = dml2rowEvent(t, decoder.NextDMLEvent()) + decodedEvent = dml2rowEvent(t, decoder.NextDMLMessage().ToDMLEvent()) require.True(t, decodedEvent.IsUpdate()) err = encoder.AppendRowChangedEvent(ctx, "", deleteEvent) @@ -117,7 +117,7 @@ func TestDMLE2E(t *testing.T) { require.True(t, hasNext) require.EqualValues(t, messageType, common.MessageTypeRow) - decodedEvent = dml2rowEvent(t, decoder.NextDMLEvent()) + decodedEvent = dml2rowEvent(t, decoder.NextDMLMessage().ToDMLEvent()) require.NoError(t, err) require.True(t, decodedEvent.IsDelete()) } @@ -156,7 +156,7 @@ func TestCanalJSONCompressionE2E(t *testing.T) { require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - decodedEvent := decoder.NextDMLEvent() + decodedEvent := decoder.NextDMLMessage().ToDMLEvent() require.Equal(t, decodedEvent.CommitTs, insertEvent.CommitTs) require.Equal(t, decodedEvent.TableInfo.GetSchemaName(), insertEvent.TableInfo.GetSchemaName()) require.Equal(t, decodedEvent.TableInfo.GetTableName(), insertEvent.TableInfo.GetTableName()) @@ -238,7 +238,7 @@ func TestCanalJSONClaimCheckE2E(t *testing.T) { require.Equal(t, messageType, common.MessageTypeRow) require.True(t, ok) - decodedLargeEvent := decoder.NextDMLEvent() + decodedLargeEvent := decoder.NextDMLMessage().ToDMLEvent() require.Equal(t, insertEvent.CommitTs, decodedLargeEvent.CommitTs) require.Equal(t, insertEvent.TableInfo.GetSchemaName(), decodedLargeEvent.TableInfo.GetSchemaName()) @@ -653,7 +653,7 @@ func TestCanalJSONContentCompatibleE2E(t *testing.T) { require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - decodedEvent := decoder.NextDMLEvent() + decodedEvent := decoder.NextDMLMessage().ToDMLEvent() require.NoError(t, err) require.Equal(t, decodedEvent.CommitTs, event.CommitTs) require.Equal(t, decodedEvent.TableInfo.GetSchemaName(), event.TableInfo.GetSchemaName()) @@ -714,7 +714,7 @@ func TestE2EPartitionTableByHash(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, tp) - decodedEvent := decoder.NextDMLEvent() + decodedEvent := decoder.NextDMLMessage().ToDMLEvent() require.NotZero(t, decodedEvent.GetTableID()) } @@ -771,7 +771,7 @@ func TestE2EPartitionTableByRange(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, tp) - decodedEvent := decoder.NextDMLEvent() + decodedEvent := decoder.NextDMLMessage().ToDMLEvent() require.NotZero(t, decodedEvent.GetTableID()) } @@ -835,7 +835,7 @@ func TestE2EPartitionTable(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, tp) - decodedEvent := decoder.NextDMLEvent() + decodedEvent := decoder.NextDMLMessage().ToDMLEvent() require.NotZero(t, decodedEvent.GetTableID()) rc, ok = insertEvent1.GetNextRow() @@ -855,7 +855,7 @@ func TestE2EPartitionTable(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, tp) - decodedEvent = decoder.NextDMLEvent() + decodedEvent = decoder.NextDMLMessage().ToDMLEvent() require.NotZero(t, decodedEvent.GetTableID()) rc, ok = insertEvent2.GetNextRow() @@ -875,7 +875,7 @@ func TestE2EPartitionTable(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, tp) - decodedEvent = decoder.NextDMLEvent() + decodedEvent = decoder.NextDMLMessage().ToDMLEvent() require.NotZero(t, decodedEvent.GetTableID()) } diff --git a/pkg/sink/codec/canal/canal_json_test.go b/pkg/sink/codec/canal/canal_json_test.go index aa6a242144..c92948d198 100644 --- a/pkg/sink/codec/canal/canal_json_test.go +++ b/pkg/sink/codec/canal/canal_json_test.go @@ -19,6 +19,7 @@ import ( "testing" "github.com/pingcap/ticdc/downstreamadapter/sink/columnselector" + commonType "github.com/pingcap/ticdc/pkg/common" commonEvent "github.com/pingcap/ticdc/pkg/common/event" "github.com/pingcap/ticdc/pkg/config" "github.com/pingcap/ticdc/pkg/config/kerneltype" @@ -84,7 +85,7 @@ func TestIntegerContentCompatible(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decodedInsert := decoder.NextDMLEvent() + decodedInsert := decoder.NextDMLMessage().ToDMLEvent() require.NotNil(t, decodedInsert) } @@ -169,7 +170,7 @@ func TestIntegerTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() if enableTiDBExtension { require.Equal(t, event.CommitTs, decoded.GetCommitTs()) } @@ -230,7 +231,7 @@ func TestFloatTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -279,7 +280,7 @@ func TestTimeTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -329,7 +330,7 @@ func TestStringTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -379,7 +380,7 @@ func TestBlobTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -429,7 +430,7 @@ func TestTextTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -488,7 +489,7 @@ func TestOtherTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -546,7 +547,7 @@ func TestDMLEventWithColumnSelector(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -605,7 +606,7 @@ func TestDMLMultiplePK(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -790,7 +791,7 @@ func TestLargeMessageClaimCheck(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decodedInsert := dec.NextDMLEvent() + decodedInsert := dec.NextDMLMessage().ToDMLEvent() require.NotNil(t, decodedInsert) change, ok := decodedInsert.GetNextRow() @@ -881,7 +882,7 @@ func TestMessageLargeHandleKeyOnly(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -972,7 +973,7 @@ func TestDMLTypeEvent(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() change, ok := decoded.GetNextRow() require.True(t, ok) @@ -998,7 +999,7 @@ func TestDMLTypeEvent(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() change, ok := decoded.GetNextRow() require.True(t, ok) @@ -1092,7 +1093,7 @@ func TestDDLSequence(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decodedInsert := dec.NextDMLEvent() + decodedInsert := dec.NextDMLMessage().ToDMLEvent() require.NotZero(t, decodedInsert.GetTableID()) addColumn := helper.DDL2Event(`alter table t add column c int`) @@ -1129,6 +1130,148 @@ func TestDDLSequence(t *testing.T) { require.Equal(t, obtained.GetBlockedTables().InfluenceType, commonEvent.InfluenceTypeNormal) } +func TestDecoderTableInfoCacheUsesDDLCommitTsAcrossColumnChanges(t *testing.T) { + ctx := context.Background() + codecConfig := common.NewConfig(config.ProtocolCanalJSON) + codecConfig.EnableTiDBExtension = true + + decoder, err := NewDecoder(ctx, codecConfig, nil) + require.NoError(t, err) + + buildRowMessage := func(commitTs uint64, mysqlTypes map[string]string) *canalJSONMessageWithTiDBExtension { + data := map[string]any{ + "data": "insert_1", + "id": "525", + } + if _, ok := mysqlTypes["new_col"]; ok { + data["new_col"] = nil + } + return &canalJSONMessageWithTiDBExtension{ + JSONMessage: &JSONMessage{ + Schema: "test", + Table: "table_5", + PKNames: []string{"id"}, + IsDDL: false, + EventType: "INSERT", + SQLType: map[string]int32{ + "data": 12, + "id": 4, + "new_col": 4, + }, + MySQLType: mysqlTypes, + Data: []map[string]any{data}, + }, + Extensions: &tidbExtension{CommitTs: commitTs}, + } + } + + decodeRow := func(message *canalJSONMessageWithTiDBExtension) *commonEvent.DMLEvent { + payload, err := json.Marshal(message) + require.NoError(t, err) + decoder.AddKeyValue(nil, payload) + messageType, hasNext := decoder.HasNext() + require.True(t, hasNext) + require.Equal(t, common.MessageTypeRow, messageType) + return decoder.NextDMLMessage().ToDMLEvent() + } + + columnNames := func(event *commonEvent.DMLEvent) []string { + columns := event.TableInfo.GetColumns() + result := make([]string, 0, len(columns)) + for _, column := range columns { + result = append(result, column.Name.O) + } + return result + } + + oldSchema := map[string]string{ + "data": "varchar(255)", + "id": "int", + } + newSchema := map[string]string{ + "data": "varchar(255)", + "id": "int", + "new_col": "int", + } + + oldEvent := decodeRow(buildRowMessage(100, oldSchema)) + require.Equal(t, []string{"data", "id"}, columnNames(oldEvent)) + + ddlMessage := &canalJSONMessageWithTiDBExtension{ + JSONMessage: &JSONMessage{ + Schema: "test", + Table: "table_5", + IsDDL: true, + EventType: "ALTER", + Query: "ALTER TABLE `test`.`table_5` ADD COLUMN `new_col` INT", + }, + Extensions: &tidbExtension{CommitTs: 200}, + } + payload, err := json.Marshal(ddlMessage) + require.NoError(t, err) + decoder.AddKeyValue(nil, payload) + messageType, hasNext := decoder.HasNext() + require.True(t, hasNext) + require.Equal(t, common.MessageTypeDDL, messageType) + ddl := decoder.NextDDLEvent() + require.Equal(t, oldEvent.GetTableID(), ddl.GetBlockedTables().TableIDs[0]) + + newEvent := decodeRow(buildRowMessage(300, newSchema)) + require.Equal(t, []string{"data", "id", "new_col"}, columnNames(newEvent)) + + lateOldEvent := decodeRow(buildRowMessage(100, oldSchema)) + require.Equal(t, []string{"data", "id"}, columnNames(lateOldEvent)) + require.Equal(t, oldEvent.GetTableID(), lateOldEvent.GetTableID()) + require.Equal(t, oldEvent.GetTableID(), newEvent.GetTableID()) +} + +func TestDecoderTableInfoCacheUsesDDLCommitTsBoundary(t *testing.T) { + tableIDAllocator.Clean() + dec := &decoder{ + tableInfoCache: make(map[tableKey]*commonType.TableInfo), + ddlCommitTs: make(map[tableNameKey][]uint64), + } + buildMessage := func(commitTs uint64, mysqlTypes map[string]string) *canalJSONMessageWithTiDBExtension { + return &canalJSONMessageWithTiDBExtension{ + JSONMessage: &JSONMessage{ + Schema: "test", + Table: "table_5", + PKNames: []string{"id"}, + MySQLType: mysqlTypes, + }, + Extensions: &tidbExtension{CommitTs: commitTs}, + } + } + columnNames := func(tableInfo *commonType.TableInfo) []string { + columns := tableInfo.GetColumns() + result := make([]string, 0, len(columns)) + for _, column := range columns { + result = append(result, column.Name.O) + } + return result + } + beforeDropSchema := map[string]string{ + "data": "varchar(255)", + "id": "int", + "new_col": "int", + } + afterDropSchema := map[string]string{ + "data": "varchar(255)", + "id": "int", + } + + dec.addDDLCommitTs("test", "table_5", 200) + sameTsAsDDL := dec.queryTableInfo(buildMessage(200, beforeDropSchema)) + afterDDL := dec.queryTableInfo(buildMessage(300, afterDropSchema)) + lateSameTsAsDDL := dec.queryTableInfo(buildMessage(200, beforeDropSchema)) + + require.Equal(t, []string{"data", "id", "new_col"}, columnNames(sameTsAsDDL)) + require.Equal(t, []string{"data", "id"}, columnNames(afterDDL)) + require.NotSame(t, sameTsAsDDL, afterDDL) + require.Same(t, sameTsAsDDL, lateSameTsAsDDL) + require.Equal(t, []uint64{200}, dec.ddlCommitTs[tableNameKey{schema: "test", table: "table_5"}]) +} + func TestCreateTableDDL(t *testing.T) { helper := commonEvent.NewEventTestHelper(t) defer helper.Close() diff --git a/pkg/sink/codec/canal/canal_json_txn_decoder.go b/pkg/sink/codec/canal/canal_json_txn_decoder.go index 02129a0fb6..9cdb70f383 100644 --- a/pkg/sink/codec/canal/canal_json_txn_decoder.go +++ b/pkg/sink/codec/canal/canal_json_txn_decoder.go @@ -96,19 +96,40 @@ func (d *txnDecoder) HasNext() (common.MessageType, bool) { return d.msg.messageType(), true } -func (d *txnDecoder) NextDMLEvent() *commonEvent.DMLEvent { +func (d *txnDecoder) NextDMLMessage() *common.DMLMessage { if d.msg == nil || d.msg.messageType() != common.MessageTypeRow { + messageType := common.MessageTypeUnknown + if d.msg != nil { + messageType = d.msg.messageType() + } log.Panic("message type is not row changed", - zap.Any("messageType", d.msg.messageType()), zap.Any("msg", d.msg)) + zap.Any("messageType", messageType), zap.Any("msg", d.msg)) } - result := d.canalJSONMessage2RowChange() - d.msg = nil - return result -} -func (d *txnDecoder) canalJSONMessage2RowChange() *commonEvent.DMLEvent { msg := d.msg + d.msg = nil + schemaName := *msg.getSchema() + tableName := *msg.getTable() + tableID := tableIDAllocator.Allocate(schemaName, tableName) + + var rowType commonType.RowType + switch msg.eventType() { + case canal.EventType_DELETE: + rowType = commonType.RowTypeDelete + case canal.EventType_INSERT: + rowType = commonType.RowTypeInsert + case canal.EventType_UPDATE: + rowType = commonType.RowTypeUpdate + default: + log.Panic("unknown event type for the DML event", zap.Any("eventType", msg.eventType())) + } + + return common.NewDMLMessage(tableID, schemaName, tableName, msg.getCommitTs(), rowType, func() *commonEvent.DMLEvent { + return d.canalJSONMessage2RowChange(msg) + }) +} +func (d *txnDecoder) canalJSONMessage2RowChange(msg canalJSONMessageInterface) *commonEvent.DMLEvent { tableInfo := newTableInfo(msg) result := new(commonEvent.DMLEvent) result.Length++ // todo: set this field correctly diff --git a/pkg/sink/codec/common/decoder.go b/pkg/sink/codec/common/decoder.go index d83219d7e7..86e5a8f5fa 100644 --- a/pkg/sink/codec/common/decoder.go +++ b/pkg/sink/codec/common/decoder.go @@ -14,9 +14,66 @@ package common import ( + commonType "github.com/pingcap/ticdc/pkg/common" commonEvent "github.com/pingcap/ticdc/pkg/common/event" ) +type DMLMessage struct { + TableID int64 + Schema string + Table string + RowType commonType.RowType + + commitTs uint64 + // toDMLEvent may be called after the decoder has consumed later messages. + // It must only use data captured by this DMLMessage and must not depend on decoder cursor state. + toDMLEvent func() *commonEvent.DMLEvent +} + +func NewDMLMessage( + tableID int64, + schema string, + table string, + commitTs uint64, + rowType commonType.RowType, + toDMLEvent func() *commonEvent.DMLEvent, +) *DMLMessage { + return &DMLMessage{ + TableID: tableID, + Schema: schema, + Table: table, + RowType: rowType, + commitTs: commitTs, + toDMLEvent: toDMLEvent, + } +} + +func NewDMLMessageFromEvent(event *commonEvent.DMLEvent) *DMLMessage { + var ( + schema string + table string + rowType commonType.RowType + ) + if event.TableInfo != nil { + schema = event.TableInfo.GetSchemaName() + table = event.TableInfo.GetTableName() + } + if len(event.RowTypes) > 0 { + rowType = event.RowTypes[0] + } + return NewDMLMessage(event.GetTableID(), schema, table, event.GetCommitTs(), rowType, func() *commonEvent.DMLEvent { + return event + }) +} + +func (m *DMLMessage) GetCommitTs() uint64 { + return m.commitTs +} + +func (m *DMLMessage) ToDMLEvent() *commonEvent.DMLEvent { + return m.toDMLEvent() +} + // Decoder is an abstraction for events decoder // this interface is only for testing now type Decoder interface { @@ -34,8 +91,8 @@ type Decoder interface { // NextResolvedEvent returns the next resolved event if exists NextResolvedEvent() uint64 - // NextDMLEvent returns the next DML event if exists - NextDMLEvent() *commonEvent.DMLEvent + // NextDMLMessage returns the next DML message if exists + NextDMLMessage() *DMLMessage // NextDDLEvent returns the next DDL event if exists NextDDLEvent() *commonEvent.DDLEvent diff --git a/pkg/sink/codec/common/table_info_cache.go b/pkg/sink/codec/common/table_info_cache.go index 9cccc22c1f..f1c7853ee6 100644 --- a/pkg/sink/codec/common/table_info_cache.go +++ b/pkg/sink/codec/common/table_info_cache.go @@ -15,6 +15,7 @@ package common import ( "strings" + "sync" "github.com/pingcap/log" "go.uber.org/zap" @@ -39,6 +40,7 @@ func newAccessKey(schema, table string) accessKey { // tableIDAllocator is a fake table id allocator type tableIDAllocator struct { + mu sync.RWMutex tableIDs map[accessKey]int64 currentTableID int64 blockedTableIDs map[accessKey]map[int64]struct{} @@ -63,11 +65,17 @@ func (a *tableIDAllocator) allocateByKey(key accessKey) int64 { // Allocate allocates a table id func (a *tableIDAllocator) Allocate(schema, table string) int64 { + a.mu.Lock() + defer a.mu.Unlock() + key := newAccessKey(schema, table) return a.allocateByKey(key) } func (a *tableIDAllocator) GetBlockedTables(schema, table string) []int64 { + a.mu.RLock() + defer a.mu.RUnlock() + key := newAccessKey(schema, table) blocked := a.blockedTableIDs[key] result := make([]int64, 0, len(blocked)) @@ -78,6 +86,9 @@ func (a *tableIDAllocator) GetBlockedTables(schema, table string) []int64 { } func (a *tableIDAllocator) AddBlockTableID(schema string, table string, physicalTableID int64) { + a.mu.Lock() + defer a.mu.Unlock() + key := newAccessKey(schema, table) if _, ok := a.blockedTableIDs[key]; !ok { a.blockedTableIDs[key] = make(map[int64]struct{}) @@ -91,6 +102,9 @@ func (a *tableIDAllocator) AddBlockTableID(schema string, table string, physical } func (a *tableIDAllocator) Clean() { + a.mu.Lock() + defer a.mu.Unlock() + a.currentTableID = 0 clear(a.tableIDs) clear(a.blockedTableIDs) diff --git a/pkg/sink/codec/csv/csv_decoder.go b/pkg/sink/codec/csv/csv_decoder.go index 05109c687a..62bfd7b56e 100644 --- a/pkg/sink/codec/csv/csv_decoder.go +++ b/pkg/sink/codec/csv/csv_decoder.go @@ -129,17 +129,34 @@ func (b *decoder) NextResolvedEvent() uint64 { return 0 } -// NextDMLEvent implements the Decoder interface. -func (b *decoder) NextDMLEvent() *commonEvent.DMLEvent { +// NextDMLMessage implements the Decoder interface. +func (b *decoder) NextDMLMessage() *common.DMLMessage { if b.closed { log.Panic("batch decoder is closed, cannot fetch the next DML event") } - e, err := csvMsg2RowChangedEvent(b.codecConfig, b.msg, b.tableInfo) - if err != nil { - log.Panic("convert message to event failed", zap.Error(err)) + msg := *b.msg + msg.columns = append([]any(nil), b.msg.columns...) + msg.preColumns = append([]any(nil), b.msg.preColumns...) + + rowType := commonType.RowTypeInsert + if msg.opType == operationDelete { + rowType = commonType.RowTypeDelete } - return e + + return common.NewDMLMessage( + b.tableInfo.TableName.TableID, + msg.schemaName, + msg.tableName, + msg.commitTs, + rowType, + func() *commonEvent.DMLEvent { + e, err := csvMsg2RowChangedEvent(b.codecConfig, &msg, b.tableInfo) + if err != nil { + log.Panic("convert message to event failed", zap.Error(err)) + } + return e + }) } // NextDDLEvent implements the Decoder interface. diff --git a/pkg/sink/codec/csv/csv_decoder_test.go b/pkg/sink/codec/csv/csv_decoder_test.go index 7faeaebdc8..d974f4ecd4 100644 --- a/pkg/sink/codec/csv/csv_decoder_test.go +++ b/pkg/sink/codec/csv/csv_decoder_test.go @@ -49,7 +49,7 @@ func TestCSVBatchDecoder(t *testing.T) { tp, hasNext := decoder.HasNext() require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, tp) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() require.NotNil(t, event) } diff --git a/pkg/sink/codec/debezium/avro_decoder.go b/pkg/sink/codec/debezium/avro_decoder.go index 083fc3481d..3e735745b1 100644 --- a/pkg/sink/codec/debezium/avro_decoder.go +++ b/pkg/sink/codec/debezium/avro_decoder.go @@ -95,8 +95,8 @@ func (d *avroDecoder) NextResolvedEvent() uint64 { return d.inner.NextResolvedEvent() } -func (d *avroDecoder) NextDMLEvent() *commonEvent.DMLEvent { - return d.inner.NextDMLEvent() +func (d *avroDecoder) NextDMLMessage() *common.DMLMessage { + return d.inner.NextDMLMessage() } func (d *avroDecoder) NextDDLEvent() *commonEvent.DDLEvent { diff --git a/pkg/sink/codec/debezium/avro_test.go b/pkg/sink/codec/debezium/avro_test.go index e5f72fff0f..123815420c 100644 --- a/pkg/sink/codec/debezium/avro_test.go +++ b/pkg/sink/codec/debezium/avro_test.go @@ -237,7 +237,7 @@ func TestDebeziumConfluentAvroDecodeRowEvent(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() require.Equal(t, commitTs, decoded.CommitTs) require.Equal(t, "test", decoded.TableInfo.GetSchemaName()) require.Equal(t, "foo", decoded.TableInfo.GetTableName()) @@ -308,7 +308,7 @@ func TestDebeziumConfluentAvroDecodeAccountDMLEvents(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() require.Equal(t, "test", decoded.TableInfo.GetSchemaName()) require.Equal(t, "tp_account", decoded.TableInfo.GetTableName()) diff --git a/pkg/sink/codec/debezium/decoder.go b/pkg/sink/codec/debezium/decoder.go index 7dc91f377a..5b80cb8047 100644 --- a/pkg/sink/codec/debezium/decoder.go +++ b/pkg/sink/codec/debezium/decoder.go @@ -48,10 +48,10 @@ type decoder struct { upstreamTiDB *sql.DB - keyPayload map[string]interface{} - keySchema map[string]interface{} - valuePayload map[string]interface{} - valueSchema map[string]interface{} + keyPayload map[string]any + keySchema map[string]any + valuePayload map[string]any + valueSchema map[string]any } // NewDecoder return an debezium decoder @@ -150,20 +150,58 @@ func (d *decoder) NextDDLEvent() *commonEvent.DDLEvent { return event } -// NextDMLEvent returns the next dml event if exists -func (d *decoder) NextDMLEvent() *commonEvent.DMLEvent { +// NextDMLMessage returns the next dml message if exists +func (d *decoder) NextDMLMessage() *common.DMLMessage { if len(d.valuePayload) == 0 { - log.Panic("next DML event failed, since value payload is empty") + log.Panic("next DML message failed, since value payload is empty") } if d.config.DebeziumDisableSchema { - log.Panic("next DML event failed, since DebeziumDisableSchema is true") + log.Panic("next DML message failed, since DebeziumDisableSchema is true") } if !d.config.EnableTiDBExtension { - log.Panic("next DML event failed, since EnableTiDBExtension is false") + log.Panic("next DML message failed, since EnableTiDBExtension is false") } - defer d.clear() - tableInfo := d.queryTableInfo() - commitTs := d.getCommitTs() + + keyPayload := d.keyPayload + valuePayload := d.valuePayload + valueSchema := d.valueSchema + commitTs := getCommitTsFromPayload(valuePayload) + schemaName := getSchemaNameFromPayload(valuePayload) + tableName := getTableNameFromPayload(valuePayload) + rowType := rowTypeFromPayload(valuePayload) + tableID := tableIDAllocator.Allocate(schemaName, tableName) + d.clear() + + return common.NewDMLMessage(tableID, schemaName, tableName, commitTs, rowType, func() *commonEvent.DMLEvent { + return d.assembleDMLEventFromPayload(keyPayload, valuePayload, valueSchema) + }) +} + +func rowTypeFromPayload(valuePayload map[string]any) commonType.RowType { + op, ok := valuePayload["op"] + if !ok { + log.Panic("DML message op not found") + } + switch op { + case "c": + return commonType.RowTypeInsert + case "u": + return commonType.RowTypeUpdate + case "d": + return commonType.RowTypeDelete + default: + log.Panic("unknown op for the DML message", zap.Any("op", op)) + } + return commonType.RowTypeInsert +} + +func (d *decoder) assembleDMLEventFromPayload( + keyPayload map[string]any, + valuePayload map[string]any, + valueSchema map[string]any, +) *commonEvent.DMLEvent { + tableInfo := queryTableInfoFromPayload(keyPayload, valuePayload, valueSchema) + commitTs := getCommitTsFromPayload(valuePayload) event := &commonEvent.DMLEvent{ Rows: chunk.NewChunkFromPoolWithCapacity(tableInfo.GetFieldSlice(), chunk.InitialCapacity), StartTs: commitTs, @@ -176,12 +214,12 @@ func (d *decoder) NextDMLEvent() *commonEvent.DMLEvent { event.Rows.Destroy(chunk.InitialCapacity, tableInfo.GetFieldSlice()) }) columns := tableInfo.GetColumns() - before, ok1 := d.valuePayload["before"].(map[string]interface{}) + before, ok1 := valuePayload["before"].(map[string]any) if ok1 { data := assembleColumnData(before, columns, d.config.TimeZone) common.AppendRow2Chunk(data, columns, event.Rows) } - after, ok2 := d.valuePayload["after"].(map[string]interface{}) + after, ok2 := valuePayload["after"].(map[string]any) if ok2 { data := assembleColumnData(after, columns, d.config.TimeZone) common.AppendRow2Chunk(data, columns, event.Rows) @@ -200,7 +238,11 @@ func (d *decoder) NextDMLEvent() *commonEvent.DMLEvent { } func (d *decoder) getCommitTs() uint64 { - source := d.valuePayload["source"].(map[string]interface{}) + return getCommitTsFromPayload(d.valuePayload) +} + +func getCommitTsFromPayload(valuePayload map[string]any) uint64 { + source := valuePayload["source"].(map[string]any) commitTs, err := source["commit_ts"].(json.Number).Int64() if err != nil { log.Error("decode value failed", zap.Error(err), zap.String("value", util.RedactAny(source))) @@ -209,15 +251,21 @@ func (d *decoder) getCommitTs() uint64 { } func (d *decoder) getSchemaName() string { - source := d.valuePayload["source"].(map[string]interface{}) - schemaName := source["db"].(string) - return schemaName + return getSchemaNameFromPayload(d.valuePayload) +} + +func getSchemaNameFromPayload(valuePayload map[string]any) string { + source := valuePayload["source"].(map[string]any) + return source["db"].(string) } func (d *decoder) getTableName() string { - source := d.valuePayload["source"].(map[string]interface{}) - tableName := source["table"].(string) - return tableName + return getTableNameFromPayload(d.valuePayload) +} + +func getTableNameFromPayload(valuePayload map[string]any) string { + source := valuePayload["source"].(map[string]any) + return source["table"].(string) } func (d *decoder) clear() { @@ -227,28 +275,31 @@ func (d *decoder) clear() { d.valueSchema = nil } -func (d *decoder) queryTableInfo() *commonType.TableInfo { - schemaName := d.getSchemaName() - tableName := d.getTableName() - +func queryTableInfoFromPayload( + keyPayload map[string]any, + valuePayload map[string]any, + valueSchema map[string]any, +) *commonType.TableInfo { + schemaName := getSchemaNameFromPayload(valuePayload) + tableName := getTableNameFromPayload(valuePayload) tidbTableInfo := new(timodel.TableInfo) tidbTableInfo.ID = tableIDAllocator.Allocate(schemaName, tableName) tableIDAllocator.AddBlockTableID(schemaName, tableName, tidbTableInfo.ID) tidbTableInfo.Name = ast.NewCIStr(tableName) - fields := d.valueSchema["fields"].([]interface{}) - after := fields[1].(map[string]interface{}) - columnsField := after["fields"].([]interface{}) - indexColumns := make([]*timodel.IndexColumn, 0, len(d.keyPayload)) + fields := valueSchema["fields"].([]any) + after := fields[1].(map[string]any) + columnsField := after["fields"].([]any) + indexColumns := make([]*timodel.IndexColumn, 0, len(keyPayload)) for idx, column := range columnsField { - col := column.(map[string]interface{}) + col := column.(map[string]any) colName := col["field"].(string) tidbType := col["tidb_type"].(string) optional := col["optional"].(bool) fieldType := parseTiDBType(tidbType, optional) switch fieldType.GetType() { case mysql.TypeEnum, mysql.TypeSet: - parameters := col["parameters"].(map[string]interface{}) + parameters := col["parameters"].(map[string]any) allowed := parameters["allowed"].(string) fieldType.SetElems(strings.Split(allowed, ",")) case mysql.TypeDatetime: @@ -257,7 +308,7 @@ func (d *decoder) queryTableInfo() *commonType.TableInfo { fieldType.SetDecimal(6) } } - if _, ok := d.keyPayload[colName]; ok { + if _, ok := keyPayload[colName]; ok { indexColumns = append(indexColumns, &timodel.IndexColumn{ Name: ast.NewCIStr(colName), Offset: idx, @@ -282,8 +333,8 @@ func (d *decoder) queryTableInfo() *commonType.TableInfo { return result } -func assembleColumnData(data map[string]interface{}, columns []*timodel.ColumnInfo, timeZone *time.Location) map[string]interface{} { - result := make(map[string]interface{}, 0) +func assembleColumnData(data map[string]any, columns []*timodel.ColumnInfo, timeZone *time.Location) map[string]any { + result := make(map[string]any, 0) for _, col := range columns { val, ok := data[col.Name.O] if !ok { @@ -294,7 +345,7 @@ func assembleColumnData(data map[string]interface{}, columns []*timodel.ColumnIn return result } -func decodeColumn(value interface{}, colInfo *timodel.ColumnInfo, timeZone *time.Location) interface{} { +func decodeColumn(value any, colInfo *timodel.ColumnInfo, timeZone *time.Location) any { if value == nil { return value } @@ -462,18 +513,18 @@ func parseTiDBType(tidbType string, optional bool) *ptypes.FieldType { return ft } -func decodeRawBytes(data []byte) (map[string]interface{}, map[string]interface{}, error) { - var v map[string]interface{} +func decodeRawBytes(data []byte) (map[string]any, map[string]any, error) { + var v map[string]any d := json.NewDecoder(bytes.NewBuffer(data)) d.UseNumber() if err := d.Decode(&v); err != nil { return nil, nil, errors.Trace(err) } - payload, ok := v["payload"].(map[string]interface{}) + payload, ok := v["payload"].(map[string]any) if !ok { return nil, nil, fmt.Errorf("decode payload failed, data: %+v", v) } - schema, ok := v["schema"].(map[string]interface{}) + schema, ok := v["schema"].(map[string]any) if !ok { return nil, nil, fmt.Errorf("decode payload failed, data: %+v", v) } diff --git a/pkg/sink/codec/open/decoder.go b/pkg/sink/codec/open/decoder.go index 1cf9cad13a..19cb2cebb1 100644 --- a/pkg/sink/codec/open/decoder.go +++ b/pkg/sink/codec/open/decoder.go @@ -192,16 +192,61 @@ func (b *decoder) NextDDLEvent() *commonEvent.DDLEvent { return result } -// NextDMLEvent implements the Decoder interface -func (b *decoder) NextDMLEvent() *commonEvent.DMLEvent { +// NextDMLMessage implements the Decoder interface +func (b *decoder) NextDMLMessage() *common.DMLMessage { if b.nextKey.Type != common.MessageTypeRow { log.Panic("message type is not row", zap.Any("messageType", b.nextKey.Type)) } + key := *b.nextKey + value := b.nextDMLValue() + b.nextKey = nil + + rowType := commonType.RowTypeInsert + if key.ClaimCheckLocation == "" { + rowType = b.rowTypeFromDMLValue(value) + } + tableID := tableIDAllocator.Allocate(key.Schema, key.Table) + return common.NewDMLMessage(tableID, key.Schema, key.Table, key.Ts, rowType, func() *commonEvent.DMLEvent { + return b.decodeDMLMessage(&key, value) + }) +} + +func (b *decoder) nextDMLValue() []byte { valueLen := binary.BigEndian.Uint64(b.valueBytes[:8]) value := b.valueBytes[8 : valueLen+8] b.valueBytes = b.valueBytes[valueLen+8:] + return append([]byte(nil), value...) +} +func (b *decoder) rowTypeFromDMLValue(value []byte) commonType.RowType { + value, err := common.Decompress(b.config.LargeMessageHandle.LargeMessageHandleCompression, value) + if err != nil { + log.Panic("decompress failed", + zap.String("compression", b.config.LargeMessageHandle.LargeMessageHandleCompression), + zap.Any("value", util.RedactAny(value)), zap.Error(err)) + } + + nextRow := new(messageRow) + nextRow.decode(value) + return rowTypeFromMessageRow(nextRow) +} + +func rowTypeFromMessageRow(row *messageRow) commonType.RowType { + if len(row.Delete) != 0 { + return commonType.RowTypeDelete + } + if len(row.Update) != 0 && len(row.PreColumns) != 0 { + return commonType.RowTypeUpdate + } + if len(row.Update) != 0 { + return commonType.RowTypeInsert + } + log.Panic("unknown event type") + return commonType.RowTypeInsert +} + +func (b *decoder) decodeDMLMessage(key *messageKey, value []byte) *commonEvent.DMLEvent { value, err := common.Decompress(b.config.LargeMessageHandle.LargeMessageHandleCompression, value) if err != nil { log.Panic("decompress failed", @@ -214,15 +259,15 @@ func (b *decoder) NextDMLEvent() *commonEvent.DMLEvent { ctx := context.Background() // claim-check message found - if b.nextKey.ClaimCheckLocation != "" { - return b.assembleEventFromClaimCheckStorage(ctx) + if key.ClaimCheckLocation != "" { + return b.assembleEventFromClaimCheckStorage(ctx, key) } - if b.nextKey.OnlyHandleKey && b.upstreamTiDB != nil { - return b.assembleHandleKeyOnlyDMLEvent(ctx, nextRow) + if key.OnlyHandleKey && b.upstreamTiDB != nil { + return b.assembleHandleKeyOnlyDMLEvent(ctx, key, nextRow) } - return b.assembleDMLEvent(nextRow) + return b.assembleDMLEvent(key, nextRow) } func buildColumns( @@ -259,8 +304,7 @@ func buildColumns( return columns } -func (b *decoder) assembleHandleKeyOnlyDMLEvent(ctx context.Context, row *messageRow) *commonEvent.DMLEvent { - key := b.nextKey +func (b *decoder) assembleHandleKeyOnlyDMLEvent(ctx context.Context, key *messageKey, row *messageRow) *commonEvent.DMLEvent { var ( schema = key.Schema table = key.Table @@ -290,13 +334,12 @@ func (b *decoder) assembleHandleKeyOnlyDMLEvent(ctx context.Context, row *messag } else { log.Panic("unknown event type") } - b.nextKey.OnlyHandleKey = false - return b.assembleDMLEvent(row) + key.OnlyHandleKey = false + return b.assembleDMLEvent(key, row) } -func (b *decoder) assembleEventFromClaimCheckStorage(ctx context.Context) *commonEvent.DMLEvent { - _, claimCheckFileName := filepath.Split(b.nextKey.ClaimCheckLocation) - b.nextKey = nil +func (b *decoder) assembleEventFromClaimCheckStorage(ctx context.Context, key *messageKey) *commonEvent.DMLEvent { + _, claimCheckFileName := filepath.Split(key.ClaimCheckLocation) data, err := b.storage.ReadFile(ctx, claimCheckFileName) if err != nil { log.Panic("read claim check file failed", zap.String("fileName", claimCheckFileName), zap.Error(err)) @@ -311,11 +354,11 @@ func (b *decoder) assembleEventFromClaimCheckStorage(ctx context.Context) *commo log.Panic("the batch version is not supported", zap.Uint64("version", version)) } - key := claimCheckM.Key[8:] - keyLen := binary.BigEndian.Uint64(key[:8]) - key = key[8 : keyLen+8] + encodedKey := claimCheckM.Key[8:] + keyLen := binary.BigEndian.Uint64(encodedKey[:8]) + encodedKey = encodedKey[8 : keyLen+8] msgKey := new(messageKey) - msgKey.Decode(key) + msgKey.Decode(encodedKey) valueLen := binary.BigEndian.Uint64(claimCheckM.Value[:8]) value := claimCheckM.Value[8 : valueLen+8] @@ -329,8 +372,7 @@ func (b *decoder) assembleEventFromClaimCheckStorage(ctx context.Context) *commo rowMsg := new(messageRow) rowMsg.decode(value) - b.nextKey = msgKey - return b.assembleDMLEvent(rowMsg) + return b.assembleDMLEvent(msgKey, rowMsg) } func (b *decoder) queryTableInfo(key *messageKey, value *messageRow) *commonType.TableInfo { @@ -490,10 +532,7 @@ func newTiIndices(columns []*timodel.ColumnInfo) []*timodel.IndexInfo { return indices } -func (b *decoder) assembleDMLEvent(value *messageRow) *commonEvent.DMLEvent { - key := b.nextKey - b.nextKey = nil - +func (b *decoder) assembleDMLEvent(key *messageKey, value *messageRow) *commonEvent.DMLEvent { tableInfo := b.queryTableInfo(key, value) result := new(commonEvent.DMLEvent) result.TableInfo = tableInfo diff --git a/pkg/sink/codec/open/encoder_test.go b/pkg/sink/codec/open/encoder_test.go index 463fa69b14..132844db09 100644 --- a/pkg/sink/codec/open/encoder_test.go +++ b/pkg/sink/codec/open/encoder_test.go @@ -83,7 +83,7 @@ func TestEncodeFlag(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() change, ok := decoded.GetNextRow() require.True(t, ok) @@ -171,7 +171,7 @@ func TestIntegerTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() require.Equal(t, event.CommitTs, decoded.GetCommitTs()) @@ -226,7 +226,7 @@ func TestFloatTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -275,7 +275,7 @@ func TestTimeTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -324,7 +324,7 @@ func TestStringTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -374,7 +374,7 @@ func TestBlobTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -424,7 +424,7 @@ func TestTextTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -471,7 +471,7 @@ func TestVectorType(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := dec.NextDMLEvent() + event := dec.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -520,7 +520,7 @@ func TestCollation(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -578,7 +578,7 @@ func TestOtherTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -701,7 +701,7 @@ func TestEncoderOneMessage(t *testing.T) { require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() change, ok := decoded.GetNextRow() require.True(t, ok) @@ -768,7 +768,7 @@ func TestEncoderMultipleMessage(t *testing.T) { require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() change, ok := decoded.GetNextRow() require.True(t, ok) @@ -778,7 +778,7 @@ func TestEncoderMultipleMessage(t *testing.T) { require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - decoded = decoder.NextDMLEvent() + decoded = decoder.NextDMLMessage().ToDMLEvent() change, ok = decoded.GetNextRow() require.True(t, ok) @@ -790,7 +790,7 @@ func TestEncoderMultipleMessage(t *testing.T) { require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - decoded = decoder.NextDMLEvent() + decoded = decoder.NextDMLMessage().ToDMLEvent() change, ok = decoded.GetNextRow() require.True(t, ok) @@ -874,7 +874,7 @@ func TestLargeMessageWithHandleEnableHandleKeyOnly(t *testing.T) { require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() change, ok := decoded.GetNextRow() require.True(t, ok) @@ -971,7 +971,7 @@ func TestDMLEventWithColumnSelector(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := decoder.NextDMLEvent() + event := decoder.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -1059,7 +1059,7 @@ func TestE2EPartitionTable(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, tp) - decodedEvent := dec.NextDMLEvent() + decodedEvent := dec.NextDMLMessage().ToDMLEvent() // table id should be set to the partition table id, the PhysicalTableID require.Equal(t, decodedEvent.GetTableID(), tableIDAllocator.Allocate(e.TableInfo.GetSchemaName(), e.TableInfo.GetTableName())) @@ -1172,7 +1172,7 @@ func TestGenerateColumn(t *testing.T) { require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - decoded := dec.NextDMLEvent() + decoded := dec.NextDMLMessage().ToDMLEvent() require.NoError(t, err) require.NotNil(t, decoded) require.Equal(t, decoded.Rows.NumCols(), 2) @@ -1202,7 +1202,7 @@ func TestGenerateColumn(t *testing.T) { require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - decoded = dec.NextDMLEvent() + decoded = dec.NextDMLMessage().ToDMLEvent() require.NoError(t, err) require.NotNil(t, decoded) require.Equal(t, decoded.Rows.NumCols(), 2) @@ -1232,8 +1232,6 @@ func TestGenerateColumn(t *testing.T) { messageType, hasNext = dec.HasNext() require.True(t, hasNext) require.Equal(t, messageType, common.MessageTypeRow) - - decoded = dec.NextDMLEvent() } // Including insert / update / delete @@ -1312,7 +1310,7 @@ func TestDMLEvent(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() change, ok := decoded.GetNextRow() require.True(t, ok) @@ -1364,7 +1362,7 @@ func TestOnlyOutputUpdatedEvent(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() change, ok := decoded.GetNextRow() require.True(t, ok) @@ -1409,7 +1407,7 @@ func TestPKWithUK(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := dec.NextDMLEvent() + event := dec.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) require.Len(t, event.TableInfo.GetIndices(), 2) @@ -1458,7 +1456,7 @@ func TestUniqueKeyWithoutPKDMLEvent(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - event := dec.NextDMLEvent() + event := dec.NextDMLMessage().ToDMLEvent() change, ok := event.GetNextRow() require.True(t, ok) @@ -1508,7 +1506,7 @@ func TestHandleOnlyEvent(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() change, ok := decoded.GetNextRow() require.True(t, ok) @@ -1573,7 +1571,7 @@ func TestRenameTable(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decodedInsert := decoder1.NextDMLEvent() + decodedInsert := decoder1.NextDMLMessage().ToDMLEvent() require.NotZero(t, decodedInsert.GetTableID()) require.Contains(t, tableIDAllocator.GetBlockedTables("test", "t"), decodedInsert.GetTableID()) @@ -1679,7 +1677,7 @@ func TestDDLSequence(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decodedInsert := decoder.NextDMLEvent() + decodedInsert := decoder.NextDMLMessage().ToDMLEvent() require.NotZero(t, decodedInsert.GetTableID()) require.Contains(t, tableIDAllocator.GetBlockedTables("test", "t"), decodedInsert.GetTableID()) diff --git a/pkg/sink/codec/simple/decoder.go b/pkg/sink/codec/simple/decoder.go index 101d73f472..76450349cb 100644 --- a/pkg/sink/codec/simple/decoder.go +++ b/pkg/sink/codec/simple/decoder.go @@ -57,8 +57,8 @@ type Decoder struct { // cachedMessages is used to store the messages which does not have received corresponding table info yet. cachedMessages *list.List - // CachedRowChangedEvents are events just decoded from the cachedMessages - CachedRowChangedEvents []*commonEvent.DMLEvent + // CachedDMLMessages are messages just released from the cachedMessages. + CachedDMLMessages []*common.DMLMessage } // NewDecoder returns a new Decoder @@ -147,34 +147,59 @@ func (d *Decoder) NextResolvedEvent() uint64 { return ts } -// NextDMLEvent returns the next dml event if exists -func (d *Decoder) NextDMLEvent() *commonEvent.DMLEvent { +// NextDMLMessage returns the next dml message if exists +func (d *Decoder) NextDMLMessage() *common.DMLMessage { if d.msg == nil || (d.msg.Data == nil && d.msg.Old == nil) { log.Panic("invalid data for the DML event", zap.String("message", util.RedactAny(d.msg))) } - if d.msg.ClaimCheckLocation != "" { - return d.assembleClaimCheckRowChangedEvent(d.msg.ClaimCheckLocation) + msg := d.msg + d.msg = nil + + tableInfo := d.memo.Read(msg.Schema, msg.Table, msg.SchemaVersion) + if tableInfo == nil { + log.Debug("table info not found for the message, "+ + "the consumer should cache this message temporarily, and update the tableInfo after it's received", + zap.String("schema", msg.Schema), + zap.String("table", msg.Table), + zap.Uint64("version", msg.SchemaVersion)) + d.cachedMessages.PushBack(msg) + return nil } - if d.msg.HandleKeyOnly { - return d.assembleHandleKeyOnlyRowChangedEvent(d.msg) + return d.newDMLMessage(msg, tableInfo) +} + +func (d *Decoder) newDMLMessage(msg *message, tableInfo *commonType.TableInfo) *common.DMLMessage { + return common.NewDMLMessage(msg.TableID, msg.Schema, msg.Table, msg.CommitTs, rowTypeFromMessageType(msg.Type), func() *commonEvent.DMLEvent { + return d.assembleDMLEvent(msg, tableInfo) + }) +} + +func rowTypeFromMessageType(tp MessageType) commonType.RowType { + switch tp { + case DMLTypeInsert: + return commonType.RowTypeInsert + case DMLTypeUpdate: + return commonType.RowTypeUpdate + case DMLTypeDelete: + return commonType.RowTypeDelete + default: + log.Panic("unknown row type for the DML message", zap.Any("type", tp)) } + return commonType.RowTypeInsert +} - tableInfo := d.memo.Read(d.msg.Schema, d.msg.Table, d.msg.SchemaVersion) - if tableInfo == nil { - log.Debug("table info not found for the event, "+ - "the consumer should cache this event temporarily, and update the tableInfo after it's received", - zap.String("schema", d.msg.Schema), - zap.String("table", d.msg.Table), - zap.Uint64("version", d.msg.SchemaVersion)) - d.cachedMessages.PushBack(d.msg) - d.msg = nil - return nil +func (d *Decoder) assembleDMLEvent(msg *message, tableInfo *commonType.TableInfo) *commonEvent.DMLEvent { + if msg.ClaimCheckLocation != "" { + return d.assembleClaimCheckRowChangedEvent(msg.ClaimCheckLocation, tableInfo) } - event := buildDMLEvent(d.msg, tableInfo, d.config.EnableRowChecksum, d.upstreamTiDB) - d.msg = nil + if msg.HandleKeyOnly { + return d.assembleHandleKeyOnlyRowChangedEvent(msg, tableInfo) + } + + event := buildDMLEvent(msg, tableInfo, d.config.EnableRowChecksum, d.upstreamTiDB) tableIDAllocator.AddBlockTableID(event.TableInfo.GetSchemaName(), event.TableInfo.GetTableName(), event.GetTableID()) @@ -182,7 +207,9 @@ func (d *Decoder) NextDMLEvent() *commonEvent.DMLEvent { return event } -func (d *Decoder) assembleClaimCheckRowChangedEvent(claimCheckLocation string) *commonEvent.DMLEvent { +func (d *Decoder) assembleClaimCheckRowChangedEvent( + claimCheckLocation string, tableInfo *commonType.TableInfo, +) *commonEvent.DMLEvent { _, claimCheckFileName := filepath.Split(claimCheckLocation) data, err := d.storage.ReadFile(context.Background(), claimCheckFileName) if err != nil { @@ -210,23 +237,12 @@ func (d *Decoder) assembleClaimCheckRowChangedEvent(claimCheckLocation string) * if err != nil { log.Panic("unmarshal claim check message failed", zap.Any("value", util.RedactAny(value)), zap.Error(err)) } - d.msg = m - return d.NextDMLEvent() + return d.assembleDMLEvent(m, tableInfo) } -func (d *Decoder) assembleHandleKeyOnlyRowChangedEvent(m *message) *commonEvent.DMLEvent { - tableInfo := d.memo.Read(m.Schema, m.Table, m.SchemaVersion) - if tableInfo == nil { - log.Debug("table info not found for the event, "+ - "the consumer should cache this event temporarily, and update the tableInfo after it's received", - zap.String("schema", d.msg.Schema), - zap.String("table", d.msg.Table), - zap.Uint64("version", d.msg.SchemaVersion)) - d.cachedMessages.PushBack(d.msg) - d.msg = nil - return nil - } - +func (d *Decoder) assembleHandleKeyOnlyRowChangedEvent( + m *message, tableInfo *commonType.TableInfo, +) *commonEvent.DMLEvent { fieldTypeMap := make(map[string]*types.FieldType, len(tableInfo.GetColumns())) for _, col := range tableInfo.GetColumns() { fieldTypeMap[col.Name.O] = &col.FieldType @@ -259,8 +275,7 @@ func (d *Decoder) assembleHandleKeyOnlyRowChangedEvent(m *message) *commonEvent. result.Old = d.buildData(holder, fieldTypeMap, timezone) } - d.msg = result - return d.NextDMLEvent() + return d.assembleDMLEvent(result, tableInfo) } func (d *Decoder) buildData( @@ -291,21 +306,22 @@ func (d *Decoder) NextDDLEvent() *commonEvent.DDLEvent { d.memo.Write(ddl.MultipleTableInfos[1]) for ele := d.cachedMessages.Front(); ele != nil; { - d.msg = ele.Value.(*message) - event := d.NextDMLEvent() - d.CachedRowChangedEvents = append(d.CachedRowChangedEvents, event) - + msg := ele.Value.(*message) next := ele.Next() - d.cachedMessages.Remove(ele) + tableInfo := d.memo.Read(msg.Schema, msg.Table, msg.SchemaVersion) + if tableInfo != nil { + d.CachedDMLMessages = append(d.CachedDMLMessages, d.newDMLMessage(msg, tableInfo)) + d.cachedMessages.Remove(ele) + } ele = next } return ddl } -// GetCachedEvents returns the cached events -func (d *Decoder) GetCachedEvents() []*commonEvent.DMLEvent { - result := d.CachedRowChangedEvents - d.CachedRowChangedEvents = nil +// GetCachedMessages returns the cached messages. +func (d *Decoder) GetCachedMessages() []*common.DMLMessage { + result := d.CachedDMLMessages + d.CachedDMLMessages = nil return result } @@ -639,14 +655,15 @@ func buildDMLEvent(msg *message, tableInfo *commonType.TableInfo, enableRowCheck } func formatAllColumnsValue(data map[string]any, columns []*timodel.ColumnInfo) map[string]any { + result := make(map[string]any, len(data)) for _, col := range columns { raw, ok := data[col.Name.O] if !ok { continue } - data[col.Name.O] = formatValue(raw, col.FieldType) + result[col.Name.O] = formatValue(raw, col.FieldType) } - return data + return result } // formatValue formats the value according to the field type diff --git a/pkg/sink/codec/simple/decoder_test.go b/pkg/sink/codec/simple/decoder_test.go new file mode 100644 index 0000000000..0272380941 --- /dev/null +++ b/pkg/sink/codec/simple/decoder_test.go @@ -0,0 +1,104 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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, +// See the License for the specific language governing permissions and +// limitations under the License. + +package simple + +import ( + "container/list" + "testing" + + commonType "github.com/pingcap/ticdc/pkg/common" + "github.com/pingcap/ticdc/pkg/config" + "github.com/pingcap/ticdc/pkg/sink/codec/common" + "github.com/stretchr/testify/require" +) + +func TestCachedDMLReturnsMessage(t *testing.T) { + const ( + schema = "test" + table = "t" + tableID = int64(1) + schemaVersion = uint64(100) + commitTs = uint64(90) + ) + + decoder := &Decoder{ + config: common.NewConfig(config.ProtocolSimple), + memo: newMemoryTableInfoProvider(), + cachedMessages: list.New(), + } + decoder.msg = &message{ + Version: defaultVersion, + Schema: schema, + Table: table, + TableID: tableID, + Type: DMLTypeInsert, + CommitTs: commitTs, + SchemaVersion: schemaVersion, + Data: map[string]any{"id": int64(1)}, + } + + require.Nil(t, decoder.NextDMLMessage()) + require.Equal(t, 1, decoder.cachedMessages.Len()) + + decoder.msg = &message{ + Version: defaultVersion, + Type: DDLTypeCreate, + CommitTs: schemaVersion, + TableSchema: &TableSchema{ + Schema: schema, + Table: table, + TableID: tableID, + Version: schemaVersion, + Columns: []*columnSchema{ + { + Name: "id", + DataType: dataType{ + MySQLType: "bigint", + Charset: "binary", + Collate: "binary", + Length: 20, + }, + }, + }, + }, + } + ddl := decoder.NextDDLEvent() + require.NotNil(t, ddl) + + cachedMessages := decoder.GetCachedMessages() + require.Len(t, cachedMessages, 1) + require.Zero(t, decoder.cachedMessages.Len()) + + dmlMessage := cachedMessages[0] + require.Equal(t, tableID, dmlMessage.TableID) + require.Equal(t, schema, dmlMessage.Schema) + require.Equal(t, table, dmlMessage.Table) + require.Equal(t, commitTs, dmlMessage.GetCommitTs()) + require.Equal(t, commonType.RowTypeInsert, dmlMessage.RowType) + + decoder.msg = &message{ + Version: defaultVersion, + Type: MessageTypeWatermark, + CommitTs: commitTs + 1, + } + dmlEvent := dmlMessage.ToDMLEvent() + require.NotNil(t, dmlEvent) + require.Equal(t, tableID, dmlEvent.GetTableID()) + require.Equal(t, commitTs, dmlEvent.GetCommitTs()) + require.Equal(t, schema, dmlEvent.TableInfo.GetSchemaName()) + require.Equal(t, table, dmlEvent.TableInfo.GetTableName()) + require.NotNil(t, decoder.msg) + require.Equal(t, MessageTypeWatermark, decoder.msg.Type) + require.Equal(t, commitTs+1, decoder.msg.CommitTs) +} diff --git a/pkg/sink/codec/simple/encoder_test.go b/pkg/sink/codec/simple/encoder_test.go index 2e92394f15..b4f8e5ed6a 100644 --- a/pkg/sink/codec/simple/encoder_test.go +++ b/pkg/sink/codec/simple/encoder_test.go @@ -146,7 +146,7 @@ func TestEncodeDMLEnableChecksum(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - // decodedRow:= decoder.NextDMLEvent() + // decodedRow:= decoder.NextDMLMessage().ToDMLEvent() // require.NoError(t, err) // require.Equal(t, updateEvent.Checksum.Current, decodedRow.Checksum.Current) // require.Equal(t, updateEvent.Checksum.Previous, decodedRow.Checksum.Previous) @@ -189,7 +189,7 @@ func TestEncodeDMLEnableChecksum(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - // decodedRow:= decoder.NextDMLEvent() + // decodedRow:= decoder.NextDMLMessage().ToDMLEvent() // require.Error(t, err) // require.Nil(t, decodedRow) } @@ -268,7 +268,7 @@ func TestE2EPartitionTable(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, tp) - decodedEvent := dec.NextDMLEvent() + decodedEvent := dec.NextDMLMessage().ToDMLEvent() // table id should be set to the partition table id, the PhysicalTableID require.Equal(t, decodedEvent.GetTableID(), e.GetTableID()) @@ -877,7 +877,7 @@ func TestEncodeDDLEvent(t *testing.T) { require.Equal(t, common.MessageTypeRow, messageType) require.NotEqual(t, 0, dec.msg.BuildTs) - decodedRow := dec.NextDMLEvent() + decodedRow := dec.NextDMLMessage().ToDMLEvent() require.Equal(t, decodedRow.CommitTs, insertEvent.GetCommitTs()) require.Equal(t, decodedRow.TableInfo.GetSchemaName(), insertEvent.TableInfo.GetSchemaName()) require.Equal(t, decodedRow.TableInfo.GetTableName(), insertEvent.TableInfo.GetTableName()) @@ -926,7 +926,7 @@ func TestEncodeDDLEvent(t *testing.T) { require.Equal(t, common.MessageTypeRow, messageType) require.NotEqual(t, 0, dec.msg.BuildTs) - decodedRow = dec.NextDMLEvent() + decodedRow = dec.NextDMLMessage().ToDMLEvent() require.Equal(t, insertEvent2.GetCommitTs(), decodedRow.GetCommitTs()) require.Equal(t, insertEvent2.TableInfo.GetSchemaName(), decodedRow.TableInfo.GetSchemaName()) require.Equal(t, insertEvent2.TableInfo.GetTableName(), decodedRow.TableInfo.GetTableName()) @@ -1080,7 +1080,7 @@ func TestEncodeIntegerTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decodedRow := dec.NextDMLEvent() + decodedRow := dec.NextDMLMessage().ToDMLEvent() require.Equal(t, decodedRow.CommitTs, event.GetCommitTs()) decoded, ok := decodedRow.GetNextRow() @@ -1155,7 +1155,7 @@ func TestEncoderOtherTypes(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decodedRow := dec.NextDMLEvent() + decodedRow := dec.NextDMLMessage().ToDMLEvent() decoded, ok := decodedRow.GetNextRow() require.True(t, ok) @@ -1223,8 +1223,8 @@ func TestE2EPartitionTableDMLBeforeDDL(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, tp) - decodedEvent := dec.NextDMLEvent() - require.Nil(t, decodedEvent) + decodedMessage := dec.NextDMLMessage() + require.Nil(t, decodedMessage) e.Rewind() } @@ -1240,8 +1240,9 @@ func TestE2EPartitionTableDMLBeforeDDL(t *testing.T) { decodedDDL := dec.NextDDLEvent() require.NotNil(t, decodedDDL) - cachedEvents := dec.(*Decoder).GetCachedEvents() - for idx, decodedRow := range cachedEvents { + cachedMessages := dec.(*Decoder).GetCachedMessages() + for idx, message := range cachedMessages { + decodedRow := message.ToDMLEvent() require.NotNil(t, decodedRow) require.NotNil(t, decodedRow.TableInfo) require.Equal(t, decodedRow.GetTableID(), events[idx].GetTableID()) @@ -1289,8 +1290,8 @@ func TestEncodeDMLBeforeDDL(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decodedRow := dec.NextDMLEvent() - require.Nil(t, decodedRow) + decodedMessage := dec.NextDMLMessage() + require.Nil(t, decodedMessage) m, err := enc.EncodeDDLEvent(ddlEvent) require.NoError(t, err) @@ -1304,8 +1305,9 @@ func TestEncodeDMLBeforeDDL(t *testing.T) { ddlEvent = dec.NextDDLEvent() require.NotNil(t, ddlEvent) - cachedEvents := dec.GetCachedEvents() - for _, decodedRow = range cachedEvents { + cachedMessages := dec.GetCachedMessages() + for _, message := range cachedMessages { + decodedRow := message.ToDMLEvent() require.NotNil(t, decodedRow) require.NotNil(t, decodedRow.TableInfo) require.Equal(t, decodedRow.TableInfo.TableName.TableID, ddlEvent.TableInfo.TableName.TableID) @@ -1394,7 +1396,7 @@ func TestEncodeBootstrapEvent(t *testing.T) { require.Equal(t, common.MessageTypeRow, messageType) require.NotEqual(t, 0, dec.msg.BuildTs) - decodedRow := dec.NextDMLEvent() + decodedRow := dec.NextDMLMessage().ToDMLEvent() decode, ok := decodedRow.GetNextRow() require.True(t, ok) require.Equal(t, decodedRow.CommitTs, dmlEvent.CommitTs) @@ -1478,7 +1480,7 @@ func TestEncodeLargeEventsNormal(t *testing.T) { require.Equal(t, dec.msg.Type, DMLTypeInsert) } - decodedRow := dec.NextDMLEvent() + decodedRow := dec.NextDMLMessage().ToDMLEvent() require.Equal(t, decodedRow.CommitTs, event.CommitTs) require.Equal(t, decodedRow.TableInfo.GetSchemaName(), event.TableInfo.GetSchemaName()) @@ -1603,7 +1605,7 @@ func TestLargerMessageHandleClaimCheck(t *testing.T) { require.Equal(t, common.MessageTypeRow, messageType) require.NotEqual(t, "", dec.msg.ClaimCheckLocation) - decodedRow := dec.NextDMLEvent() + decodedRow := dec.NextDMLMessage().ToDMLEvent() require.Equal(t, decodedRow.CommitTs, updateEvent.CommitTs) require.Equal(t, decodedRow.TableInfo.GetSchemaName(), updateEvent.TableInfo.GetSchemaName()) @@ -1675,8 +1677,8 @@ func TestLargeMessageHandleKeyOnly(t *testing.T) { require.Equal(t, common.MessageTypeRow, messageType) require.True(t, dec.msg.HandleKeyOnly) - decodedRow := dec.NextDMLEvent() - require.Nil(t, decodedRow) + decodedMessage := dec.NextDMLMessage() + require.Nil(t, decodedMessage) } enc.(*Encoder).config.MaxMessageBytes = config.DefaultMaxMessageBytes @@ -1710,8 +1712,9 @@ func TestLargeMessageHandleKeyOnly(t *testing.T) { } _ = dec.NextDDLEvent() - decodedRows := dec.GetCachedEvents() - for idx, decodedRow := range decodedRows { + decodedMessages := dec.GetCachedMessages() + for idx, message := range decodedMessages { + decodedRow := message.ToDMLEvent() event := events[idx] require.Equal(t, decodedRow.CommitTs, event.CommitTs)