From 74261574c5b48e65444d92329474068e2078f56d Mon Sep 17 00:00:00 2001 From: nhsmw Date: Tue, 21 Jul 2026 12:48:08 +0800 Subject: [PATCH 1/3] This is an automated cherry-pick of #5590 Signed-off-by: ti-chi-bot --- cmd/kafka-consumer/writer.go | 132 +++- cmd/kafka-consumer/writer_test.go | 59 ++ cmd/pulsar-consumer/consumer.go | 4 + cmd/pulsar-consumer/writer.go | 104 ++- cmd/pulsar-consumer/writer_test.go | 195 ++++- cmd/storage-consumer/consumer.go | 48 +- cmd/util/event_group.go | 100 ++- cmd/util/event_group_test.go | 116 ++- pkg/sink/codec/avro/avro_test.go | 2 +- pkg/sink/codec/avro/decoder.go | 105 ++- pkg/sink/codec/avro/encoder_test.go | 4 +- pkg/sink/codec/canal/canal_json_decoder.go | 111 ++- .../codec/canal/canal_json_encoder_test.go | 24 +- pkg/sink/codec/canal/canal_json_test.go | 173 ++++- .../codec/canal/canal_json_txn_decoder.go | 35 +- pkg/sink/codec/common/decoder.go | 61 +- pkg/sink/codec/common/table_info_cache.go | 14 + pkg/sink/codec/csv/csv_decoder.go | 29 +- pkg/sink/codec/csv/csv_decoder_test.go | 2 +- pkg/sink/codec/debezium/avro_decoder.go | 731 ++++++++++++++++++ pkg/sink/codec/debezium/avro_test.go | 682 ++++++++++++++++ pkg/sink/codec/debezium/debezium_test.go | 2 +- pkg/sink/codec/debezium/decoder.go | 129 +++- pkg/sink/codec/open/decoder.go | 87 ++- pkg/sink/codec/open/encoder_test.go | 56 +- pkg/sink/codec/simple/decoder.go | 115 +-- pkg/sink/codec/simple/decoder_test.go | 104 +++ pkg/sink/codec/simple/encoder_test.go | 49 +- 28 files changed, 2921 insertions(+), 352 deletions(-) create mode 100644 pkg/sink/codec/debezium/avro_decoder.go create mode 100644 pkg/sink/codec/debezium/avro_test.go create mode 100644 pkg/sink/codec/simple/decoder_test.go diff --git a/cmd/kafka-consumer/writer.go b/cmd/kafka-consumer/writer.go index ed40a534bc..9c300fb6f9 100644 --- a/cmd/kafka-consumer/writer.go +++ b/cmd/kafka-consumer/writer.go @@ -152,7 +152,6 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *event.DDLEvent) error { var ( done = make(chan struct{}, 1) - total int flushed atomic.Int64 ) @@ -171,6 +170,7 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *event.DDLEvent) error { if !ok { continue } +<<<<<<< HEAD before := len(resolvedEvents) resolvedEvents = g.ResolveInto(commitTs, resolvedEvents) resolvedCount := len(resolvedEvents) - before @@ -186,9 +186,18 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *event.DDLEvent) error { maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), }) total += resolvedCount +======= + messages := g.ResolveInto(commitTs, nil) + events := make([]*event.DMLEvent, 0, len(messages)) + for _, message := range messages { + events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) + } + resolvedEvents = append(resolvedEvents, events...) +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } } + total := len(resolvedEvents) if total == 0 { return w.mysqlSink.WriteBlockEvent(ddl) } @@ -289,7 +298,6 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { var ( done = make(chan struct{}, 1) - total int flushed atomic.Int64 ) @@ -303,6 +311,7 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { }, 0) for _, p := range w.progresses { for _, group := range p.eventsGroup { +<<<<<<< HEAD before := len(resolvedEvents) resolvedEvents = group.ResolveInto(watermark, resolvedEvents) resolvedCount := len(resolvedEvents) - before @@ -318,8 +327,17 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), }) total += resolvedCount +======= + messages := group.ResolveInto(watermark, nil) + events := make([]*event.DMLEvent, 0, len(messages)) + for _, message := range messages { + events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) + } + resolvedEvents = append(resolvedEvents, events...) +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } } + total := len(resolvedEvents) if total == 0 { return nil } @@ -391,12 +409,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) } } @@ -420,25 +438,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. @@ -590,15 +616,22 @@ func (w *writer) checkPartition(row *event.DMLEvent, partition int32, offset kaf } } -func (w *writer) appendRow2Group(dml *event.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 { @@ -618,14 +651,22 @@ func (w *writer) appendRow2Group(dml *event.DMLEvent, progress *partitionProgres zap.String("schema", schema), zap.String("table", table), zap.Any("protocol", w.protocol)) return } +<<<<<<< HEAD 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", +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) 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.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), +<<<<<<< HEAD zap.Stringer("eventType", dml.RowTypes[0]), zap.Any("protocol", w.protocol), zap.Bool("IsPartition", dml.TableInfo.TableName.IsPartition)) group.Append(dml, true) @@ -638,6 +679,59 @@ func (w *writer) appendRow2Group(dml *event.DMLEvent, progress *partitionProgres zap.Uint64("appliedWatermark", group.AppliedWatermark), 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 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)) + } +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } 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 70fed688de..c40871e0b5 100644 --- a/cmd/kafka-consumer/writer_test.go +++ b/cmd/kafka-consumer/writer_test.go @@ -317,6 +317,11 @@ func TestAppendRow2Group_DoesNotDropCommitTsFallbackBeforeApplied(t *testing.T) } progress := w.progresses[0] +<<<<<<< HEAD +======= + w.appendMessage2Group(codeccommon.NewDMLMessageFromEvent(newDMLEvent(200)), progress, kafka.Offset(10)) + w.appendMessage2Group(codeccommon.NewDMLMessageFromEvent(newDMLEvent(100)), progress, kafka.Offset(11)) +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) // Step 1: observe a larger commitTs first (e.g. produced before restart). w.appendRow2Group(newDMLEvent(1, 200), progress, kafka.Offset(10)) @@ -331,6 +336,7 @@ func TestAppendRow2Group_DoesNotDropCommitTsFallbackBeforeApplied(t *testing.T) // Expect: commitTs=100 is still kept and can be resolved. resolved := group.ResolveInto(150, nil) require.Len(t, resolved, 1) +<<<<<<< HEAD require.Equal(t, uint64(100), resolved[0].CommitTs) // Step 3: once downstream has flushed beyond commitTs=100, the replay is safe to ignore. @@ -339,4 +345,57 @@ func TestAppendRow2Group_DoesNotDropCommitTsFallbackBeforeApplied(t *testing.T) w.appendRow2Group(newDMLEvent(1, 100), progress, kafka.Offset(12)) resolved = group.ResolveInto(150, resolvedEvents) require.Empty(t, resolved) +======= + require.Equal(t, uint64(100), resolved[0].GetCommitTs()) +} + +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()) + }) + } +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } diff --git a/cmd/pulsar-consumer/consumer.go b/cmd/pulsar-consumer/consumer.go index d23ed61675..62f2feba91 100644 --- a/cmd/pulsar-consumer/consumer.go +++ b/cmd/pulsar-consumer/consumer.go @@ -115,7 +115,11 @@ func (c *consumer) readMessage(ctx context.Context) error { if !needCommit { continue } +<<<<<<< HEAD err := c.pulsarConsumer.AckID(consumerMsg.Message.ID()) +======= + err := c.pulsarConsumer.AckIDCumulative(consumerMsg.ID()) +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) 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 fbbf94c3aa..43d3429309 100644 --- a/cmd/pulsar-consumer/writer.go +++ b/cmd/pulsar-consumer/writer.go @@ -143,7 +143,6 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *commonEvent.DDLEvent) e var ( done = make(chan struct{}, 1) - total int flushed atomic.Int64 ) @@ -162,6 +161,7 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *commonEvent.DDLEvent) e if !ok { continue } +<<<<<<< HEAD before := len(resolvedEvents) resolvedEvents = g.ResolveInto(commitTs, resolvedEvents) resolvedCount := len(resolvedEvents) - before @@ -177,9 +177,18 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *commonEvent.DDLEvent) e maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), }) total += resolvedCount +======= + messages := g.ResolveInto(commitTs, nil) + events := make([]*commonEvent.DMLEvent, 0, len(messages)) + for _, message := range messages { + events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) + } + resolvedEvents = append(resolvedEvents, events...) +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } } + total := len(resolvedEvents) if total == 0 { return w.mysqlSink.WriteBlockEvent(ddl) } @@ -280,7 +289,6 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { var ( done = make(chan struct{}, 1) - total int flushed atomic.Int64 ) @@ -294,6 +302,7 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { }, 0) for _, p := range w.progresses { for _, group := range p.eventsGroup { +<<<<<<< HEAD before := len(resolvedEvents) resolvedEvents = group.ResolveInto(watermark, resolvedEvents) resolvedCount := len(resolvedEvents) - before @@ -309,8 +318,17 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), }) total += resolvedCount +======= + messages := group.ResolveInto(watermark, nil) + events := make([]*commonEvent.DMLEvent, 0, len(messages)) + for _, message := range messages { + events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) + } + resolvedEvents = append(resolvedEvents, events...) +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } } + total := len(resolvedEvents) if total == 0 { return nil } @@ -387,12 +405,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)) } @@ -498,12 +515,34 @@ func (w *writer) onDDL(ddl *commonEvent.DDLEvent) { } } +<<<<<<< HEAD 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.GetTargetSchemaName(), ddl.TableInfo.GetTargetTableName()) + 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) { +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) 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 { @@ -519,14 +558,21 @@ func (w *writer) appendRow2Group(dml *commonEvent.DMLEvent, progress *partitionP zap.String("schema", schema), zap.String("table", table), zap.Any("protocol", w.protocol)) return } +<<<<<<< HEAD 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 commitTs >= group.HighWatermark { + group.AppendMessage(message, false) + log.Debug("DML event append to the group", +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) 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), +<<<<<<< HEAD zap.Stringer("eventType", dml.RowTypes[0]), zap.Any("protocol", w.protocol), zap.Bool("IsPartition", dml.TableInfo.TableName.IsPartition)) group.Append(dml, true) @@ -539,4 +585,40 @@ func (w *writer) appendRow2Group(dml *commonEvent.DMLEvent, progress *partitionP zap.Uint64("appliedWatermark", group.AppliedWatermark), 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 w.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", message.RowType)) + group.AppendMessage(message, true) + return + } + 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)) + } +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } diff --git a/cmd/pulsar-consumer/writer_test.go b/cmd/pulsar-consumer/writer_test.go index 8fa315d686..0ee23c2b22 100644 --- a/cmd/pulsar-consumer/writer_test.go +++ b/cmd/pulsar-consumer/writer_test.go @@ -16,7 +16,13 @@ package main import ( "context" "testing" + "time" +<<<<<<< HEAD +======= + "github.com/apache/pulsar-client-go/pulsar" + "github.com/golang/mock/gomock" +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) "github.com/pingcap/ticdc/cmd/util" "github.com/pingcap/ticdc/downstreamadapter/sink" "github.com/pingcap/ticdc/pkg/common" @@ -24,7 +30,6 @@ import ( "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" ) @@ -302,6 +307,7 @@ func TestAppendRow2Group_DoesNotDropCommitTsFallbackBeforeApplied(t *testing.T) partitionTableAccessor: codeccommon.NewPartitionTableAccessor(), } +<<<<<<< HEAD newDMLEvent := func(tableID int64, commitTs uint64) *commonEvent.DMLEvent { return &commonEvent.DMLEvent{ PhysicalTableID: tableID, @@ -315,6 +321,33 @@ func TestAppendRow2Group_DoesNotDropCommitTsFallbackBeforeApplied(t *testing.T) } progress := w.progresses[0] +======= + ddl := &commonEvent.DDLEvent{ + Query: "CREATE TABLE `target`.`dst` LIKE `target`.`src`", + SchemaName: "source", + TableName: "dst", + Type: byte(timodel.ActionCreateTable), + TableInfo: &common.TableInfo{ + TableName: common.TableName{ + Schema: "source", + Table: "dst", + IsPartition: true, + TargetSchema: "target", + TargetTable: "dst", + }, + }, + } + w.onDDL(ddl) + require.True(t, w.partitionTableAccessor.IsPartitionTable("target", "dst")) + + newDMLMessage := func(commitTs uint64) *codeccommon.DMLMessage { + return codeccommon.NewDMLMessage(1, "target", "dst", commitTs, common.RowTypeUpdate, nil) + } + + progress := w.progresses[0] + w.appendMessage2Group(newDMLMessage(200), progress) + w.appendMessage2Group(newDMLMessage(100), progress) +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) // Step 1: observe a larger commitTs first (e.g. produced before restart). w.appendRow2Group(newDMLEvent(1, 200), progress) @@ -328,6 +361,7 @@ func TestAppendRow2Group_DoesNotDropCommitTsFallbackBeforeApplied(t *testing.T) // Expect: commitTs=100 is still kept and can be resolved. resolved := group.ResolveInto(150, nil) require.Len(t, resolved, 1) +<<<<<<< HEAD require.Equal(t, uint64(100), resolved[0].CommitTs) // Step 3: once downstream has flushed beyond commitTs=100, replay is safe to ignore. @@ -335,4 +369,163 @@ func TestAppendRow2Group_DoesNotDropCommitTsFallbackBeforeApplied(t *testing.T) w.appendRow2Group(newDMLEvent(1, 100), progress) resolved = group.ResolveInto(150, nil) require.Empty(t, resolved) +======= + require.Equal(t, uint64(100), resolved[0].GetCommitTs()) +} + +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) + + decoder := &deferredDMLDecoder{ + row: &commonEvent.DMLEvent{ + PhysicalTableID: 1, + CommitTs: 100, + RowTypes: []common.RowType{common.RowTypeInsert}, + TableInfo: &common.TableInfo{ + 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.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) +} + +type deferredDMLDecoder struct { + row *commonEvent.DMLEvent + + addKeyValueCount int + hasNextCount int + nextDMLMessageCount int + toDMLEventCount int + lastValue []byte +} + +func (d *deferredDMLDecoder) AddKeyValue(_, value []byte) { + d.addKeyValueCount++ + d.lastValue = append(d.lastValue[:0], value...) +} + +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 +} + +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 +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } diff --git a/cmd/storage-consumer/consumer.go b/cmd/storage-consumer/consumer.go index e3ee2d6eb5..12cc90b335 100644 --- a/cmd/storage-consumer/consumer.go +++ b/cmd/storage-consumer/consumer.go @@ -239,12 +239,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 { @@ -252,25 +252,31 @@ func (c *consumer) appendRow2Group(dml *event.DMLEvent, enableTableAcrossNodes b c.eventsGroup[tableID] = group } if commitTs >= group.HighWatermark { +<<<<<<< HEAD group.Append(dml, false) log.Info("DML event append to the group", +======= + group.AppendMessage(message, false) + log.Debug("DML event append to the group", +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) 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), ) } @@ -325,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++ } } @@ -340,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..d3bf980b5b 100644 --- a/cmd/util/event_group.go +++ b/cmd/util/event_group.go @@ -14,26 +14,30 @@ 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 +<<<<<<< HEAD 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 +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) HighWatermark uint64 // AppliedWatermark is the maximum CommitTs that has been successfully flushed // to the downstream for this group. @@ -49,19 +53,78 @@ 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 +// 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 } + var lastMessage *codeccommon.DMLMessage + if len(g.messages) > 0 { + lastMessage = g.messages[len(g.messages)-1] + } + + if lastMessage == nil || lastMessage.GetCommitTs() <= commitTs { + g.messages = append(g.messages, message) + return + } + + if force { + i := sort.Search(len(g.messages), func(i int) bool { + return g.messages[i].GetCommitTs() > commitTs + }) + 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", lastMessage.GetCommitTs()), zap.Uint64("commitTs", commitTs)) +} + +// 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.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.messages)), + zap.Uint64("resolveTs", resolve), zap.Uint64("firstCommitTs", g.messages[0].GetCommitTs())) + } + return dst +} + +// 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(g.events) > 0 { - lastDMLEvent = g.events[len(g.events)-1] + if len(events) > 0 { + lastDMLEvent = events[len(events)-1] } mergeDMLEvent := func(dst, src *commonEvent.DMLEvent) { @@ -72,11 +135,11 @@ func (g *EventsGroup) Append(row *commonEvent.DMLEvent, force bool) { } if lastDMLEvent == nil || lastDMLEvent.GetCommitTs() < row.GetCommitTs() { - g.events = append(g.events, row) - return + return append(events, row) } if lastDMLEvent.GetCommitTs() == row.GetCommitTs() { +<<<<<<< HEAD mergeDMLEvent(lastDMLEvent, row) return } @@ -108,9 +171,19 @@ func (g *EventsGroup) Append(row *commonEvent.DMLEvent, force bool) { g.events = slices.Insert(g.events, i, row) return } +======= + 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 + } + +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) log.Panic("append event with smaller commit ts", - zap.Int32("partition", g.Partition), zap.Int64("tableID", g.tableID), + zap.Int64("tableID", row.GetTableID()), zap.Uint64("lastCommitTs", lastDMLEvent.GetCommitTs()), zap.Uint64("commitTs", row.GetCommitTs())) +<<<<<<< HEAD } func compareTableInfo(previous, now *commonEvent.DMLEvent) bool { @@ -153,4 +226,7 @@ func (g *EventsGroup) GetAllEvents() []*commonEvent.DMLEvent { result := g.events g.events = nil return result +======= + return events +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } diff --git a/cmd/util/event_group_test.go b/cmd/util/event_group_test.go index 5b8816ea50..2641236bd9 100644 --- a/cmd/util/event_group_test.go +++ b/cmd/util/event_group_test.go @@ -18,12 +18,17 @@ import ( "github.com/pingcap/ticdc/pkg/common" commonEvent "github.com/pingcap/ticdc/pkg/common/event" +<<<<<<< HEAD timodel "github.com/pingcap/tidb/pkg/meta/model" parser_model "github.com/pingcap/tidb/pkg/parser/model" +======= + codeccommon "github.com/pingcap/ticdc/pkg/sink/codec/common" +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) "github.com/pingcap/tidb/pkg/util/chunk" "github.com/stretchr/testify/require" ) +<<<<<<< HEAD func TestEventsGroupAppendForceMergesExistingCommitTs(t *testing.T) { // Scenario: // 1) An upstream transaction (commitTs=100) is split into multiple messages. @@ -62,6 +67,20 @@ func TestEventsGroupAppendForceMergesExistingCommitTs(t *testing.T) { require.Len(t, dst, 1) require.Equal(t, uint64(100), dst[0].CommitTs) require.Len(t, dst[0].RowTypes, 2) +======= +func newTestDMLMessage(commitTs uint64) *codeccommon.DMLMessage { + return codeccommon.NewDMLMessage(1, "test", "t", commitTs, common.RowTypeInsert, nil) +} + +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), + } +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } func TestEventsGroupResolveIntoAppendsAndClearsResolvedPrefix(t *testing.T) { @@ -75,75 +94,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/avro_test.go b/pkg/sink/codec/avro/avro_test.go index 798bc611d8..c4e5883e7c 100644 --- a/pkg/sink/codec/avro/avro_test.go +++ b/pkg/sink/codec/avro/avro_test.go @@ -187,7 +187,7 @@ func TestAvroEncodeDeleteEventWithWatermarkCarriesCommitTs(t *testing.T) { require.True(t, exists) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() require.NotNil(t, decoded) require.Equal(t, event.CommitTs, decoded.GetCommitTs()) require.Len(t, decoded.RowTypes, 1) diff --git a/pkg/sink/codec/avro/decoder.go b/pkg/sink/codec/avro/decoder.go index d590dee0e5..180fc0d8a1 100644 --- a/pkg/sink/codec/avro/decoder.go +++ b/pkg/sink/codec/avro/decoder.go @@ -113,24 +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 || d.isDeleteValue() - deleteCommitTs := uint64(0) + isDelete = len(d.value) == 0 || d.isDeleteValue() if isDelete { // delete event only have key part, treat it as the value part also. if d.isDeleteValue() { @@ -145,6 +166,16 @@ func (d *decoder) NextDMLEvent() *commonEvent.DMLEvent { } } + 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)) @@ -213,18 +244,18 @@ func (d *decoder) decodeDeleteCommitTs() uint64 { // 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") } @@ -235,18 +266,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"])) } @@ -273,10 +304,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 { @@ -307,12 +335,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) @@ -335,7 +368,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 @@ -347,7 +380,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 @@ -366,12 +399,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 } @@ -536,7 +569,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 @@ -570,12 +603,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) } @@ -583,13 +616,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 83797b551e..30a998a04a 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()) @@ -131,7 +131,7 @@ func TestEncodeRoutedDMLEventUsesTargetNames(t *testing.T) { require.True(t, exists) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() require.Equal(t, "target_db", decoded.TableInfo.GetSchemaName()) require.Equal(t, "target_table", decoded.TableInfo.GetTableName()) } diff --git a/pkg/sink/codec/canal/canal_json_decoder.go b/pkg/sink/codec/canal/canal_json_decoder.go index a75b34fd21..52179bc0ad 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 storage.ExternalStorage 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,8 +414,13 @@ func (d *decoder) NextDDLEvent() *commonEvent.DDLEvent { tableIDAllocator.AddBlockTableID(result.SchemaName, result.TableName, tableIDAllocator.Allocate(result.SchemaName, result.TableName)) result.BlockedTables = common.GetBlockedTables(tableIDAllocator, result) +<<<<<<< HEAD // 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}) +======= + d.addDDLCommitTs(result.SchemaName, result.TableName, result.GetCommitTs()) + d.addDDLCommitTs(result.ExtraSchemaName, result.ExtraTableName, result.GetCommitTs()) +>>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) return result } @@ -399,14 +440,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 { @@ -535,9 +577,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 { @@ -556,6 +602,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 2953642f89..65191626bc 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()) @@ -229,7 +229,7 @@ func TestEncodeRoutedDMLEventUsesTargetNames(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := dml2rowEvent(t, decoder.NextDMLEvent()) + decoded := dml2rowEvent(t, decoder.NextDMLMessage().ToDMLEvent()) require.Equal(t, "target_db", decoded.TableInfo.GetSchemaName()) require.Equal(t, "target_table", decoded.TableInfo.GetTableName()) } @@ -296,7 +296,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()) @@ -711,7 +711,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()) @@ -772,7 +772,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()) } @@ -829,7 +829,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()) } @@ -893,7 +893,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() @@ -913,7 +913,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() @@ -933,7 +933,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 4fa008a27b..4d2514e783 100644 --- a/pkg/sink/codec/canal/canal_json_test.go +++ b/pkg/sink/codec/canal/canal_json_test.go @@ -20,6 +20,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/errors" @@ -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 1e8610c6cc..201eea1aba 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 5d94a1a6bc..b6b65eb2cc 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 new file mode 100644 index 0000000000..3e735745b1 --- /dev/null +++ b/pkg/sink/codec/debezium/avro_decoder.go @@ -0,0 +1,731 @@ +// 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 debezium + +import ( + "context" + "database/sql" + "encoding/binary" + "encoding/json" + "io" + "math/big" + "net/http" + "strconv" + "strings" + "sync" + + "github.com/linkedin/goavro/v2" + "github.com/pingcap/log" + commonEvent "github.com/pingcap/ticdc/pkg/common/event" + "github.com/pingcap/ticdc/pkg/errors" + codecavro "github.com/pingcap/ticdc/pkg/sink/codec/avro" + "github.com/pingcap/ticdc/pkg/sink/codec/common" + "go.uber.org/zap" +) + +const confluentAvroHeaderLen = 5 + +type avroDecoder struct { + ctx context.Context + registryURL string + httpClient *http.Client + inner *decoder + + mu sync.RWMutex + schemas map[int]*registeredDebeziumAvroSchema +} + +type registeredDebeziumAvroSchema struct { + schema any + namedSchemas map[string]any + codec *goavro.Codec +} + +// NewAvroDecoder returns a Debezium decoder for Confluent Avro wire-format +// messages. It decodes the Avro payload and then delegates Debezium event +// semantics to the JSON decoder. +func NewAvroDecoder( + ctx context.Context, + config *common.Config, + idx int, + db *sql.DB, +) (common.Decoder, error) { + registryURL := strings.TrimRight(config.AvroConfluentSchemaRegistry, "/") + if registryURL == "" { + return nil, errors.ErrAvroSchemaAPIError.GenWithStackByArgs("schema registry URI is empty") + } + + return &avroDecoder{ + ctx: ctx, + registryURL: registryURL, + httpClient: http.DefaultClient, + inner: NewDecoder(config, idx, db).(*decoder), + schemas: make(map[int]*registeredDebeziumAvroSchema), + }, nil +} + +func (d *avroDecoder) AddKeyValue(key, value []byte) { + keyJSON, err := d.toDebeziumJSON(key) + if err != nil { + log.Panic("decode Debezium Avro key failed", zap.Error(err), zap.Int("keySize", len(key))) + } + valueJSON, err := d.toDebeziumJSON(value) + if err != nil { + log.Panic("decode Debezium Avro value failed", zap.Error(err), zap.Int("valueSize", len(value))) + } + d.inner.AddKeyValue(keyJSON, valueJSON) +} + +func (d *avroDecoder) HasNext() (common.MessageType, bool) { + return d.inner.HasNext() +} + +func (d *avroDecoder) NextResolvedEvent() uint64 { + return d.inner.NextResolvedEvent() +} + +func (d *avroDecoder) NextDMLMessage() *common.DMLMessage { + return d.inner.NextDMLMessage() +} + +func (d *avroDecoder) NextDDLEvent() *commonEvent.DDLEvent { + return d.inner.NextDDLEvent() +} + +func (d *avroDecoder) toDebeziumJSON(data []byte) ([]byte, error) { + payload, schema, err := d.decodeConfluentAvroMessage(data) + if err != nil { + return nil, err + } + message := map[string]any{ + "schema": schema, + "payload": payload, + } + result, err := json.Marshal(message) + if err != nil { + return nil, errors.WrapError(errors.ErrDebeziumInvalidMessage, err) + } + return result, nil +} + +func (d *avroDecoder) decodeConfluentAvroMessage(data []byte) (any, map[string]any, error) { + if len(data) == 0 { + return nil, nil, errors.ErrDebeziumEmptyValueMessage.GenWithStackByArgs() + } + if len(data) < confluentAvroHeaderLen { + return nil, nil, errors.ErrAvroInvalidMessage.GenWithStackByArgs("confluent header is too short") + } + if data[0] != 0 { + return nil, nil, errors.ErrAvroInvalidMessage.GenWithStackByArgs("invalid confluent magic byte") + } + + schemaID := int(binary.BigEndian.Uint32(data[1:confluentAvroHeaderLen])) + registeredSchema, err := d.getSchema(schemaID) + if err != nil { + return nil, nil, err + } + + native, _, err := registeredSchema.codec.NativeFromBinary(data[confluentAvroHeaderLen:]) + if err != nil { + return nil, nil, errors.WrapError(errors.ErrAvroInvalidMessage, err) + } + + payload, err := avroNativeToConnectPayload( + registeredSchema.schema, + native, + registeredSchema.namedSchemas, + ) + if err != nil { + return nil, nil, err + } + schema, err := avroSchemaToConnectSchema( + registeredSchema.schema, + "", + nil, + registeredSchema.namedSchemas, + ) + if err != nil { + return nil, nil, err + } + return payload, schema, nil +} + +func (d *avroDecoder) getSchema(schemaID int) (*registeredDebeziumAvroSchema, error) { + d.mu.RLock() + schema, ok := d.schemas[schemaID] + d.mu.RUnlock() + if ok { + return schema, nil + } + + uri := d.registryURL + "/schemas/ids/" + strconv.Itoa(schemaID) + req, err := http.NewRequestWithContext(d.ctx, http.MethodGet, uri, nil) + if err != nil { + return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) + } + req.Header.Add( + "Accept", + "application/vnd.schemaregistry.v1+json, application/vnd.schemaregistry+json, "+ + "application/json", + ) + + resp, err := d.httpClient.Do(req) + if err != nil { + return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) + } + defer func() { + _ = resp.Body.Close() + }() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) + } + if resp.StatusCode != http.StatusOK { + return nil, errors.ErrAvroSchemaAPIError.GenWithStackByArgs( + "failed to query schema id " + strconv.Itoa(schemaID)) + } + + var lookupResp struct { + Schema string `json:"schema"` + } + if err := json.Unmarshal(body, &lookupResp); err != nil { + return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) + } + + codec, err := codecavro.GenCodec(lookupResp.Schema) + if err != nil { + return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) + } + + decoder := json.NewDecoder(strings.NewReader(lookupResp.Schema)) + decoder.UseNumber() + var schemaDef any + if err := decoder.Decode(&schemaDef); err != nil { + return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) + } + + namedSchemas := make(map[string]any) + collectAvroNamedSchemas(schemaDef, namedSchemas) + + schema = ®isteredDebeziumAvroSchema{ + schema: schemaDef, + namedSchemas: namedSchemas, + codec: codec, + } + d.mu.Lock() + d.schemas[schemaID] = schema + d.mu.Unlock() + return schema, nil +} + +func avroNativeToConnectPayload(schema any, value any, namedSchemas map[string]any) (any, error) { + switch typedSchema := schema.(type) { + case []any: + if value == nil { + return nil, nil + } + branchSchema, branchValue, err := avroUnionBranch(typedSchema, value) + if err != nil { + return nil, err + } + return avroNativeToConnectPayload(branchSchema, branchValue, namedSchemas) + case map[string]any: + rawType, ok := typedSchema["type"] + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema is missing type") + } + if unionType, ok := rawType.([]any); ok { + return avroNativeToConnectPayload(unionType, value, namedSchemas) + } + typeName, ok := rawType.(string) + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema type is invalid") + } + switch typeName { + case "record": + valueMap, ok := value.(map[string]any) + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro record payload is invalid") + } + fields, ok := typedSchema["fields"].([]any) + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro record schema is missing fields") + } + result := make(map[string]any, len(fields)) + for _, rawField := range fields { + field, ok := rawField.(map[string]any) + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro field schema is invalid") + } + avroFieldName, ok := field["name"].(string) + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro field is missing name") + } + connectFieldName := avroConnectFieldName(field, avroFieldName) + rawValue, exists := valueMap[avroFieldName] + if !exists && connectFieldName != avroFieldName { + rawValue, exists = valueMap[connectFieldName] + } + if !exists { + rawValue, exists = avroMissingFieldValue(field) + } + if !exists { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs( + "avro record payload is missing field " + avroFieldName) + } + fieldValue, err := avroNativeToConnectPayload( + field["type"], + rawValue, + namedSchemas, + ) + if err != nil { + return nil, err + } + result[connectFieldName] = fieldValue + } + return result, nil + case "array": + items, ok := typedSchema["items"] + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro array schema is missing items") + } + if value == nil { + return []any{}, nil + } + values, ok := value.([]any) + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro array payload is invalid") + } + result := make([]any, 0, len(values)) + for _, item := range values { + itemValue, err := avroNativeToConnectPayload(items, item, namedSchemas) + if err != nil { + return nil, err + } + result = append(result, itemValue) + } + return result, nil + case "bytes": + if avroSchemaIsDecimal(typedSchema) { + return avroDecimalNativeToString(typedSchema, value) + } + return value, nil + default: + return value, nil + } + case string: + if namedSchema, ok := namedSchemas[typedSchema]; ok { + return avroNativeToConnectPayload(namedSchema, value, namedSchemas) + } + return value, nil + default: + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema is invalid") + } +} + +func avroSchemaToConnectSchema( + schema any, + fieldName string, + fieldMeta map[string]any, + namedSchemas map[string]any, +) (map[string]any, error) { + switch typedSchema := schema.(type) { + case []any: + branchSchema, _, err := avroNonNullUnionBranch(typedSchema) + if err != nil { + return nil, err + } + connectSchema, err := avroSchemaToConnectSchema( + branchSchema, + fieldName, + fieldMeta, + namedSchemas, + ) + if err != nil { + return nil, err + } + connectSchema["optional"] = true + return connectSchema, nil + case map[string]any: + rawType, ok := typedSchema["type"] + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema is missing type") + } + if unionType, ok := rawType.([]any); ok { + return avroSchemaToConnectSchema(unionType, fieldName, fieldMeta, namedSchemas) + } + typeName, ok := rawType.(string) + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema type is invalid") + } + switch typeName { + case "record": + connectSchema := newConnectSchema("struct", false, fieldName, typedSchema, fieldMeta) + fields, ok := typedSchema["fields"].([]any) + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro record schema is missing fields") + } + connectFields := make([]any, 0, len(fields)) + for _, rawField := range fields { + field, ok := rawField.(map[string]any) + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro field schema is invalid") + } + avroFieldName, ok := field["name"].(string) + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro field is missing name") + } + fieldSchema, err := avroSchemaToConnectSchema( + field["type"], + avroConnectFieldName(field, avroFieldName), + field, + namedSchemas, + ) + if err != nil { + return nil, err + } + connectFields = append(connectFields, fieldSchema) + } + connectSchema["fields"] = connectFields + return connectSchema, nil + case "array": + connectSchema := newConnectSchema("array", false, fieldName, typedSchema, fieldMeta) + items, ok := typedSchema["items"] + if !ok { + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro array schema is missing items") + } + connectItems, err := avroSchemaToConnectSchema(items, "", nil, namedSchemas) + if err != nil { + return nil, err + } + connectSchema["items"] = connectItems + return connectSchema, nil + default: + connectType, err := avroPrimitiveToConnectType(typeName, typedSchema) + if err != nil { + return nil, err + } + return newConnectSchema(connectType, false, fieldName, typedSchema, fieldMeta), nil + } + case string: + if namedSchema, ok := namedSchemas[typedSchema]; ok { + return avroSchemaToConnectSchema(namedSchema, fieldName, fieldMeta, namedSchemas) + } + connectType, err := avroPrimitiveToConnectType(typedSchema, nil) + if err != nil { + return nil, err + } + return newConnectSchema(connectType, false, fieldName, nil, fieldMeta), nil + default: + return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema is invalid") + } +} + +func collectAvroNamedSchemas(schema any, namedSchemas map[string]any) { + switch typedSchema := schema.(type) { + case []any: + for _, branch := range typedSchema { + collectAvroNamedSchemas(branch, namedSchemas) + } + case map[string]any: + rawType := typedSchema["type"] + if unionType, ok := rawType.([]any); ok { + collectAvroNamedSchemas(unionType, namedSchemas) + return + } + typeName, _ := rawType.(string) + switch typeName { + case "record": + name := avroBranchName(typedSchema) + if name != "" { + namedSchemas[name] = typedSchema + shortName := avroShortBranchName(name) + if _, exists := namedSchemas[shortName]; shortName != "" && !exists { + namedSchemas[shortName] = typedSchema + } + } + fields, _ := typedSchema["fields"].([]any) + for _, rawField := range fields { + field, ok := rawField.(map[string]any) + if !ok { + continue + } + collectAvroNamedSchemas(field["type"], namedSchemas) + } + case "array": + collectAvroNamedSchemas(typedSchema["items"], namedSchemas) + } + } +} + +func newConnectSchema( + connectType string, + optional bool, + fieldName string, + schemaMeta map[string]any, + fieldMeta map[string]any, +) map[string]any { + connectSchema := map[string]any{ + "type": connectType, + "optional": optional, + } + if fieldName != "" { + connectSchema["field"] = fieldName + } + addConnectSchemaMetadata(connectSchema, schemaMeta) + addConnectFieldMetadata(connectSchema, fieldMeta) + return connectSchema +} + +func addConnectSchemaMetadata(connectSchema map[string]any, schemaMeta map[string]any) { + if schemaMeta == nil { + return + } + if name, ok := schemaMeta["connect.name"].(string); ok && name != "" { + connectSchema["name"] = name + } + if version, ok := schemaMeta["connect.version"]; ok { + connectSchema["version"] = version + } + if parameters, ok := schemaMeta["connect.parameters"].(map[string]any); ok { + connectSchema["parameters"] = parameters + } + if tidbType, ok := schemaMeta[debeziumAvroTiDBTypeKey].(string); ok && tidbType != "" { + connectSchema[debeziumAvroTiDBTypeKey] = tidbType + } +} + +func addConnectFieldMetadata(connectSchema map[string]any, fieldMeta map[string]any) { + if fieldMeta == nil { + return + } + if tidbType, ok := fieldMeta[debeziumAvroTiDBTypeKey].(string); ok && tidbType != "" { + connectSchema[debeziumAvroTiDBTypeKey] = tidbType + } +} + +func avroPrimitiveToConnectType(avroType string, schemaMeta map[string]any) (string, error) { + if schemaMeta != nil { + if connectType, ok := schemaMeta["connect.type"].(string); ok && connectType != "" { + return connectType, nil + } + } + switch avroType { + case "boolean": + return "boolean", nil + case "string": + return "string", nil + case "bytes": + return "bytes", nil + case "int": + return "int32", nil + case "long": + return "int64", nil + case "float": + return "float", nil + case "double": + return "double", nil + default: + return "", errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("unsupported avro type " + avroType) + } +} + +func avroSchemaIsDecimal(schema map[string]any) bool { + typeName, _ := schema["type"].(string) + logicalType, _ := schema["logicalType"].(string) + return typeName == "bytes" && logicalType == "decimal" +} + +func avroDecimalNativeToString(schema map[string]any, value any) (string, error) { + scale, err := avroDecimalScale(schema) + if err != nil { + return "", err + } + switch v := value.(type) { + case *big.Rat: + return v.FloatString(scale), nil + case big.Rat: + return v.FloatString(scale), nil + case string: + return v, nil + default: + return "", errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("decimal payload is invalid") + } +} + +func avroDecimalScale(schema map[string]any) (int, error) { + switch scale := schema["scale"].(type) { + case float64: + return int(scale), nil + case int: + return scale, nil + case int32: + return int(scale), nil + case int64: + return int(scale), nil + case json.Number: + value, err := scale.Int64() + if err != nil { + return 0, errors.WrapError(errors.ErrDebeziumInvalidMessage, err) + } + return int(value), nil + default: + return 0, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("decimal schema is missing scale") + } +} + +func avroUnionBranch(union []any, value any) (any, any, error) { + if value == nil { + return nil, nil, nil + } + var wrappedBranchName string + var wrappedBranchValue any + hasWrappedBranch := false + if branchValueMap, ok := value.(map[string]any); ok && len(branchValueMap) == 1 { + for branchName, branchValue := range branchValueMap { + wrappedBranchName = branchName + wrappedBranchValue = branchValue + hasWrappedBranch = true + for _, branchSchema := range union { + if avroBranchName(branchSchema) == branchName { + return branchSchema, branchValue, nil + } + } + } + } + + branchSchema, isSingleNonNullBranch, err := avroNonNullUnionBranch(union) + if err != nil { + return nil, nil, err + } + if hasWrappedBranch && + isSingleNonNullBranch && + avroShortBranchName(branchSchema) == avroShortBranchName(wrappedBranchName) { + return branchSchema, wrappedBranchValue, nil + } + return branchSchema, value, nil +} + +func avroNonNullUnionBranch(union []any) (any, bool, error) { + var result any + count := 0 + for _, branch := range union { + if avroBranchName(branch) != "null" { + if count == 0 { + result = branch + } + count++ + } + } + if count > 0 { + return result, count == 1, nil + } + return nil, false, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro union has no non-null branch") +} + +func avroBranchName(schema any) string { + switch typedSchema := schema.(type) { + case string: + return typedSchema + case map[string]any: + typeName, _ := typedSchema["type"].(string) + switch typeName { + case "record": + name, _ := typedSchema["name"].(string) + namespace, _ := typedSchema["namespace"].(string) + if namespace != "" && name != "" { + return namespace + "." + name + } + return name + case "array": + return "array" + default: + if avroSchemaIsDecimal(typedSchema) { + return "bytes.decimal" + } + return typeName + } + default: + return "" + } +} + +func avroShortBranchName(schema any) string { + switch typedSchema := schema.(type) { + case string: + if idx := strings.LastIndex(typedSchema, "."); idx >= 0 { + return typedSchema[idx+1:] + } + return typedSchema + case map[string]any: + typeName, _ := typedSchema["type"].(string) + if typeName == "record" { + name, _ := typedSchema["name"].(string) + return name + } + return avroBranchName(schema) + default: + return "" + } +} + +func avroFieldAllowsMissing(field map[string]any) bool { + if _, hasDefault := field["default"]; hasDefault { + return true + } + return avroSchemaAllowsNull(field["type"]) +} + +func avroMissingFieldValue(field map[string]any) (any, bool) { + if avroFieldAllowsMissing(field) { + return nil, true + } + if avroSchemaIsArray(field["type"]) { + return []any{}, true + } + return nil, false +} + +func avroSchemaIsArray(schema any) bool { + switch typedSchema := schema.(type) { + case map[string]any: + typeName, _ := typedSchema["type"].(string) + return typeName == "array" + case string: + return typedSchema == "array" + default: + return false + } +} + +func avroSchemaAllowsNull(schema any) bool { + union, ok := schema.([]any) + if !ok { + return false + } + for _, branch := range union { + if avroBranchName(branch) == "null" { + return true + } + } + return false +} + +func avroConnectFieldName(field map[string]any, fallback string) string { + if fieldName, ok := field[debeziumAvroConnectFieldKey].(string); ok && fieldName != "" { + return fieldName + } + return fallback +} diff --git a/pkg/sink/codec/debezium/avro_test.go b/pkg/sink/codec/debezium/avro_test.go new file mode 100644 index 0000000000..a61e2094ed --- /dev/null +++ b/pkg/sink/codec/debezium/avro_test.go @@ -0,0 +1,682 @@ +// 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 debezium + +import ( + "context" + "encoding/binary" + "encoding/json" + "fmt" + "io" + "net/http" + "testing" + "time" + + "github.com/pingcap/ticdc/downstreamadapter/sink/columnselector" + commonEvent "github.com/pingcap/ticdc/pkg/common/event" + "github.com/pingcap/ticdc/pkg/config" + "github.com/pingcap/ticdc/pkg/sink/codec/avro" + "github.com/pingcap/ticdc/pkg/sink/codec/common" + "github.com/stretchr/testify/require" +) + +func TestDebeziumConfluentAvroEncodeRowEvent(t *testing.T) { + ctx := context.Background() + _, err := avro.SetupEncoderAndSchemaRegistry4Testing( + ctx, + common.NewConfig(config.ProtocolAvro), + ) + require.NoError(t, err) + defer avro.TeardownEncoderAndSchemaRegistry4Testing() + + helper := NewSQLTestHelper(t, "foo", ` + create table foo( + id int primary key, + name varchar(16), + bin varbinary(16), + price decimal(10, 4), + ubig bigint unsigned, + v bigint null + )`) + defer helper.Close() + + dmls := helper.helper.DML2Event("test", "foo", + "insert into foo values (1, 'alice', x'010203', 12.3400, 18446744073709551615, null)") + row, ok := dmls.GetNextRow() + require.True(t, ok) + + cfg := common.NewConfig(config.ProtocolDebeziumAvro) + cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" + cfg.AvroBigintUnsignedHandlingMode = common.BigintUnsignedHandlingModeString + cfg.DebeziumDisableSchema = true + cfg.TimeZone = time.UTC + + encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") + require.NoError(t, err) + require.NoError(t, encoder.AppendRowChangedEvent(ctx, "dbserver1.test.foo", &commonEvent.RowEvent{ + TableInfo: helper.tableInfo, + CommitTs: 1, + Event: row, + ColumnSelector: columnselector.NewDefaultColumnSelector(), + Callback: func() {}, + })) + + messages := encoder.Build() + require.Len(t, messages, 1) + require.Equal(t, byte(0), messages[0].Key[0]) + require.Equal(t, byte(0), messages[0].Value[0]) + + key := decodeConfluentAvroForTest(t, messages[0].Key) + require.Equal(t, int32(1), unwrapAvroUnionForTest(t, key["id"], "int")) + + value := decodeConfluentAvroForTest(t, messages[0].Value) + require.Equal(t, "c", value["op"]) + require.Nil(t, value["before"]) + require.NotContains(t, value, "transaction") + require.IsType(t, int64(0), value["ts_ms"]) + + afterUnion, ok := value["after"].(map[string]any) + require.True(t, ok) + after, ok := afterUnion["dbserver1.test.foo"].(map[string]any) + require.True(t, ok) + require.Equal(t, int32(1), unwrapAvroUnionForTest(t, after["id"], "int")) + require.Equal(t, "alice", unwrapAvroUnionForTest(t, after["name"], "string")) + require.Equal(t, []byte{1, 2, 3}, unwrapAvroUnionForTest(t, after["bin"], "bytes")) + require.Equal(t, "18446744073709551615", unwrapAvroUnionForTest(t, after["ubig"], "string")) + require.Nil(t, after["v"]) + + source, ok := value["source"].(map[string]any) + require.True(t, ok) + require.Equal(t, "test", source["db"]) + require.Equal(t, "foo", source["table"]) + require.Nil(t, source["snapshot"]) + require.Nil(t, source["thread"]) + require.Equal(t, "dbserver1", source["name"]) + + valueSchema := decodeConfluentAvroSchemaForTest(t, messages[0].Value) + require.Contains(t, valueSchema, `"name":"fooEnvelope"`) + require.Contains(t, valueSchema, `"name":"foo"`) + require.Contains(t, valueSchema, `"name":"Source"`) + require.Contains(t, valueSchema, `"logicalType":"decimal"`) + require.NotContains(t, valueSchema, `"field":"transaction"`) +} + +func TestDebeziumConfluentAvroSanitizesFullNameAndUnionBranch(t *testing.T) { + ctx := context.Background() + _, err := avro.SetupEncoderAndSchemaRegistry4Testing( + ctx, + common.NewConfig(config.ProtocolAvro), + ) + require.NoError(t, err) + defer avro.TeardownEncoderAndSchemaRegistry4Testing() + + cfg := common.NewConfig(config.ProtocolDebeziumAvro) + cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" + cfg.TimeZone = time.UTC + + eventEncoder, err := NewAvroBatchEncoder(ctx, cfg, "db-server") + require.NoError(t, err) + encoder, ok := eventEncoder.(*BatchEncoder) + require.True(t, ok) + + payload, err := encoder.encodeAvroPayload( + ctx, + "topic-with-invalid-schema-name", + debeziumAvroValueSchemaSuffix, + &debeziumAvroMessage{ + Schema: &debeziumConnectSchema{ + Type: "struct", + Name: "db-server.test-db.foo-barEnvelope", + Fields: []*debeziumConnectSchema{ + { + Type: "struct", + Optional: true, + Name: "db-server.test-db.foo-bar", + Field: "after", + Fields: []*debeziumConnectSchema{ + { + Type: "int32", + Field: "id", + }, + }, + }, + { + Type: "string", + Field: "op", + }, + }, + }, + Payload: map[string]any{ + "after": map[string]any{ + "id": int32(1), + }, + "op": "c", + }, + }, + 1, + ) + require.NoError(t, err) + + schema := decodeConfluentAvroSchemaForTest(t, payload) + require.Contains(t, schema, `"name":"foo_barEnvelope"`) + require.Contains(t, schema, `"namespace":"db_server.test_db"`) + require.Contains(t, schema, `"name":"foo_bar"`) + require.NotContains(t, schema, `"namespace":"db-server.test-db"`) + + value := decodeConfluentAvroForTest(t, payload) + after := unwrapAvroUnionForTest(t, value["after"], "db_server.test_db.foo_bar") + require.Equal(t, map[string]any{"id": int32(1)}, after) +} + +func TestDebeziumConfluentAvroDecodeRowEvent(t *testing.T) { + ctx := context.Background() + _, err := avro.SetupEncoderAndSchemaRegistry4Testing( + ctx, + common.NewConfig(config.ProtocolAvro), + ) + require.NoError(t, err) + defer avro.TeardownEncoderAndSchemaRegistry4Testing() + + helper := NewSQLTestHelper(t, "foo", ` + create table foo( + id int primary key, + name varchar(16), + bin varbinary(16), + price decimal(10, 4), + ubig bigint unsigned, + v bigint null + )`) + defer helper.Close() + + dmls := helper.helper.DML2Event("test", "foo", + "insert into foo values (1, 'alice', x'010203', 12.3400, 18446744073709551615, null)") + row, ok := dmls.GetNextRow() + require.True(t, ok) + + cfg := common.NewConfig(config.ProtocolDebeziumAvro) + cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" + cfg.AvroBigintUnsignedHandlingMode = common.BigintUnsignedHandlingModeString + cfg.EnableTiDBExtension = true + cfg.TimeZone = time.UTC + + commitTs := uint64(123) + encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") + require.NoError(t, err) + require.NoError(t, encoder.AppendRowChangedEvent(ctx, "dbserver1.test.foo", &commonEvent.RowEvent{ + TableInfo: helper.tableInfo, + CommitTs: commitTs, + Event: row, + ColumnSelector: columnselector.NewDefaultColumnSelector(), + Callback: func() {}, + })) + + messages := encoder.Build() + require.Len(t, messages, 1) + + decoder, err := NewAvroDecoder(ctx, cfg, 0, nil) + require.NoError(t, err) + decoder.AddKeyValue(messages[0].Key, messages[0].Value) + + messageType, hasNext := decoder.HasNext() + require.True(t, hasNext) + require.Equal(t, common.MessageTypeRow, messageType) + + decoded := decoder.NextDMLMessage().ToDMLEvent() + require.Equal(t, commitTs, decoded.CommitTs) + require.Equal(t, "test", decoded.TableInfo.GetSchemaName()) + require.Equal(t, "foo", decoded.TableInfo.GetTableName()) + + change, ok := decoded.GetNextRow() + require.True(t, ok) + common.CompareRow(t, row, helper.tableInfo, change, decoded.TableInfo) +} + +func TestDebeziumConfluentAvroDecodeAccountDMLEvents(t *testing.T) { + ctx := context.Background() + _, err := avro.SetupEncoderAndSchemaRegistry4Testing( + ctx, + common.NewConfig(config.ProtocolAvro), + ) + require.NoError(t, err) + defer avro.TeardownEncoderAndSchemaRegistry4Testing() + + helper := NewSQLTestHelper(t, "tp_account", ` + create table tp_account( + id int primary key, + account_id int not null + )`) + defer helper.Close() + + insertDML := helper.helper.DML2Event("test", "tp_account", + "insert into tp_account values (12, 34)") + updateDML, _ := helper.helper.DML2UpdateEvent("test", "tp_account", + "insert into tp_account values (13, 34)", + "update tp_account set account_id = 35 where id = 13") + deleteDML := helper.helper.DML2DeleteEvent("test", "tp_account", + "insert into tp_account values (14, 34)", + "delete from tp_account where id = 14") + + cfg := common.NewConfig(config.ProtocolDebeziumAvro) + cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" + cfg.EnableTiDBExtension = true + cfg.TimeZone = time.UTC + + encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") + require.NoError(t, err) + + rows := make([]commonEvent.RowChange, 0, 3) + for _, dml := range []*commonEvent.DMLEvent{insertDML, updateDML, deleteDML} { + row, ok := dml.GetNextRow() + if !ok { + continue + } + rows = append(rows, row) + require.NoError(t, encoder.AppendRowChangedEvent(ctx, "dbserver1.test.tp_account", &commonEvent.RowEvent{ + TableInfo: helper.tableInfo, + CommitTs: 1, + Event: row, + ColumnSelector: columnselector.NewDefaultColumnSelector(), + Callback: func() {}, + })) + } + require.Len(t, rows, 3) + + messages := encoder.Build() + require.Len(t, messages, 3) + for idx, message := range messages { + decoder, err := NewAvroDecoder(ctx, cfg, 0, nil) + require.NoError(t, err) + decoder.AddKeyValue(message.Key, message.Value) + + messageType, hasNext := decoder.HasNext() + require.True(t, hasNext) + require.Equal(t, common.MessageTypeRow, messageType) + + decoded := decoder.NextDMLMessage().ToDMLEvent() + require.Equal(t, "test", decoded.TableInfo.GetSchemaName()) + require.Equal(t, "tp_account", decoded.TableInfo.GetTableName()) + + change, ok := decoded.GetNextRow() + require.True(t, ok) + common.CompareRow(t, rows[idx], helper.tableInfo, change, decoded.TableInfo) + } +} + +func TestDebeziumConfluentAvroDecodeShortNamedUnionBranch(t *testing.T) { + valueSchema := map[string]any{ + "type": "record", + "name": "Value", + "namespace": "dbserver1.test.tp_account", + "fields": []any{ + map[string]any{ + "name": "id", + "type": "int", + debeziumAvroConnectFieldKey: "id", + }, + map[string]any{ + "name": "account_id", + "type": "int", + debeziumAvroConnectFieldKey: "account_id", + }, + }, + } + namedSchemas := map[string]any{ + "dbserver1.test.tp_account.Value": valueSchema, + } + + payload, err := avroNativeToConnectPayload( + []any{"null", "dbserver1.test.tp_account.Value"}, + map[string]any{ + "Value": map[string]any{ + "id": int32(12), + "account_id": int32(34), + }, + }, + namedSchemas, + ) + require.NoError(t, err) + require.Equal(t, map[string]any{ + "id": int32(12), + "account_id": int32(34), + }, payload) +} + +func TestDebeziumConfluentAvroDecodeFullNamedWrapperForShortUnionBranch(t *testing.T) { + valueSchema := map[string]any{ + "type": "record", + "name": "Value", + "namespace": "default.test.tp_account", + "fields": []any{ + map[string]any{ + "name": "id", + "type": "int", + debeziumAvroConnectFieldKey: "id", + }, + map[string]any{ + "name": "account_id", + "type": "int", + debeziumAvroConnectFieldKey: "account_id", + }, + }, + } + envelopeSchema := map[string]any{ + "type": "record", + "name": "Envelope", + "namespace": "default.test.tp_account", + "fields": []any{ + map[string]any{ + "name": "after", + "type": []any{"null", "Value"}, + }, + }, + } + namedSchemas := map[string]any{} + collectAvroNamedSchemas(valueSchema, namedSchemas) + collectAvroNamedSchemas(envelopeSchema, namedSchemas) + + payload, err := avroNativeToConnectPayload( + []any{"null", "Value"}, + map[string]any{ + "default.test.tp_account.Value": map[string]any{ + "id": int32(12), + "account_id": int32(34), + }, + }, + namedSchemas, + ) + require.NoError(t, err) + require.Equal(t, map[string]any{ + "id": int32(12), + "account_id": int32(34), + }, payload) +} + +func TestDebeziumConfluentAvroDecodeSingleFieldUnionRecord(t *testing.T) { + valueSchema := map[string]any{ + "type": "record", + "name": "Value", + "namespace": "dbserver1.test.only_pk", + "fields": []any{ + map[string]any{ + "name": "id", + "type": "int", + debeziumAvroConnectFieldKey: "id", + }, + }, + } + namedSchemas := map[string]any{ + "dbserver1.test.only_pk.Value": valueSchema, + } + + payload, err := avroNativeToConnectPayload( + []any{"null", "dbserver1.test.only_pk.Value"}, + map[string]any{ + "id": int32(12), + }, + namedSchemas, + ) + require.NoError(t, err) + require.Equal(t, map[string]any{ + "id": int32(12), + }, payload) +} + +func TestDebeziumConfluentAvroDecodeMissingRecordField(t *testing.T) { + valueSchema := map[string]any{ + "type": "record", + "name": "Value", + "namespace": "dbserver1.test.tp_account", + "fields": []any{ + map[string]any{ + "name": "id", + "type": "int", + debeziumAvroConnectFieldKey: "id", + }, + map[string]any{ + "name": "account_id", + "type": "int", + debeziumAvroConnectFieldKey: "account_id", + }, + }, + } + + _, err := avroNativeToConnectPayload( + valueSchema, + map[string]any{ + "id": int32(12), + }, + nil, + ) + require.Error(t, err) + require.Contains(t, err.Error(), "avro record payload is missing field account_id") +} + +func TestDebeziumConfluentAvroDecodeMissingOptionalRecordField(t *testing.T) { + valueSchema := map[string]any{ + "type": "record", + "name": "Table", + "namespace": "io.debezium.connector.schema", + "fields": []any{ + map[string]any{ + "name": "defaultCharsetName", + "type": []any{"null", "string"}, + "default": nil, + }, + map[string]any{ + "name": "columns", + "type": map[string]any{ + "type": "array", + "items": "string", + }, + }, + }, + } + + payload, err := avroNativeToConnectPayload( + valueSchema, + map[string]any{ + "columns": []any{"id"}, + }, + nil, + ) + require.NoError(t, err) + require.Equal(t, map[string]any{ + "defaultCharsetName": nil, + "columns": []any{"id"}, + }, payload) +} + +func TestDebeziumConfluentAvroDecodeMissingArrayRecordField(t *testing.T) { + valueSchema := map[string]any{ + "type": "record", + "name": "Table", + "namespace": "io.debezium.connector.schema", + "fields": []any{ + map[string]any{ + "name": "columns", + "type": map[string]any{ + "type": "array", + "items": "string", + }, + }, + }, + } + + payload, err := avroNativeToConnectPayload( + valueSchema, + map[string]any{}, + nil, + ) + require.NoError(t, err) + require.Equal(t, map[string]any{ + "columns": []any{}, + }, payload) +} + +func TestDebeziumConfluentAvroEncodeDDLEvent(t *testing.T) { + ctx := context.Background() + _, err := avro.SetupEncoderAndSchemaRegistry4Testing( + ctx, + common.NewConfig(config.ProtocolAvro), + ) + require.NoError(t, err) + defer avro.TeardownEncoderAndSchemaRegistry4Testing() + + cfg := common.NewConfig(config.ProtocolDebeziumAvro) + cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" + cfg.EnableTiDBExtension = true + cfg.TimeZone = time.UTC + + routedDDL := common.NewRoutedDDLEvent4Test() + + encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") + require.NoError(t, err) + message, err := encoder.EncodeDDLEvent(routedDDL) + require.NoError(t, err) + require.Nil(t, message) + + cfg.AvroEnableWatermark = true + encoder, err = NewAvroBatchEncoder(ctx, cfg, "dbserver1") + require.NoError(t, err) + message, err = encoder.EncodeDDLEvent(routedDDL) + require.NoError(t, err) + require.NotNil(t, message) + require.Equal(t, byte(0), message.Key[0]) + require.Equal(t, byte(0), message.Value[0]) + + decoder, err := NewAvroDecoder(ctx, cfg, 0, nil) + require.NoError(t, err) + decoder.AddKeyValue(message.Key, message.Value) + + messageType, hasNext := decoder.HasNext() + require.True(t, hasNext) + require.Equal(t, common.MessageTypeDDL, messageType) + + decoded := decoder.NextDDLEvent() + require.Equal(t, routedDDL.GetCommitTs(), decoded.GetCommitTs()) + require.Equal(t, routedDDL.GetDDLType(), decoded.GetDDLType()) + require.Equal(t, "target_db", decoded.SchemaName) + require.Equal(t, "target_table", decoded.TableName) + require.Equal(t, routedDDL.Query, decoded.Query) +} + +func TestDebeziumConfluentAvroEncodeCheckpointEvent(t *testing.T) { + ctx := context.Background() + _, err := avro.SetupEncoderAndSchemaRegistry4Testing( + ctx, + common.NewConfig(config.ProtocolAvro), + ) + require.NoError(t, err) + defer avro.TeardownEncoderAndSchemaRegistry4Testing() + + cfg := common.NewConfig(config.ProtocolDebeziumAvro) + cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" + cfg.EnableTiDBExtension = true + cfg.AvroEnableWatermark = true + cfg.TimeZone = time.UTC + + encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") + require.NoError(t, err) + + message, err := encoder.EncodeCheckpointEvent(100) + require.NoError(t, err) + require.NotNil(t, message) + + decoder, err := NewAvroDecoder(ctx, cfg, 0, nil) + require.NoError(t, err) + decoder.AddKeyValue(message.Key, message.Value) + + messageType, hasNext := decoder.HasNext() + require.True(t, hasNext) + require.Equal(t, common.MessageTypeResolved, messageType) + require.Equal(t, uint64(100), decoder.NextResolvedEvent()) +} + +func TestDebeziumConfluentAvroDoesNotEncodeCheckpointEventByDefault(t *testing.T) { + ctx := context.Background() + _, err := avro.SetupEncoderAndSchemaRegistry4Testing( + ctx, + common.NewConfig(config.ProtocolAvro), + ) + require.NoError(t, err) + defer avro.TeardownEncoderAndSchemaRegistry4Testing() + + cfg := common.NewConfig(config.ProtocolDebeziumAvro) + cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" + cfg.EnableTiDBExtension = true + cfg.TimeZone = time.UTC + + encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") + require.NoError(t, err) + + message, err := encoder.EncodeCheckpointEvent(100) + require.NoError(t, err) + require.Nil(t, message) +} + +func decodeConfluentAvroForTest(t *testing.T, data []byte) map[string]any { + t.Helper() + + schema, binaryData := decodeConfluentAvroEnvelopeForTest(t, data) + codec, err := avro.GenCodec(schema) + require.NoError(t, err) + + native, _, err := codec.NativeFromBinary(binaryData) + require.NoError(t, err) + + result, ok := native.(map[string]any) + require.True(t, ok) + return result +} + +func decodeConfluentAvroSchemaForTest(t *testing.T, data []byte) string { + t.Helper() + + schema, _ := decodeConfluentAvroEnvelopeForTest(t, data) + return schema +} + +func decodeConfluentAvroEnvelopeForTest(t *testing.T, data []byte) (string, []byte) { + t.Helper() + + require.GreaterOrEqual(t, len(data), 5) + require.Equal(t, byte(0), data[0]) + schemaID := int(binary.BigEndian.Uint32(data[1:5])) + binaryData := data[5:] + + resp, err := http.Get(fmt.Sprintf("http://127.0.0.1:8081/schemas/ids/%d", schemaID)) + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusOK, resp.StatusCode) + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + var schemaResp struct { + Schema string `json:"schema"` + } + require.NoError(t, json.Unmarshal(body, &schemaResp)) + + return schemaResp.Schema, binaryData +} + +func unwrapAvroUnionForTest(t *testing.T, value any, branch string) any { + t.Helper() + + union, ok := value.(map[string]any) + require.True(t, ok) + result, ok := union[branch] + require.True(t, ok) + return result +} diff --git a/pkg/sink/codec/debezium/debezium_test.go b/pkg/sink/codec/debezium/debezium_test.go index 4b716b404a..80380bfebf 100644 --- a/pkg/sink/codec/debezium/debezium_test.go +++ b/pkg/sink/codec/debezium/debezium_test.go @@ -120,7 +120,7 @@ func TestEncodeRoutedDMLEventUsesTargetNames(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() require.Equal(t, "target_db", decoded.TableInfo.GetSchemaName()) require.Equal(t, "target_table", decoded.TableInfo.GetTableName()) diff --git a/pkg/sink/codec/debezium/decoder.go b/pkg/sink/codec/debezium/decoder.go index ae44351003..2d2ae1675d 100644 --- a/pkg/sink/codec/debezium/decoder.go +++ b/pkg/sink/codec/debezium/decoder.go @@ -47,10 +47,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 @@ -149,20 +149,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, @@ -175,12 +213,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) @@ -199,7 +237,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))) @@ -208,15 +250,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() { @@ -226,28 +274,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 = parser_model.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: @@ -256,7 +307,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: parser_model.NewCIStr(colName), Offset: idx, @@ -281,8 +332,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 { @@ -293,7 +344,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 } @@ -446,18 +497,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 acd2d8cb77..04580dae5e 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 9b02366709..984e100ccb 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) @@ -673,7 +673,7 @@ func TestEncodeRoutedDMLEventUsesTargetNames(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decoded := decoder.NextDMLEvent() + decoded := decoder.NextDMLMessage().ToDMLEvent() require.Equal(t, "target_db", decoded.TableInfo.GetSchemaName()) require.Equal(t, "target_table", decoded.TableInfo.GetTableName()) @@ -757,7 +757,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) @@ -824,7 +824,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) @@ -834,7 +834,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) @@ -846,7 +846,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) @@ -930,7 +930,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) @@ -1027,7 +1027,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) @@ -1115,7 +1115,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())) @@ -1228,7 +1228,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) @@ -1258,7 +1258,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) @@ -1288,8 +1288,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 @@ -1368,7 +1366,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) @@ -1420,7 +1418,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) @@ -1465,7 +1463,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) @@ -1514,7 +1512,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) @@ -1564,7 +1562,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) @@ -1629,7 +1627,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()) @@ -1735,7 +1733,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 93aa7b8cf9..075fa38969 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 33997736af..181a314a40 100644 --- a/pkg/sink/codec/simple/encoder_test.go +++ b/pkg/sink/codec/simple/encoder_test.go @@ -135,7 +135,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) @@ -178,7 +178,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) } @@ -227,7 +227,7 @@ func TestEncodeRoutedEventsUsesTargetNames(t *testing.T) { require.True(t, hasNext) require.Equal(t, common.MessageTypeRow, messageType) - decodedDML := decoder.NextDMLEvent() + decodedDML := decoder.NextDMLMessage().ToDMLEvent() require.Equal(t, "target_db", decodedDML.TableInfo.GetSchemaName()) require.Equal(t, "target_table", decodedDML.TableInfo.GetTableName()) } @@ -307,7 +307,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()) @@ -916,7 +916,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()) @@ -965,7 +965,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()) @@ -1119,7 +1119,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() @@ -1194,7 +1194,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) @@ -1262,8 +1262,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() } @@ -1279,8 +1279,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()) @@ -1328,8 +1329,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) @@ -1343,8 +1344,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) @@ -1433,7 +1435,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) @@ -1517,7 +1519,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()) @@ -1642,7 +1644,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()) @@ -1714,8 +1716,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 @@ -1749,8 +1751,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) From 2bcf84200c8549c4eec38bb31cccc4f49c1be835 Mon Sep 17 00:00:00 2001 From: wk989898 Date: Mon, 3 Aug 2026 03:24:06 +0000 Subject: [PATCH 2/3] update Signed-off-by: wk989898 --- cmd/kafka-consumer/writer.go | 100 ++++++++------------- cmd/kafka-consumer/writer_test.go | 81 +++++++---------- cmd/pulsar-consumer/consumer.go | 4 - cmd/pulsar-consumer/writer.go | 84 ++++------------- cmd/pulsar-consumer/writer_test.go | 52 +---------- cmd/storage-consumer/consumer.go | 5 -- cmd/util/event_group.go | 96 -------------------- cmd/util/event_group_test.go | 46 ---------- pkg/sink/codec/canal/canal_json_decoder.go | 5 -- 9 files changed, 94 insertions(+), 379 deletions(-) diff --git a/cmd/kafka-consumer/writer.go b/cmd/kafka-consumer/writer.go index 9c300fb6f9..c8a3783c43 100644 --- a/cmd/kafka-consumer/writer.go +++ b/cmd/kafka-consumer/writer.go @@ -170,30 +170,12 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *event.DDLEvent) error { if !ok { continue } -<<<<<<< HEAD - before := len(resolvedEvents) - resolvedEvents = g.ResolveInto(commitTs, resolvedEvents) - resolvedCount := len(resolvedEvents) - before - if resolvedCount == 0 { - continue - } - - resolvedGroups = append(resolvedGroups, struct { - group *util.EventsGroup - maxCommitTs uint64 - }{ - group: g, - maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), - }) - total += resolvedCount -======= messages := g.ResolveInto(commitTs, nil) events := make([]*event.DMLEvent, 0, len(messages)) for _, message := range messages { events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) } resolvedEvents = append(resolvedEvents, events...) ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } } @@ -311,30 +293,12 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { }, 0) for _, p := range w.progresses { for _, group := range p.eventsGroup { -<<<<<<< HEAD - before := len(resolvedEvents) - resolvedEvents = group.ResolveInto(watermark, resolvedEvents) - resolvedCount := len(resolvedEvents) - before - if resolvedCount == 0 { - continue - } - - resolvedGroups = append(resolvedGroups, struct { - group *util.EventsGroup - maxCommitTs uint64 - }{ - group: group, - maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), - }) - total += resolvedCount -======= messages := group.ResolveInto(watermark, nil) events := make([]*event.DMLEvent, 0, len(messages)) for _, message := range messages { events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) } resolvedEvents = append(resolvedEvents, events...) ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } } total := len(resolvedEvents) @@ -567,7 +531,8 @@ func (w *writer) onDDL(ddl *event.DDLEvent) { return } switch w.protocol { - case config.ProtocolCanalJSON, config.ProtocolOpen, config.ProtocolAvro, config.ProtocolSimple, config.ProtocolDebezium: + case config.ProtocolCanalJSON, config.ProtocolOpen, config.ProtocolAvro, config.ProtocolSimple, + config.ProtocolDebezium, config.ProtocolDebeziumAvro: default: return } @@ -575,20 +540,54 @@ func (w *writer) onDDL(ddl *event.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 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 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) 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.GetTargetSchemaName(), ddl.TableInfo.GetTargetTableName()) + 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()) @@ -651,35 +650,15 @@ func (w *writer) appendMessage2Group(message *common.DMLMessage, progress *parti zap.String("schema", schema), zap.String("table", table), zap.Any("protocol", w.protocol)) return } -<<<<<<< HEAD - 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", ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) 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.String("schema", schema), zap.String("table", table), zap.Int64("tableID", tableID), -<<<<<<< HEAD - zap.Stringer("eventType", dml.RowTypes[0]), zap.Any("protocol", w.protocol), - zap.Bool("IsPartition", dml.TableInfo.TableName.IsPartition)) - group.Append(dml, true) - 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])) -======= zap.Stringer("eventType", message.RowType)) return } @@ -731,7 +710,6 @@ func (w *writer) appendMessage2Group(message *common.DMLMessage, progress *parti default: log.Panic("unknown protocol", zap.Any("protocol", w.protocol)) } ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } 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 c40871e0b5..3776eb135b 100644 --- a/cmd/kafka-consumer/writer_test.go +++ b/cmd/kafka-consumer/writer_test.go @@ -24,7 +24,7 @@ import ( "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" + 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" @@ -98,7 +98,7 @@ func TestWriterWrite_executesIndependentCreateTableWithoutWatermark(t *testing.T }, } - w.Write(ctx, codecCommon.MessageTypeDDL) + w.Write(ctx, codeccommon.MessageTypeDDL) require.Equal(t, []string{"CREATE TABLE `test`.`t` (`id` INT PRIMARY KEY)"}, s.ddls) require.Empty(t, w.ddlList) @@ -144,12 +144,12 @@ func TestWriterWrite_preservesOrderWhenBlockedDDLNotReady(t *testing.T) { }, } - w.Write(ctx, codecCommon.MessageTypeDDL) + w.Write(ctx, codeccommon.MessageTypeDDL) require.Empty(t, s.ddls) require.Len(t, w.ddlList, 2) p.watermark = 200 - w.Write(ctx, codecCommon.MessageTypeDDL) + w.Write(ctx, codeccommon.MessageTypeDDL) require.Equal(t, []string{ "ALTER TABLE `test`.`t` ADD COLUMN `c2` INT", "CREATE TABLE `test`.`t2` (`id` INT PRIMARY KEY)", @@ -190,12 +190,12 @@ func TestWriterWrite_doesNotBypassWatermarkForCreateTableLike(t *testing.T) { }, } - w.Write(ctx, codecCommon.MessageTypeDDL) + w.Write(ctx, codeccommon.MessageTypeDDL) require.Empty(t, s.ddls) require.Len(t, w.ddlList, 1) p.watermark = 200 - w.Write(ctx, codecCommon.MessageTypeDDL) + w.Write(ctx, codeccommon.MessageTypeDDL) require.Equal(t, []string{"CREATE TABLE `test`.`t2` LIKE `test`.`t1`"}, s.ddls) require.Empty(t, w.ddlList) } @@ -272,7 +272,7 @@ func TestWriterWrite_handlesOutOfOrderDDLsByCommitTs(t *testing.T) { }, } - w.Write(ctx, codecCommon.MessageTypeDDL) + w.Write(ctx, codeccommon.MessageTypeDDL) require.Equal(t, []string{ "CREATE TABLE `common_1`.`add_and_drop_columns` (`id` INT(11) NOT NULL PRIMARY KEY)", @@ -284,68 +284,54 @@ 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 Kafka in commitTs order. - // 2) Under network partition / changefeed restart, TiCDC may replay older commitTs, - // which will be appended to Kafka at a larger offset (commitTs appears to go backwards). - // - // The kafka-consumer must not drop these "fallback commitTs" events unless they have - // already been flushed to downstream (AppliedWatermark), otherwise the replay cannot - // heal the missing window. +func TestOnDDLMarksRoutedCreateTableLikePartitionTableForAvro(t *testing.T) { replicaCfg := config.GetDefaultReplicaConfig() - eventRouter, err := eventrouter.NewEventRouter(replicaCfg.Sink, "test-topic", false, false) + eventRouter, err := eventrouter.NewEventRouter(replicaCfg.Sink, "test-topic", false, true) require.NoError(t, err) w := &writer{ progresses: []*partitionProgress{{partition: 0, eventsGroup: make(map[int64]*util.EventsGroup)}}, eventRouter: eventRouter, - protocol: config.ProtocolCanalJSON, - partitionTableAccessor: codecCommon.NewPartitionTableAccessor(), + protocol: config.ProtocolAvro, + partitionTableAccessor: codeccommon.NewPartitionTableAccessor(), + } + + ddl := &commonEvent.DDLEvent{ + Query: "CREATE TABLE `target`.`dst` LIKE `target`.`src`", + SchemaName: "source", + TableName: "dst", + Type: byte(timodel.ActionCreateTable), + TableInfo: &common.TableInfo{ + TableName: common.TableName{ + Schema: "source", + Table: "dst", + IsPartition: true, + TargetSchema: "target", + TargetTable: "dst", + }, + }, } + w.onDDL(ddl) + require.True(t, w.partitionTableAccessor.IsPartitionTable("target", "dst")) - newDMLEvent := func(tableID int64, commitTs uint64) *commonEvent.DMLEvent { + newDMLEvent := func(commitTs uint64) *commonEvent.DMLEvent { return &commonEvent.DMLEvent{ - PhysicalTableID: tableID, + PhysicalTableID: 1, CommitTs: commitTs, RowTypes: []common.RowType{common.RowTypeUpdate}, Rows: chunk.NewChunkWithCapacity(nil, 0), TableInfo: &common.TableInfo{ - TableName: common.TableName{Schema: "test", Table: "t"}, + TableName: common.TableName{Schema: "target", Table: "dst"}, }, } } progress := w.progresses[0] -<<<<<<< HEAD -======= w.appendMessage2Group(codeccommon.NewDMLMessageFromEvent(newDMLEvent(200)), progress, kafka.Offset(10)) w.appendMessage2Group(codeccommon.NewDMLMessageFromEvent(newDMLEvent(100)), progress, kafka.Offset(11)) ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) - - // Step 1: observe a larger commitTs first (e.g. produced before restart). - w.appendRow2Group(newDMLEvent(1, 200), progress, kafka.Offset(10)) - - // Step 2: observe a smaller commitTs later (e.g. replayed after restart). - w.appendRow2Group(newDMLEvent(1, 100), progress, kafka.Offset(11)) - - group := progress.eventsGroup[1] - require.NotNil(t, group) - resolvedEvents := make([]*commonEvent.DMLEvent, 0) - // Expect: commitTs=100 is still kept and can be resolved. - resolved := group.ResolveInto(150, nil) + resolved := progress.eventsGroup[1].ResolveInto(150, nil) require.Len(t, resolved, 1) -<<<<<<< HEAD - require.Equal(t, uint64(100), resolved[0].CommitTs) - - // Step 3: once downstream has flushed beyond commitTs=100, the replay is safe to ignore. - resolvedEvents = make([]*commonEvent.DMLEvent, 0) - group.AppliedWatermark = 200 - w.appendRow2Group(newDMLEvent(1, 100), progress, kafka.Offset(12)) - resolved = group.ResolveInto(150, resolvedEvents) - require.Empty(t, resolved) -======= require.Equal(t, uint64(100), resolved[0].GetCommitTs()) } @@ -397,5 +383,4 @@ func TestAppendRow2GroupKeepsDebeziumPartitionTableFallback(t *testing.T) { require.Equal(t, uint64(100), resolved[0].GetCommitTs()) }) } ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } diff --git a/cmd/pulsar-consumer/consumer.go b/cmd/pulsar-consumer/consumer.go index 62f2feba91..2ae884ae5d 100644 --- a/cmd/pulsar-consumer/consumer.go +++ b/cmd/pulsar-consumer/consumer.go @@ -115,11 +115,7 @@ func (c *consumer) readMessage(ctx context.Context) error { if !needCommit { continue } -<<<<<<< HEAD - err := c.pulsarConsumer.AckID(consumerMsg.Message.ID()) -======= err := c.pulsarConsumer.AckIDCumulative(consumerMsg.ID()) ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) 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 43d3429309..1210dd62a8 100644 --- a/cmd/pulsar-consumer/writer.go +++ b/cmd/pulsar-consumer/writer.go @@ -161,30 +161,12 @@ func (w *writer) flushDDLEvent(ctx context.Context, ddl *commonEvent.DDLEvent) e if !ok { continue } -<<<<<<< HEAD - before := len(resolvedEvents) - resolvedEvents = g.ResolveInto(commitTs, resolvedEvents) - resolvedCount := len(resolvedEvents) - before - if resolvedCount == 0 { - continue - } - - resolvedGroups = append(resolvedGroups, struct { - group *util.EventsGroup - maxCommitTs uint64 - }{ - group: g, - maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), - }) - total += resolvedCount -======= messages := g.ResolveInto(commitTs, nil) events := make([]*commonEvent.DMLEvent, 0, len(messages)) for _, message := range messages { events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) } resolvedEvents = append(resolvedEvents, events...) ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } } @@ -302,30 +284,12 @@ func (w *writer) flushDMLEventsByWatermark(ctx context.Context) error { }, 0) for _, p := range w.progresses { for _, group := range p.eventsGroup { -<<<<<<< HEAD - before := len(resolvedEvents) - resolvedEvents = group.ResolveInto(watermark, resolvedEvents) - resolvedCount := len(resolvedEvents) - before - if resolvedCount == 0 { - continue - } - - resolvedGroups = append(resolvedGroups, struct { - group *util.EventsGroup - maxCommitTs uint64 - }{ - group: group, - maxCommitTs: resolvedEvents[len(resolvedEvents)-1].GetCommitTs(), - }) - total += resolvedCount -======= messages := group.ResolveInto(watermark, nil) events := make([]*commonEvent.DMLEvent, 0, len(messages)) for _, message := range messages { events = util.AppendOrMergeDMLEvent(events, message.ToDMLEvent()) } resolvedEvents = append(resolvedEvents, events...) ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } } total := len(resolvedEvents) @@ -501,23 +465,36 @@ 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) } } -<<<<<<< HEAD -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 @@ -537,7 +514,6 @@ func (w *writer) addPartitionTable(schema, table string) { } func (w *writer) appendMessage2Group(message *common.DMLMessage, progress *partitionProgress) { ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) var ( tableID = message.TableID schema = message.Schema @@ -558,34 +534,13 @@ func (w *writer) appendMessage2Group(message *common.DMLMessage, progress *parti zap.String("schema", schema), zap.String("table", table), zap.Any("protocol", w.protocol)) return } -<<<<<<< HEAD - 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 commitTs >= group.HighWatermark { group.AppendMessage(message, false) log.Debug("DML event append to the group", ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) 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), -<<<<<<< HEAD - zap.Stringer("eventType", dml.RowTypes[0]), zap.Any("protocol", w.protocol), - zap.Bool("IsPartition", dml.TableInfo.TableName.IsPartition)) - group.Append(dml, 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])) -======= zap.Stringer("eventType", message.RowType)) return } @@ -620,5 +575,4 @@ func (w *writer) appendMessage2Group(message *common.DMLMessage, progress *parti default: log.Panic("unknown protocol", zap.Any("protocol", w.protocol)) } ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } diff --git a/cmd/pulsar-consumer/writer_test.go b/cmd/pulsar-consumer/writer_test.go index 0ee23c2b22..39a302c8e1 100644 --- a/cmd/pulsar-consumer/writer_test.go +++ b/cmd/pulsar-consumer/writer_test.go @@ -18,13 +18,11 @@ import ( "testing" "time" -<<<<<<< HEAD -======= "github.com/apache/pulsar-client-go/pulsar" "github.com/golang/mock/gomock" ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) "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" @@ -287,15 +285,7 @@ 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. +func TestOnDDLMarksRoutedCreateTableLikePartitionTable(t *testing.T) { w := &writer{ progresses: []*partitionProgress{ { @@ -307,21 +297,6 @@ func TestAppendRow2Group_DoesNotDropCommitTsFallbackBeforeApplied(t *testing.T) partitionTableAccessor: codeccommon.NewPartitionTableAccessor(), } -<<<<<<< HEAD - 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), - TableInfo: &common.TableInfo{ - TableName: common.TableName{Schema: "test", Table: "t"}, - }, - } - } - - progress := w.progresses[0] -======= ddl := &commonEvent.DDLEvent{ Query: "CREATE TABLE `target`.`dst` LIKE `target`.`src`", SchemaName: "source", @@ -347,29 +322,9 @@ func TestAppendRow2Group_DoesNotDropCommitTsFallbackBeforeApplied(t *testing.T) progress := w.progresses[0] w.appendMessage2Group(newDMLMessage(200), progress) w.appendMessage2Group(newDMLMessage(100), progress) ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) - - // Step 1: observe a larger commitTs first (e.g. produced before restart). - w.appendRow2Group(newDMLEvent(1, 200), progress) - - // Step 2: observe a smaller commitTs later (e.g. replayed after restart). - w.appendRow2Group(newDMLEvent(1, 100), progress) - - group := progress.eventsGroup[1] - require.NotNil(t, group) - // Expect: commitTs=100 is still kept and can be resolved. - resolved := group.ResolveInto(150, nil) + resolved := progress.eventsGroup[1].ResolveInto(150, nil) require.Len(t, resolved, 1) -<<<<<<< HEAD - require.Equal(t, uint64(100), resolved[0].CommitTs) - - // 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) -======= require.Equal(t, uint64(100), resolved[0].GetCommitTs()) } @@ -527,5 +482,4 @@ func (m fakePulsarMessage) Index() *uint64 { func (m fakePulsarMessage) BrokerPublishTime() *time.Time { return nil ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } diff --git a/cmd/storage-consumer/consumer.go b/cmd/storage-consumer/consumer.go index 12cc90b335..57fec1cb47 100644 --- a/cmd/storage-consumer/consumer.go +++ b/cmd/storage-consumer/consumer.go @@ -252,13 +252,8 @@ func (c *consumer) appendMessage2Group(message *common.DMLMessage, enableTableAc c.eventsGroup[tableID] = group } if commitTs >= group.HighWatermark { -<<<<<<< HEAD - group.Append(dml, false) - log.Info("DML event append to the group", -======= group.AppendMessage(message, false) log.Debug("DML event append to the group", ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) 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)) diff --git a/cmd/util/event_group.go b/cmd/util/event_group.go index d3bf980b5b..31b3c740d5 100644 --- a/cmd/util/event_group.go +++ b/cmd/util/event_group.go @@ -17,7 +17,6 @@ import ( "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" @@ -28,16 +27,7 @@ type EventsGroup struct { Partition int32 tableID int64 -<<<<<<< HEAD - 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 ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) HighWatermark uint64 // AppliedWatermark is the maximum CommitTs that has been successfully flushed // to the downstream for this group. @@ -127,51 +117,11 @@ func AppendOrMergeDMLEvent(events []*commonEvent.DMLEvent, row *commonEvent.DMLE 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() { -<<<<<<< HEAD - mergeDMLEvent(lastDMLEvent, row) - 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 - }) - 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) - return - } -======= lastDMLEvent.Rows.Append(row.Rows, 0, row.Rows.NumRows()) lastDMLEvent.RowTypes = append(lastDMLEvent.RowTypes, row.RowTypes...) lastDMLEvent.Length += row.Length @@ -179,54 +129,8 @@ func AppendOrMergeDMLEvent(events []*commonEvent.DMLEvent, row *commonEvent.DMLE return events } ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) log.Panic("append event with smaller commit ts", zap.Int64("tableID", row.GetTableID()), zap.Uint64("lastCommitTs", lastDMLEvent.GetCommitTs()), zap.Uint64("commitTs", row.GetCommitTs())) -<<<<<<< HEAD -} - -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 - }) - 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 { - 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)) - } - return dst -} - -// GetAllEvents will get all events. -func (g *EventsGroup) GetAllEvents() []*commonEvent.DMLEvent { - result := g.events - g.events = nil - return result -======= return events ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } diff --git a/cmd/util/event_group_test.go b/cmd/util/event_group_test.go index 2641236bd9..a8c5fd83ab 100644 --- a/cmd/util/event_group_test.go +++ b/cmd/util/event_group_test.go @@ -18,56 +18,11 @@ import ( "github.com/pingcap/ticdc/pkg/common" commonEvent "github.com/pingcap/ticdc/pkg/common/event" -<<<<<<< HEAD - timodel "github.com/pingcap/tidb/pkg/meta/model" - parser_model "github.com/pingcap/tidb/pkg/parser/model" -======= codeccommon "github.com/pingcap/ticdc/pkg/sink/codec/common" ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) "github.com/pingcap/tidb/pkg/util/chunk" "github.com/stretchr/testify/require" ) -<<<<<<< HEAD -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) - - 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: parser_model.NewCIStr("t"), - Columns: []*timodel.ColumnInfo{ - {Name: parser_model.NewCIStr("a")}, - }, - }), - } - } - - 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 newTestDMLMessage(commitTs uint64) *codeccommon.DMLMessage { return codeccommon.NewDMLMessage(1, "test", "t", commitTs, common.RowTypeInsert, nil) } @@ -80,7 +35,6 @@ func newTestDMLEvent(commitTs uint64, rowTypes ...common.RowType) *commonEvent.D RowTypes: rowTypes, Rows: chunk.NewChunkWithCapacity(nil, 0), } ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) } func TestEventsGroupResolveIntoAppendsAndClearsResolvedPrefix(t *testing.T) { diff --git a/pkg/sink/codec/canal/canal_json_decoder.go b/pkg/sink/codec/canal/canal_json_decoder.go index 52179bc0ad..7e92f43e8e 100644 --- a/pkg/sink/codec/canal/canal_json_decoder.go +++ b/pkg/sink/codec/canal/canal_json_decoder.go @@ -414,13 +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) -<<<<<<< HEAD - // 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}) -======= d.addDDLCommitTs(result.SchemaName, result.TableName, result.GetCommitTs()) d.addDDLCommitTs(result.ExtraSchemaName, result.ExtraTableName, result.GetCommitTs()) ->>>>>>> 5573f0194 (consumer: use dml message instead of dml event (#5590)) return result } From 3b5a2db6a576a38ca8dc20958ba70ddade4c9997 Mon Sep 17 00:00:00 2001 From: wk989898 Date: Mon, 3 Aug 2026 05:22:35 +0000 Subject: [PATCH 3/3] update Signed-off-by: wk989898 --- cmd/kafka-consumer/writer.go | 4 +- cmd/kafka-consumer/writer_test.go | 1 - pkg/sink/codec/debezium/avro_decoder.go | 731 ------------------------ pkg/sink/codec/debezium/avro_test.go | 682 ---------------------- 4 files changed, 2 insertions(+), 1416 deletions(-) delete mode 100644 pkg/sink/codec/debezium/avro_decoder.go delete mode 100644 pkg/sink/codec/debezium/avro_test.go diff --git a/cmd/kafka-consumer/writer.go b/cmd/kafka-consumer/writer.go index c8a3783c43..0cf1356254 100644 --- a/cmd/kafka-consumer/writer.go +++ b/cmd/kafka-consumer/writer.go @@ -532,7 +532,7 @@ func (w *writer) onDDL(ddl *event.DDLEvent) { } switch w.protocol { case config.ProtocolCanalJSON, config.ProtocolOpen, config.ProtocolAvro, config.ProtocolSimple, - config.ProtocolDebezium, config.ProtocolDebeziumAvro: + config.ProtocolDebezium: default: return } @@ -687,7 +687,7 @@ func (w *writer) appendMessage2Group(message *common.DMLMessage, progress *parti // 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: + config.ProtocolDebezium: // 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) { diff --git a/cmd/kafka-consumer/writer_test.go b/cmd/kafka-consumer/writer_test.go index 3776eb135b..616b450b26 100644 --- a/cmd/kafka-consumer/writer_test.go +++ b/cmd/kafka-consumer/writer_test.go @@ -338,7 +338,6 @@ func TestOnDDLMarksRoutedCreateTableLikePartitionTableForAvro(t *testing.T) { 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() diff --git a/pkg/sink/codec/debezium/avro_decoder.go b/pkg/sink/codec/debezium/avro_decoder.go deleted file mode 100644 index 3e735745b1..0000000000 --- a/pkg/sink/codec/debezium/avro_decoder.go +++ /dev/null @@ -1,731 +0,0 @@ -// 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 debezium - -import ( - "context" - "database/sql" - "encoding/binary" - "encoding/json" - "io" - "math/big" - "net/http" - "strconv" - "strings" - "sync" - - "github.com/linkedin/goavro/v2" - "github.com/pingcap/log" - commonEvent "github.com/pingcap/ticdc/pkg/common/event" - "github.com/pingcap/ticdc/pkg/errors" - codecavro "github.com/pingcap/ticdc/pkg/sink/codec/avro" - "github.com/pingcap/ticdc/pkg/sink/codec/common" - "go.uber.org/zap" -) - -const confluentAvroHeaderLen = 5 - -type avroDecoder struct { - ctx context.Context - registryURL string - httpClient *http.Client - inner *decoder - - mu sync.RWMutex - schemas map[int]*registeredDebeziumAvroSchema -} - -type registeredDebeziumAvroSchema struct { - schema any - namedSchemas map[string]any - codec *goavro.Codec -} - -// NewAvroDecoder returns a Debezium decoder for Confluent Avro wire-format -// messages. It decodes the Avro payload and then delegates Debezium event -// semantics to the JSON decoder. -func NewAvroDecoder( - ctx context.Context, - config *common.Config, - idx int, - db *sql.DB, -) (common.Decoder, error) { - registryURL := strings.TrimRight(config.AvroConfluentSchemaRegistry, "/") - if registryURL == "" { - return nil, errors.ErrAvroSchemaAPIError.GenWithStackByArgs("schema registry URI is empty") - } - - return &avroDecoder{ - ctx: ctx, - registryURL: registryURL, - httpClient: http.DefaultClient, - inner: NewDecoder(config, idx, db).(*decoder), - schemas: make(map[int]*registeredDebeziumAvroSchema), - }, nil -} - -func (d *avroDecoder) AddKeyValue(key, value []byte) { - keyJSON, err := d.toDebeziumJSON(key) - if err != nil { - log.Panic("decode Debezium Avro key failed", zap.Error(err), zap.Int("keySize", len(key))) - } - valueJSON, err := d.toDebeziumJSON(value) - if err != nil { - log.Panic("decode Debezium Avro value failed", zap.Error(err), zap.Int("valueSize", len(value))) - } - d.inner.AddKeyValue(keyJSON, valueJSON) -} - -func (d *avroDecoder) HasNext() (common.MessageType, bool) { - return d.inner.HasNext() -} - -func (d *avroDecoder) NextResolvedEvent() uint64 { - return d.inner.NextResolvedEvent() -} - -func (d *avroDecoder) NextDMLMessage() *common.DMLMessage { - return d.inner.NextDMLMessage() -} - -func (d *avroDecoder) NextDDLEvent() *commonEvent.DDLEvent { - return d.inner.NextDDLEvent() -} - -func (d *avroDecoder) toDebeziumJSON(data []byte) ([]byte, error) { - payload, schema, err := d.decodeConfluentAvroMessage(data) - if err != nil { - return nil, err - } - message := map[string]any{ - "schema": schema, - "payload": payload, - } - result, err := json.Marshal(message) - if err != nil { - return nil, errors.WrapError(errors.ErrDebeziumInvalidMessage, err) - } - return result, nil -} - -func (d *avroDecoder) decodeConfluentAvroMessage(data []byte) (any, map[string]any, error) { - if len(data) == 0 { - return nil, nil, errors.ErrDebeziumEmptyValueMessage.GenWithStackByArgs() - } - if len(data) < confluentAvroHeaderLen { - return nil, nil, errors.ErrAvroInvalidMessage.GenWithStackByArgs("confluent header is too short") - } - if data[0] != 0 { - return nil, nil, errors.ErrAvroInvalidMessage.GenWithStackByArgs("invalid confluent magic byte") - } - - schemaID := int(binary.BigEndian.Uint32(data[1:confluentAvroHeaderLen])) - registeredSchema, err := d.getSchema(schemaID) - if err != nil { - return nil, nil, err - } - - native, _, err := registeredSchema.codec.NativeFromBinary(data[confluentAvroHeaderLen:]) - if err != nil { - return nil, nil, errors.WrapError(errors.ErrAvroInvalidMessage, err) - } - - payload, err := avroNativeToConnectPayload( - registeredSchema.schema, - native, - registeredSchema.namedSchemas, - ) - if err != nil { - return nil, nil, err - } - schema, err := avroSchemaToConnectSchema( - registeredSchema.schema, - "", - nil, - registeredSchema.namedSchemas, - ) - if err != nil { - return nil, nil, err - } - return payload, schema, nil -} - -func (d *avroDecoder) getSchema(schemaID int) (*registeredDebeziumAvroSchema, error) { - d.mu.RLock() - schema, ok := d.schemas[schemaID] - d.mu.RUnlock() - if ok { - return schema, nil - } - - uri := d.registryURL + "/schemas/ids/" + strconv.Itoa(schemaID) - req, err := http.NewRequestWithContext(d.ctx, http.MethodGet, uri, nil) - if err != nil { - return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) - } - req.Header.Add( - "Accept", - "application/vnd.schemaregistry.v1+json, application/vnd.schemaregistry+json, "+ - "application/json", - ) - - resp, err := d.httpClient.Do(req) - if err != nil { - return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) - } - defer func() { - _ = resp.Body.Close() - }() - - body, err := io.ReadAll(resp.Body) - if err != nil { - return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) - } - if resp.StatusCode != http.StatusOK { - return nil, errors.ErrAvroSchemaAPIError.GenWithStackByArgs( - "failed to query schema id " + strconv.Itoa(schemaID)) - } - - var lookupResp struct { - Schema string `json:"schema"` - } - if err := json.Unmarshal(body, &lookupResp); err != nil { - return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) - } - - codec, err := codecavro.GenCodec(lookupResp.Schema) - if err != nil { - return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) - } - - decoder := json.NewDecoder(strings.NewReader(lookupResp.Schema)) - decoder.UseNumber() - var schemaDef any - if err := decoder.Decode(&schemaDef); err != nil { - return nil, errors.WrapError(errors.ErrAvroSchemaAPIError, err) - } - - namedSchemas := make(map[string]any) - collectAvroNamedSchemas(schemaDef, namedSchemas) - - schema = ®isteredDebeziumAvroSchema{ - schema: schemaDef, - namedSchemas: namedSchemas, - codec: codec, - } - d.mu.Lock() - d.schemas[schemaID] = schema - d.mu.Unlock() - return schema, nil -} - -func avroNativeToConnectPayload(schema any, value any, namedSchemas map[string]any) (any, error) { - switch typedSchema := schema.(type) { - case []any: - if value == nil { - return nil, nil - } - branchSchema, branchValue, err := avroUnionBranch(typedSchema, value) - if err != nil { - return nil, err - } - return avroNativeToConnectPayload(branchSchema, branchValue, namedSchemas) - case map[string]any: - rawType, ok := typedSchema["type"] - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema is missing type") - } - if unionType, ok := rawType.([]any); ok { - return avroNativeToConnectPayload(unionType, value, namedSchemas) - } - typeName, ok := rawType.(string) - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema type is invalid") - } - switch typeName { - case "record": - valueMap, ok := value.(map[string]any) - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro record payload is invalid") - } - fields, ok := typedSchema["fields"].([]any) - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro record schema is missing fields") - } - result := make(map[string]any, len(fields)) - for _, rawField := range fields { - field, ok := rawField.(map[string]any) - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro field schema is invalid") - } - avroFieldName, ok := field["name"].(string) - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro field is missing name") - } - connectFieldName := avroConnectFieldName(field, avroFieldName) - rawValue, exists := valueMap[avroFieldName] - if !exists && connectFieldName != avroFieldName { - rawValue, exists = valueMap[connectFieldName] - } - if !exists { - rawValue, exists = avroMissingFieldValue(field) - } - if !exists { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs( - "avro record payload is missing field " + avroFieldName) - } - fieldValue, err := avroNativeToConnectPayload( - field["type"], - rawValue, - namedSchemas, - ) - if err != nil { - return nil, err - } - result[connectFieldName] = fieldValue - } - return result, nil - case "array": - items, ok := typedSchema["items"] - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro array schema is missing items") - } - if value == nil { - return []any{}, nil - } - values, ok := value.([]any) - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro array payload is invalid") - } - result := make([]any, 0, len(values)) - for _, item := range values { - itemValue, err := avroNativeToConnectPayload(items, item, namedSchemas) - if err != nil { - return nil, err - } - result = append(result, itemValue) - } - return result, nil - case "bytes": - if avroSchemaIsDecimal(typedSchema) { - return avroDecimalNativeToString(typedSchema, value) - } - return value, nil - default: - return value, nil - } - case string: - if namedSchema, ok := namedSchemas[typedSchema]; ok { - return avroNativeToConnectPayload(namedSchema, value, namedSchemas) - } - return value, nil - default: - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema is invalid") - } -} - -func avroSchemaToConnectSchema( - schema any, - fieldName string, - fieldMeta map[string]any, - namedSchemas map[string]any, -) (map[string]any, error) { - switch typedSchema := schema.(type) { - case []any: - branchSchema, _, err := avroNonNullUnionBranch(typedSchema) - if err != nil { - return nil, err - } - connectSchema, err := avroSchemaToConnectSchema( - branchSchema, - fieldName, - fieldMeta, - namedSchemas, - ) - if err != nil { - return nil, err - } - connectSchema["optional"] = true - return connectSchema, nil - case map[string]any: - rawType, ok := typedSchema["type"] - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema is missing type") - } - if unionType, ok := rawType.([]any); ok { - return avroSchemaToConnectSchema(unionType, fieldName, fieldMeta, namedSchemas) - } - typeName, ok := rawType.(string) - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema type is invalid") - } - switch typeName { - case "record": - connectSchema := newConnectSchema("struct", false, fieldName, typedSchema, fieldMeta) - fields, ok := typedSchema["fields"].([]any) - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro record schema is missing fields") - } - connectFields := make([]any, 0, len(fields)) - for _, rawField := range fields { - field, ok := rawField.(map[string]any) - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro field schema is invalid") - } - avroFieldName, ok := field["name"].(string) - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro field is missing name") - } - fieldSchema, err := avroSchemaToConnectSchema( - field["type"], - avroConnectFieldName(field, avroFieldName), - field, - namedSchemas, - ) - if err != nil { - return nil, err - } - connectFields = append(connectFields, fieldSchema) - } - connectSchema["fields"] = connectFields - return connectSchema, nil - case "array": - connectSchema := newConnectSchema("array", false, fieldName, typedSchema, fieldMeta) - items, ok := typedSchema["items"] - if !ok { - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro array schema is missing items") - } - connectItems, err := avroSchemaToConnectSchema(items, "", nil, namedSchemas) - if err != nil { - return nil, err - } - connectSchema["items"] = connectItems - return connectSchema, nil - default: - connectType, err := avroPrimitiveToConnectType(typeName, typedSchema) - if err != nil { - return nil, err - } - return newConnectSchema(connectType, false, fieldName, typedSchema, fieldMeta), nil - } - case string: - if namedSchema, ok := namedSchemas[typedSchema]; ok { - return avroSchemaToConnectSchema(namedSchema, fieldName, fieldMeta, namedSchemas) - } - connectType, err := avroPrimitiveToConnectType(typedSchema, nil) - if err != nil { - return nil, err - } - return newConnectSchema(connectType, false, fieldName, nil, fieldMeta), nil - default: - return nil, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro schema is invalid") - } -} - -func collectAvroNamedSchemas(schema any, namedSchemas map[string]any) { - switch typedSchema := schema.(type) { - case []any: - for _, branch := range typedSchema { - collectAvroNamedSchemas(branch, namedSchemas) - } - case map[string]any: - rawType := typedSchema["type"] - if unionType, ok := rawType.([]any); ok { - collectAvroNamedSchemas(unionType, namedSchemas) - return - } - typeName, _ := rawType.(string) - switch typeName { - case "record": - name := avroBranchName(typedSchema) - if name != "" { - namedSchemas[name] = typedSchema - shortName := avroShortBranchName(name) - if _, exists := namedSchemas[shortName]; shortName != "" && !exists { - namedSchemas[shortName] = typedSchema - } - } - fields, _ := typedSchema["fields"].([]any) - for _, rawField := range fields { - field, ok := rawField.(map[string]any) - if !ok { - continue - } - collectAvroNamedSchemas(field["type"], namedSchemas) - } - case "array": - collectAvroNamedSchemas(typedSchema["items"], namedSchemas) - } - } -} - -func newConnectSchema( - connectType string, - optional bool, - fieldName string, - schemaMeta map[string]any, - fieldMeta map[string]any, -) map[string]any { - connectSchema := map[string]any{ - "type": connectType, - "optional": optional, - } - if fieldName != "" { - connectSchema["field"] = fieldName - } - addConnectSchemaMetadata(connectSchema, schemaMeta) - addConnectFieldMetadata(connectSchema, fieldMeta) - return connectSchema -} - -func addConnectSchemaMetadata(connectSchema map[string]any, schemaMeta map[string]any) { - if schemaMeta == nil { - return - } - if name, ok := schemaMeta["connect.name"].(string); ok && name != "" { - connectSchema["name"] = name - } - if version, ok := schemaMeta["connect.version"]; ok { - connectSchema["version"] = version - } - if parameters, ok := schemaMeta["connect.parameters"].(map[string]any); ok { - connectSchema["parameters"] = parameters - } - if tidbType, ok := schemaMeta[debeziumAvroTiDBTypeKey].(string); ok && tidbType != "" { - connectSchema[debeziumAvroTiDBTypeKey] = tidbType - } -} - -func addConnectFieldMetadata(connectSchema map[string]any, fieldMeta map[string]any) { - if fieldMeta == nil { - return - } - if tidbType, ok := fieldMeta[debeziumAvroTiDBTypeKey].(string); ok && tidbType != "" { - connectSchema[debeziumAvroTiDBTypeKey] = tidbType - } -} - -func avroPrimitiveToConnectType(avroType string, schemaMeta map[string]any) (string, error) { - if schemaMeta != nil { - if connectType, ok := schemaMeta["connect.type"].(string); ok && connectType != "" { - return connectType, nil - } - } - switch avroType { - case "boolean": - return "boolean", nil - case "string": - return "string", nil - case "bytes": - return "bytes", nil - case "int": - return "int32", nil - case "long": - return "int64", nil - case "float": - return "float", nil - case "double": - return "double", nil - default: - return "", errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("unsupported avro type " + avroType) - } -} - -func avroSchemaIsDecimal(schema map[string]any) bool { - typeName, _ := schema["type"].(string) - logicalType, _ := schema["logicalType"].(string) - return typeName == "bytes" && logicalType == "decimal" -} - -func avroDecimalNativeToString(schema map[string]any, value any) (string, error) { - scale, err := avroDecimalScale(schema) - if err != nil { - return "", err - } - switch v := value.(type) { - case *big.Rat: - return v.FloatString(scale), nil - case big.Rat: - return v.FloatString(scale), nil - case string: - return v, nil - default: - return "", errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("decimal payload is invalid") - } -} - -func avroDecimalScale(schema map[string]any) (int, error) { - switch scale := schema["scale"].(type) { - case float64: - return int(scale), nil - case int: - return scale, nil - case int32: - return int(scale), nil - case int64: - return int(scale), nil - case json.Number: - value, err := scale.Int64() - if err != nil { - return 0, errors.WrapError(errors.ErrDebeziumInvalidMessage, err) - } - return int(value), nil - default: - return 0, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("decimal schema is missing scale") - } -} - -func avroUnionBranch(union []any, value any) (any, any, error) { - if value == nil { - return nil, nil, nil - } - var wrappedBranchName string - var wrappedBranchValue any - hasWrappedBranch := false - if branchValueMap, ok := value.(map[string]any); ok && len(branchValueMap) == 1 { - for branchName, branchValue := range branchValueMap { - wrappedBranchName = branchName - wrappedBranchValue = branchValue - hasWrappedBranch = true - for _, branchSchema := range union { - if avroBranchName(branchSchema) == branchName { - return branchSchema, branchValue, nil - } - } - } - } - - branchSchema, isSingleNonNullBranch, err := avroNonNullUnionBranch(union) - if err != nil { - return nil, nil, err - } - if hasWrappedBranch && - isSingleNonNullBranch && - avroShortBranchName(branchSchema) == avroShortBranchName(wrappedBranchName) { - return branchSchema, wrappedBranchValue, nil - } - return branchSchema, value, nil -} - -func avroNonNullUnionBranch(union []any) (any, bool, error) { - var result any - count := 0 - for _, branch := range union { - if avroBranchName(branch) != "null" { - if count == 0 { - result = branch - } - count++ - } - } - if count > 0 { - return result, count == 1, nil - } - return nil, false, errors.ErrDebeziumInvalidMessage.GenWithStackByArgs("avro union has no non-null branch") -} - -func avroBranchName(schema any) string { - switch typedSchema := schema.(type) { - case string: - return typedSchema - case map[string]any: - typeName, _ := typedSchema["type"].(string) - switch typeName { - case "record": - name, _ := typedSchema["name"].(string) - namespace, _ := typedSchema["namespace"].(string) - if namespace != "" && name != "" { - return namespace + "." + name - } - return name - case "array": - return "array" - default: - if avroSchemaIsDecimal(typedSchema) { - return "bytes.decimal" - } - return typeName - } - default: - return "" - } -} - -func avroShortBranchName(schema any) string { - switch typedSchema := schema.(type) { - case string: - if idx := strings.LastIndex(typedSchema, "."); idx >= 0 { - return typedSchema[idx+1:] - } - return typedSchema - case map[string]any: - typeName, _ := typedSchema["type"].(string) - if typeName == "record" { - name, _ := typedSchema["name"].(string) - return name - } - return avroBranchName(schema) - default: - return "" - } -} - -func avroFieldAllowsMissing(field map[string]any) bool { - if _, hasDefault := field["default"]; hasDefault { - return true - } - return avroSchemaAllowsNull(field["type"]) -} - -func avroMissingFieldValue(field map[string]any) (any, bool) { - if avroFieldAllowsMissing(field) { - return nil, true - } - if avroSchemaIsArray(field["type"]) { - return []any{}, true - } - return nil, false -} - -func avroSchemaIsArray(schema any) bool { - switch typedSchema := schema.(type) { - case map[string]any: - typeName, _ := typedSchema["type"].(string) - return typeName == "array" - case string: - return typedSchema == "array" - default: - return false - } -} - -func avroSchemaAllowsNull(schema any) bool { - union, ok := schema.([]any) - if !ok { - return false - } - for _, branch := range union { - if avroBranchName(branch) == "null" { - return true - } - } - return false -} - -func avroConnectFieldName(field map[string]any, fallback string) string { - if fieldName, ok := field[debeziumAvroConnectFieldKey].(string); ok && fieldName != "" { - return fieldName - } - return fallback -} diff --git a/pkg/sink/codec/debezium/avro_test.go b/pkg/sink/codec/debezium/avro_test.go deleted file mode 100644 index a61e2094ed..0000000000 --- a/pkg/sink/codec/debezium/avro_test.go +++ /dev/null @@ -1,682 +0,0 @@ -// 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 debezium - -import ( - "context" - "encoding/binary" - "encoding/json" - "fmt" - "io" - "net/http" - "testing" - "time" - - "github.com/pingcap/ticdc/downstreamadapter/sink/columnselector" - commonEvent "github.com/pingcap/ticdc/pkg/common/event" - "github.com/pingcap/ticdc/pkg/config" - "github.com/pingcap/ticdc/pkg/sink/codec/avro" - "github.com/pingcap/ticdc/pkg/sink/codec/common" - "github.com/stretchr/testify/require" -) - -func TestDebeziumConfluentAvroEncodeRowEvent(t *testing.T) { - ctx := context.Background() - _, err := avro.SetupEncoderAndSchemaRegistry4Testing( - ctx, - common.NewConfig(config.ProtocolAvro), - ) - require.NoError(t, err) - defer avro.TeardownEncoderAndSchemaRegistry4Testing() - - helper := NewSQLTestHelper(t, "foo", ` - create table foo( - id int primary key, - name varchar(16), - bin varbinary(16), - price decimal(10, 4), - ubig bigint unsigned, - v bigint null - )`) - defer helper.Close() - - dmls := helper.helper.DML2Event("test", "foo", - "insert into foo values (1, 'alice', x'010203', 12.3400, 18446744073709551615, null)") - row, ok := dmls.GetNextRow() - require.True(t, ok) - - cfg := common.NewConfig(config.ProtocolDebeziumAvro) - cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" - cfg.AvroBigintUnsignedHandlingMode = common.BigintUnsignedHandlingModeString - cfg.DebeziumDisableSchema = true - cfg.TimeZone = time.UTC - - encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") - require.NoError(t, err) - require.NoError(t, encoder.AppendRowChangedEvent(ctx, "dbserver1.test.foo", &commonEvent.RowEvent{ - TableInfo: helper.tableInfo, - CommitTs: 1, - Event: row, - ColumnSelector: columnselector.NewDefaultColumnSelector(), - Callback: func() {}, - })) - - messages := encoder.Build() - require.Len(t, messages, 1) - require.Equal(t, byte(0), messages[0].Key[0]) - require.Equal(t, byte(0), messages[0].Value[0]) - - key := decodeConfluentAvroForTest(t, messages[0].Key) - require.Equal(t, int32(1), unwrapAvroUnionForTest(t, key["id"], "int")) - - value := decodeConfluentAvroForTest(t, messages[0].Value) - require.Equal(t, "c", value["op"]) - require.Nil(t, value["before"]) - require.NotContains(t, value, "transaction") - require.IsType(t, int64(0), value["ts_ms"]) - - afterUnion, ok := value["after"].(map[string]any) - require.True(t, ok) - after, ok := afterUnion["dbserver1.test.foo"].(map[string]any) - require.True(t, ok) - require.Equal(t, int32(1), unwrapAvroUnionForTest(t, after["id"], "int")) - require.Equal(t, "alice", unwrapAvroUnionForTest(t, after["name"], "string")) - require.Equal(t, []byte{1, 2, 3}, unwrapAvroUnionForTest(t, after["bin"], "bytes")) - require.Equal(t, "18446744073709551615", unwrapAvroUnionForTest(t, after["ubig"], "string")) - require.Nil(t, after["v"]) - - source, ok := value["source"].(map[string]any) - require.True(t, ok) - require.Equal(t, "test", source["db"]) - require.Equal(t, "foo", source["table"]) - require.Nil(t, source["snapshot"]) - require.Nil(t, source["thread"]) - require.Equal(t, "dbserver1", source["name"]) - - valueSchema := decodeConfluentAvroSchemaForTest(t, messages[0].Value) - require.Contains(t, valueSchema, `"name":"fooEnvelope"`) - require.Contains(t, valueSchema, `"name":"foo"`) - require.Contains(t, valueSchema, `"name":"Source"`) - require.Contains(t, valueSchema, `"logicalType":"decimal"`) - require.NotContains(t, valueSchema, `"field":"transaction"`) -} - -func TestDebeziumConfluentAvroSanitizesFullNameAndUnionBranch(t *testing.T) { - ctx := context.Background() - _, err := avro.SetupEncoderAndSchemaRegistry4Testing( - ctx, - common.NewConfig(config.ProtocolAvro), - ) - require.NoError(t, err) - defer avro.TeardownEncoderAndSchemaRegistry4Testing() - - cfg := common.NewConfig(config.ProtocolDebeziumAvro) - cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" - cfg.TimeZone = time.UTC - - eventEncoder, err := NewAvroBatchEncoder(ctx, cfg, "db-server") - require.NoError(t, err) - encoder, ok := eventEncoder.(*BatchEncoder) - require.True(t, ok) - - payload, err := encoder.encodeAvroPayload( - ctx, - "topic-with-invalid-schema-name", - debeziumAvroValueSchemaSuffix, - &debeziumAvroMessage{ - Schema: &debeziumConnectSchema{ - Type: "struct", - Name: "db-server.test-db.foo-barEnvelope", - Fields: []*debeziumConnectSchema{ - { - Type: "struct", - Optional: true, - Name: "db-server.test-db.foo-bar", - Field: "after", - Fields: []*debeziumConnectSchema{ - { - Type: "int32", - Field: "id", - }, - }, - }, - { - Type: "string", - Field: "op", - }, - }, - }, - Payload: map[string]any{ - "after": map[string]any{ - "id": int32(1), - }, - "op": "c", - }, - }, - 1, - ) - require.NoError(t, err) - - schema := decodeConfluentAvroSchemaForTest(t, payload) - require.Contains(t, schema, `"name":"foo_barEnvelope"`) - require.Contains(t, schema, `"namespace":"db_server.test_db"`) - require.Contains(t, schema, `"name":"foo_bar"`) - require.NotContains(t, schema, `"namespace":"db-server.test-db"`) - - value := decodeConfluentAvroForTest(t, payload) - after := unwrapAvroUnionForTest(t, value["after"], "db_server.test_db.foo_bar") - require.Equal(t, map[string]any{"id": int32(1)}, after) -} - -func TestDebeziumConfluentAvroDecodeRowEvent(t *testing.T) { - ctx := context.Background() - _, err := avro.SetupEncoderAndSchemaRegistry4Testing( - ctx, - common.NewConfig(config.ProtocolAvro), - ) - require.NoError(t, err) - defer avro.TeardownEncoderAndSchemaRegistry4Testing() - - helper := NewSQLTestHelper(t, "foo", ` - create table foo( - id int primary key, - name varchar(16), - bin varbinary(16), - price decimal(10, 4), - ubig bigint unsigned, - v bigint null - )`) - defer helper.Close() - - dmls := helper.helper.DML2Event("test", "foo", - "insert into foo values (1, 'alice', x'010203', 12.3400, 18446744073709551615, null)") - row, ok := dmls.GetNextRow() - require.True(t, ok) - - cfg := common.NewConfig(config.ProtocolDebeziumAvro) - cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" - cfg.AvroBigintUnsignedHandlingMode = common.BigintUnsignedHandlingModeString - cfg.EnableTiDBExtension = true - cfg.TimeZone = time.UTC - - commitTs := uint64(123) - encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") - require.NoError(t, err) - require.NoError(t, encoder.AppendRowChangedEvent(ctx, "dbserver1.test.foo", &commonEvent.RowEvent{ - TableInfo: helper.tableInfo, - CommitTs: commitTs, - Event: row, - ColumnSelector: columnselector.NewDefaultColumnSelector(), - Callback: func() {}, - })) - - messages := encoder.Build() - require.Len(t, messages, 1) - - decoder, err := NewAvroDecoder(ctx, cfg, 0, nil) - require.NoError(t, err) - decoder.AddKeyValue(messages[0].Key, messages[0].Value) - - messageType, hasNext := decoder.HasNext() - require.True(t, hasNext) - require.Equal(t, common.MessageTypeRow, messageType) - - decoded := decoder.NextDMLMessage().ToDMLEvent() - require.Equal(t, commitTs, decoded.CommitTs) - require.Equal(t, "test", decoded.TableInfo.GetSchemaName()) - require.Equal(t, "foo", decoded.TableInfo.GetTableName()) - - change, ok := decoded.GetNextRow() - require.True(t, ok) - common.CompareRow(t, row, helper.tableInfo, change, decoded.TableInfo) -} - -func TestDebeziumConfluentAvroDecodeAccountDMLEvents(t *testing.T) { - ctx := context.Background() - _, err := avro.SetupEncoderAndSchemaRegistry4Testing( - ctx, - common.NewConfig(config.ProtocolAvro), - ) - require.NoError(t, err) - defer avro.TeardownEncoderAndSchemaRegistry4Testing() - - helper := NewSQLTestHelper(t, "tp_account", ` - create table tp_account( - id int primary key, - account_id int not null - )`) - defer helper.Close() - - insertDML := helper.helper.DML2Event("test", "tp_account", - "insert into tp_account values (12, 34)") - updateDML, _ := helper.helper.DML2UpdateEvent("test", "tp_account", - "insert into tp_account values (13, 34)", - "update tp_account set account_id = 35 where id = 13") - deleteDML := helper.helper.DML2DeleteEvent("test", "tp_account", - "insert into tp_account values (14, 34)", - "delete from tp_account where id = 14") - - cfg := common.NewConfig(config.ProtocolDebeziumAvro) - cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" - cfg.EnableTiDBExtension = true - cfg.TimeZone = time.UTC - - encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") - require.NoError(t, err) - - rows := make([]commonEvent.RowChange, 0, 3) - for _, dml := range []*commonEvent.DMLEvent{insertDML, updateDML, deleteDML} { - row, ok := dml.GetNextRow() - if !ok { - continue - } - rows = append(rows, row) - require.NoError(t, encoder.AppendRowChangedEvent(ctx, "dbserver1.test.tp_account", &commonEvent.RowEvent{ - TableInfo: helper.tableInfo, - CommitTs: 1, - Event: row, - ColumnSelector: columnselector.NewDefaultColumnSelector(), - Callback: func() {}, - })) - } - require.Len(t, rows, 3) - - messages := encoder.Build() - require.Len(t, messages, 3) - for idx, message := range messages { - decoder, err := NewAvroDecoder(ctx, cfg, 0, nil) - require.NoError(t, err) - decoder.AddKeyValue(message.Key, message.Value) - - messageType, hasNext := decoder.HasNext() - require.True(t, hasNext) - require.Equal(t, common.MessageTypeRow, messageType) - - decoded := decoder.NextDMLMessage().ToDMLEvent() - require.Equal(t, "test", decoded.TableInfo.GetSchemaName()) - require.Equal(t, "tp_account", decoded.TableInfo.GetTableName()) - - change, ok := decoded.GetNextRow() - require.True(t, ok) - common.CompareRow(t, rows[idx], helper.tableInfo, change, decoded.TableInfo) - } -} - -func TestDebeziumConfluentAvroDecodeShortNamedUnionBranch(t *testing.T) { - valueSchema := map[string]any{ - "type": "record", - "name": "Value", - "namespace": "dbserver1.test.tp_account", - "fields": []any{ - map[string]any{ - "name": "id", - "type": "int", - debeziumAvroConnectFieldKey: "id", - }, - map[string]any{ - "name": "account_id", - "type": "int", - debeziumAvroConnectFieldKey: "account_id", - }, - }, - } - namedSchemas := map[string]any{ - "dbserver1.test.tp_account.Value": valueSchema, - } - - payload, err := avroNativeToConnectPayload( - []any{"null", "dbserver1.test.tp_account.Value"}, - map[string]any{ - "Value": map[string]any{ - "id": int32(12), - "account_id": int32(34), - }, - }, - namedSchemas, - ) - require.NoError(t, err) - require.Equal(t, map[string]any{ - "id": int32(12), - "account_id": int32(34), - }, payload) -} - -func TestDebeziumConfluentAvroDecodeFullNamedWrapperForShortUnionBranch(t *testing.T) { - valueSchema := map[string]any{ - "type": "record", - "name": "Value", - "namespace": "default.test.tp_account", - "fields": []any{ - map[string]any{ - "name": "id", - "type": "int", - debeziumAvroConnectFieldKey: "id", - }, - map[string]any{ - "name": "account_id", - "type": "int", - debeziumAvroConnectFieldKey: "account_id", - }, - }, - } - envelopeSchema := map[string]any{ - "type": "record", - "name": "Envelope", - "namespace": "default.test.tp_account", - "fields": []any{ - map[string]any{ - "name": "after", - "type": []any{"null", "Value"}, - }, - }, - } - namedSchemas := map[string]any{} - collectAvroNamedSchemas(valueSchema, namedSchemas) - collectAvroNamedSchemas(envelopeSchema, namedSchemas) - - payload, err := avroNativeToConnectPayload( - []any{"null", "Value"}, - map[string]any{ - "default.test.tp_account.Value": map[string]any{ - "id": int32(12), - "account_id": int32(34), - }, - }, - namedSchemas, - ) - require.NoError(t, err) - require.Equal(t, map[string]any{ - "id": int32(12), - "account_id": int32(34), - }, payload) -} - -func TestDebeziumConfluentAvroDecodeSingleFieldUnionRecord(t *testing.T) { - valueSchema := map[string]any{ - "type": "record", - "name": "Value", - "namespace": "dbserver1.test.only_pk", - "fields": []any{ - map[string]any{ - "name": "id", - "type": "int", - debeziumAvroConnectFieldKey: "id", - }, - }, - } - namedSchemas := map[string]any{ - "dbserver1.test.only_pk.Value": valueSchema, - } - - payload, err := avroNativeToConnectPayload( - []any{"null", "dbserver1.test.only_pk.Value"}, - map[string]any{ - "id": int32(12), - }, - namedSchemas, - ) - require.NoError(t, err) - require.Equal(t, map[string]any{ - "id": int32(12), - }, payload) -} - -func TestDebeziumConfluentAvroDecodeMissingRecordField(t *testing.T) { - valueSchema := map[string]any{ - "type": "record", - "name": "Value", - "namespace": "dbserver1.test.tp_account", - "fields": []any{ - map[string]any{ - "name": "id", - "type": "int", - debeziumAvroConnectFieldKey: "id", - }, - map[string]any{ - "name": "account_id", - "type": "int", - debeziumAvroConnectFieldKey: "account_id", - }, - }, - } - - _, err := avroNativeToConnectPayload( - valueSchema, - map[string]any{ - "id": int32(12), - }, - nil, - ) - require.Error(t, err) - require.Contains(t, err.Error(), "avro record payload is missing field account_id") -} - -func TestDebeziumConfluentAvroDecodeMissingOptionalRecordField(t *testing.T) { - valueSchema := map[string]any{ - "type": "record", - "name": "Table", - "namespace": "io.debezium.connector.schema", - "fields": []any{ - map[string]any{ - "name": "defaultCharsetName", - "type": []any{"null", "string"}, - "default": nil, - }, - map[string]any{ - "name": "columns", - "type": map[string]any{ - "type": "array", - "items": "string", - }, - }, - }, - } - - payload, err := avroNativeToConnectPayload( - valueSchema, - map[string]any{ - "columns": []any{"id"}, - }, - nil, - ) - require.NoError(t, err) - require.Equal(t, map[string]any{ - "defaultCharsetName": nil, - "columns": []any{"id"}, - }, payload) -} - -func TestDebeziumConfluentAvroDecodeMissingArrayRecordField(t *testing.T) { - valueSchema := map[string]any{ - "type": "record", - "name": "Table", - "namespace": "io.debezium.connector.schema", - "fields": []any{ - map[string]any{ - "name": "columns", - "type": map[string]any{ - "type": "array", - "items": "string", - }, - }, - }, - } - - payload, err := avroNativeToConnectPayload( - valueSchema, - map[string]any{}, - nil, - ) - require.NoError(t, err) - require.Equal(t, map[string]any{ - "columns": []any{}, - }, payload) -} - -func TestDebeziumConfluentAvroEncodeDDLEvent(t *testing.T) { - ctx := context.Background() - _, err := avro.SetupEncoderAndSchemaRegistry4Testing( - ctx, - common.NewConfig(config.ProtocolAvro), - ) - require.NoError(t, err) - defer avro.TeardownEncoderAndSchemaRegistry4Testing() - - cfg := common.NewConfig(config.ProtocolDebeziumAvro) - cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" - cfg.EnableTiDBExtension = true - cfg.TimeZone = time.UTC - - routedDDL := common.NewRoutedDDLEvent4Test() - - encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") - require.NoError(t, err) - message, err := encoder.EncodeDDLEvent(routedDDL) - require.NoError(t, err) - require.Nil(t, message) - - cfg.AvroEnableWatermark = true - encoder, err = NewAvroBatchEncoder(ctx, cfg, "dbserver1") - require.NoError(t, err) - message, err = encoder.EncodeDDLEvent(routedDDL) - require.NoError(t, err) - require.NotNil(t, message) - require.Equal(t, byte(0), message.Key[0]) - require.Equal(t, byte(0), message.Value[0]) - - decoder, err := NewAvroDecoder(ctx, cfg, 0, nil) - require.NoError(t, err) - decoder.AddKeyValue(message.Key, message.Value) - - messageType, hasNext := decoder.HasNext() - require.True(t, hasNext) - require.Equal(t, common.MessageTypeDDL, messageType) - - decoded := decoder.NextDDLEvent() - require.Equal(t, routedDDL.GetCommitTs(), decoded.GetCommitTs()) - require.Equal(t, routedDDL.GetDDLType(), decoded.GetDDLType()) - require.Equal(t, "target_db", decoded.SchemaName) - require.Equal(t, "target_table", decoded.TableName) - require.Equal(t, routedDDL.Query, decoded.Query) -} - -func TestDebeziumConfluentAvroEncodeCheckpointEvent(t *testing.T) { - ctx := context.Background() - _, err := avro.SetupEncoderAndSchemaRegistry4Testing( - ctx, - common.NewConfig(config.ProtocolAvro), - ) - require.NoError(t, err) - defer avro.TeardownEncoderAndSchemaRegistry4Testing() - - cfg := common.NewConfig(config.ProtocolDebeziumAvro) - cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" - cfg.EnableTiDBExtension = true - cfg.AvroEnableWatermark = true - cfg.TimeZone = time.UTC - - encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") - require.NoError(t, err) - - message, err := encoder.EncodeCheckpointEvent(100) - require.NoError(t, err) - require.NotNil(t, message) - - decoder, err := NewAvroDecoder(ctx, cfg, 0, nil) - require.NoError(t, err) - decoder.AddKeyValue(message.Key, message.Value) - - messageType, hasNext := decoder.HasNext() - require.True(t, hasNext) - require.Equal(t, common.MessageTypeResolved, messageType) - require.Equal(t, uint64(100), decoder.NextResolvedEvent()) -} - -func TestDebeziumConfluentAvroDoesNotEncodeCheckpointEventByDefault(t *testing.T) { - ctx := context.Background() - _, err := avro.SetupEncoderAndSchemaRegistry4Testing( - ctx, - common.NewConfig(config.ProtocolAvro), - ) - require.NoError(t, err) - defer avro.TeardownEncoderAndSchemaRegistry4Testing() - - cfg := common.NewConfig(config.ProtocolDebeziumAvro) - cfg.AvroConfluentSchemaRegistry = "http://127.0.0.1:8081" - cfg.EnableTiDBExtension = true - cfg.TimeZone = time.UTC - - encoder, err := NewAvroBatchEncoder(ctx, cfg, "dbserver1") - require.NoError(t, err) - - message, err := encoder.EncodeCheckpointEvent(100) - require.NoError(t, err) - require.Nil(t, message) -} - -func decodeConfluentAvroForTest(t *testing.T, data []byte) map[string]any { - t.Helper() - - schema, binaryData := decodeConfluentAvroEnvelopeForTest(t, data) - codec, err := avro.GenCodec(schema) - require.NoError(t, err) - - native, _, err := codec.NativeFromBinary(binaryData) - require.NoError(t, err) - - result, ok := native.(map[string]any) - require.True(t, ok) - return result -} - -func decodeConfluentAvroSchemaForTest(t *testing.T, data []byte) string { - t.Helper() - - schema, _ := decodeConfluentAvroEnvelopeForTest(t, data) - return schema -} - -func decodeConfluentAvroEnvelopeForTest(t *testing.T, data []byte) (string, []byte) { - t.Helper() - - require.GreaterOrEqual(t, len(data), 5) - require.Equal(t, byte(0), data[0]) - schemaID := int(binary.BigEndian.Uint32(data[1:5])) - binaryData := data[5:] - - resp, err := http.Get(fmt.Sprintf("http://127.0.0.1:8081/schemas/ids/%d", schemaID)) - require.NoError(t, err) - defer resp.Body.Close() - require.Equal(t, http.StatusOK, resp.StatusCode) - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - var schemaResp struct { - Schema string `json:"schema"` - } - require.NoError(t, json.Unmarshal(body, &schemaResp)) - - return schemaResp.Schema, binaryData -} - -func unwrapAvroUnionForTest(t *testing.T, value any, branch string) any { - t.Helper() - - union, ok := value.(map[string]any) - require.True(t, ok) - result, ok := union[branch] - require.True(t, ok) - return result -}