Files
milvus/internal/streamingcoord/server/broadcaster/broadcaster_test.go
T
5eebaa9ad4 enhance: add WAL trace propagation (#50796)
issue: #47404

- message trace context: add trace context serialization and
restore/inject helpers for WAL messages and msgstream conversion
- WAL append trace: normalize WAL spans for autocommit, txn, broadcast,
append, appendimpl, and broadcast callback paths
- trace propagation: restore message trace context in producer,
broadcast retry, replicate primary/secondary, recovery, flusher, and
query/data flowgraph consumers

---------

Signed-off-by: chyezh <chyezh@outlook.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-07-15 02:14:37 +08:00

1208 lines
45 KiB
Go

package broadcaster
import (
"context"
"math/rand"
"sync"
"testing"
"time"
"github.com/cockroachdb/errors"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/require"
"go.uber.org/atomic"
"google.golang.org/protobuf/proto"
"github.com/milvus-io/milvus-proto/go-api/v3/commonpb"
"github.com/milvus-io/milvus-proto/go-api/v3/msgpb"
"github.com/milvus-io/milvus-proto/go-api/v3/schemapb"
"github.com/milvus-io/milvus/internal/distributed/streaming"
"github.com/milvus-io/milvus/internal/mocks/distributed/mock_streaming"
"github.com/milvus-io/milvus/internal/mocks/mock_metastore"
"github.com/milvus-io/milvus/internal/mocks/streamingcoord/server/mock_balancer"
"github.com/milvus-io/milvus/internal/streamingcoord/server/balancer"
"github.com/milvus-io/milvus/internal/streamingcoord/server/balancer/balance"
"github.com/milvus-io/milvus/internal/streamingcoord/server/broadcaster/registry"
"github.com/milvus-io/milvus/internal/streamingcoord/server/resource"
internaltypes "github.com/milvus-io/milvus/internal/types"
"github.com/milvus-io/milvus/internal/util/idalloc"
streamingstatus "github.com/milvus-io/milvus/internal/util/streamingutil/status"
"github.com/milvus-io/milvus/pkg/v3/mlog"
"github.com/milvus-io/milvus/pkg/v3/mocks/streaming/util/mock_message"
"github.com/milvus-io/milvus/pkg/v3/proto/messagespb"
"github.com/milvus-io/milvus/pkg/v3/proto/streamingpb"
"github.com/milvus-io/milvus/pkg/v3/streaming/util/message"
"github.com/milvus-io/milvus/pkg/v3/streaming/util/types"
"github.com/milvus-io/milvus/pkg/v3/streaming/walimpls/impls/walimplstest"
"github.com/milvus-io/milvus/pkg/v3/util/paramtable"
"github.com/milvus-io/milvus/pkg/v3/util/replicateutil"
"github.com/milvus-io/milvus/pkg/v3/util/syncutil"
"github.com/milvus-io/milvus/pkg/v3/util/typeutil"
)
func TestBroadcaster(t *testing.T) {
registry.ResetRegistration()
paramtable.Init()
paramtable.Get().StreamingCfg.WALBroadcasterTombstoneCheckInternal.SwapTempValue("10ms")
paramtable.Get().StreamingCfg.WALBroadcasterTombstoneMaxCount.SwapTempValue("2")
paramtable.Get().StreamingCfg.WALBroadcasterTombstoneMaxLifetime.SwapTempValue("20ms")
mb := mock_balancer.NewMockBalancer(t)
mb.EXPECT().ReplicateRole().Return(replicateutil.RolePrimary)
mb.EXPECT().WatchChannelAssignments(mock.Anything, mock.Anything).RunAndReturn(func(ctx context.Context, cb balancer.WatchChannelAssignmentsCallback) error {
<-ctx.Done()
return ctx.Err()
})
balance.Register(mb)
registry.RegisterDropCollectionV1AckCallback(func(ctx context.Context, msg message.BroadcastResultDropCollectionMessageV1) error {
return nil
})
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
meta.EXPECT().ListBroadcastTask(mock.Anything).
RunAndReturn(func(ctx context.Context) ([]*streamingpb.BroadcastTask, error) {
return []*streamingpb.BroadcastTask{
createNewBroadcastTask(8, []string{"v1"}, message.NewCollectionNameResourceKey("c1")),
createNewBroadcastTask(9, []string{"v1", "v2"}, message.NewCollectionNameResourceKey("c2")),
createNewBroadcastTask(3, []string{"v1", "v2", "v3"}),
createNewWaitAckBroadcastTaskFromMessage(
createNewBroadcastMsg([]string{"v1", "v2", "v3"}).WithBroadcastID(4),
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x00, 0x01, 0x00}),
createNewWaitAckBroadcastTaskFromMessage(
createNewBroadcastMsg([]string{"v1", "v2", "v3"}).WithBroadcastID(5),
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01, 0x01, 0x00}),
createNewWaitAckBroadcastTaskFromMessage(
createNewBroadcastMsg([]string{"v1", "v2", "v3"}).WithBroadcastID(6), // will be done directly.
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01, 0x01, 0x01}),
createNewWaitAckBroadcastTaskFromMessage(
createNewBroadcastMsg([]string{"v1", "v2", "v3"},
message.NewCollectionNameResourceKey("c3"),
message.NewCollectionNameResourceKey("c4")).WithBroadcastID(7),
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_REPLICATED,
[]byte{0x00, 0x00, 0x00}),
}, nil
}).Times(1)
done := typeutil.NewConcurrentSet[uint64]()
meta.EXPECT().SaveBroadcastTask(mock.Anything, mock.Anything, mock.Anything).RunAndReturn(func(ctx context.Context, broadcastID uint64, bt *streamingpb.BroadcastTask) error {
if ctx.Err() != nil {
return ctx.Err()
}
if bt.State == streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_TOMBSTONE {
done.Insert(broadcastID)
}
return nil
})
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
fbc := syncutil.NewFuture[Broadcaster]()
appended := createOpeartor(t, fbc)
bc, err := RecoverBroadcaster(context.Background())
fbc.Set(bc)
assert.NoError(t, err)
assert.NotNil(t, bc)
assert.Eventually(t, func() bool {
return appended.Load() == 9 && len(done.Collect()) == 6
}, 30*time.Second, 10*time.Millisecond)
// only task 7 is not done.
ack(t, bc, 7, "v1")
ack(t, bc, 7, "v1") // test already acked, make the idempotent.
assert.Equal(t, len(done.Collect()), 6)
ack(t, bc, 7, "v2")
ack(t, bc, 7, "v2")
assert.Equal(t, len(done.Collect()), 6)
ack(t, bc, 7, "v3")
ack(t, bc, 7, "v3")
assert.Eventually(t, func() bool {
return appended.Load() == 9 && len(done.Collect()) == 7
}, 30*time.Second, 10*time.Millisecond)
// Test broadcast here.
broadcastWithSameRK := func() {
var result *types.BroadcastAppendResult
var err error
b, err := bc.WithResourceKeys(context.Background(), message.NewCollectionNameResourceKey("c7"))
assert.NoError(t, err)
result, err = b.Broadcast(context.Background(), createNewBroadcastMsg([]string{"v1", "v2", "v3"}, message.NewCollectionNameResourceKey("c7")))
assert.Equal(t, len(result.AppendResults), 3)
assert.NoError(t, err)
}
go broadcastWithSameRK()
go broadcastWithSameRK()
assert.Eventually(t, func() bool {
return appended.Load() == 15 && len(done.Collect()) == 9
}, 30*time.Second, 10*time.Millisecond)
// Test close befor broadcast
broadcastAPI, err := bc.WithResourceKeys(context.Background(), message.NewExclusiveClusterResourceKey())
assert.NoError(t, err)
broadcastAPI.Close()
broadcastAPI, err = bc.WithResourceKeys(context.Background(), message.NewExclusiveClusterResourceKey())
assert.NoError(t, err)
broadcastAPI.Close()
bc.Close()
broadcastAPI, err = bc.WithResourceKeys(context.Background())
assert.NoError(t, err)
_, err = broadcastAPI.Broadcast(context.Background(), createNewBroadcastMsg([]string{"v1"}))
assert.Error(t, err)
err = bc.Ack(context.Background(), mock_message.NewMockImmutableMessage(t))
assert.Error(t, err)
}
func ack(t *testing.T, broadcaster Broadcaster, broadcastID uint64, vchannel string) {
for {
msg := message.NewDropCollectionMessageBuilderV1().
WithHeader(&message.DropCollectionMessageHeader{}).
WithBody(&msgpb.DropCollectionRequest{}).
WithBroadcast([]string{vchannel}).
MustBuildBroadcast().
WithBroadcastID(broadcastID).
SplitIntoMutableMessage()[0].
WithTimeTick(100).
WithLastConfirmed(walimplstest.NewTestMessageID(1)).
IntoImmutableMessage(walimplstest.NewTestMessageID(1))
if err := broadcaster.Ack(context.Background(), msg); err == nil {
break
}
}
}
func createOpeartor(t *testing.T, broadcaster *syncutil.Future[Broadcaster]) *atomic.Int64 {
id := atomic.NewInt64(1)
appended := atomic.NewInt64(0)
operator := mock_streaming.NewMockWALAccesser(t)
f := func(ctx context.Context, msgs ...message.MutableMessage) types.AppendResponses {
resps := types.AppendResponses{
Responses: make([]types.AppendResponse, len(msgs)),
}
for idx, msg := range msgs {
newID := walimplstest.NewTestMessageID(id.Inc())
if rand.Int31n(10) < 3 {
resps.Responses[idx] = types.AppendResponse{
Error: errors.New("append failed"),
}
continue
}
resps.Responses[idx] = types.AppendResponse{
AppendResult: &types.AppendResult{
MessageID: newID,
TimeTick: uint64(time.Now().UnixMilli()),
},
Error: nil,
}
appended.Inc()
broadcastID := msg.BroadcastHeader().BroadcastID
vchannel := msg.VChannel()
go func() {
time.Sleep(time.Duration(rand.Int31n(100)) * time.Millisecond)
ack(t, broadcaster.Get(), broadcastID, vchannel)
}()
}
return resps
}
operator.EXPECT().AppendMessages(mock.Anything, mock.Anything).RunAndReturn(f)
operator.EXPECT().AppendMessages(mock.Anything, mock.Anything, mock.Anything).RunAndReturn(f)
operator.EXPECT().AppendMessages(mock.Anything, mock.Anything, mock.Anything, mock.Anything).RunAndReturn(f)
operator.EXPECT().AppendMessages(mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).RunAndReturn(f)
streaming.SetWALForTest(operator)
return appended
}
func createNewBroadcastMsg(vchannels []string, rks ...message.ResourceKey) message.BroadcastMutableMessage {
msg, err := message.NewDropCollectionMessageBuilderV1().
WithHeader(&messagespb.DropCollectionMessageHeader{}).
WithBody(&msgpb.DropCollectionRequest{}).
WithBroadcast(vchannels).
BuildBroadcast()
if err != nil {
panic(err)
}
return msg.OverwriteBroadcastHeader(0, rks...)
}
func TestBroadcastTaskNotCreatedOnStoppedBroadcaster(t *testing.T) {
locker := newResourceKeyLocker()
rk := message.NewExclusiveCollectionNameResourceKey("db", "collection")
guards := locker.Lock(rk)
bm := &broadcastTaskManager{
lifetime: typeutil.NewLifetime(),
mu: &sync.Mutex{},
tasks: map[uint64]*broadcastTask{},
}
bm.lifetime.SetState(typeutil.LifetimeStateStopped)
_, err := bm.broadcast(context.Background(), createNewBroadcastMsg([]string{"v1"}, rk), 1, guards)
require.Error(t, err)
require.True(t, IsBroadcastTaskNotCreated(err))
require.True(t, streamingstatus.AsStreamingError(err).IsOnShutdown())
require.Empty(t, bm.tasks)
nextGuards, lockErr := locker.FastLock(rk)
require.NoError(t, lockErr)
nextGuards.Unlock()
}
func createNewBroadcastTask(broadcastID uint64, vchannels []string, rks ...message.ResourceKey) *streamingpb.BroadcastTask {
msg := createNewBroadcastMsg(vchannels).OverwriteBroadcastHeader(broadcastID, rks...)
pb := msg.IntoMessageProto()
return &streamingpb.BroadcastTask{
Message: &messagespb.Message{
Payload: pb.Payload,
Properties: pb.Properties,
},
State: streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
AckedVchannelBitmap: make([]byte, len(vchannels)),
}
}
func createNewWaitAckBroadcastTaskFromMessage(
msg message.BroadcastMutableMessage,
state streamingpb.BroadcastTaskState,
bitmap []byte,
) *streamingpb.BroadcastTask {
pb := msg.IntoMessageProto()
acks := make([]*streamingpb.AckedCheckpoint, len(bitmap))
for i := 0; i < len(bitmap); i++ {
if bitmap[i] != 0 {
messageID := walimplstest.NewTestMessageID(int64(i))
lastConfirmedMessageID := walimplstest.NewTestMessageID(int64(i))
acks[i] = &streamingpb.AckedCheckpoint{
MessageId: messageID.IntoProto(),
LastConfirmedMessageId: lastConfirmedMessageID.IntoProto(),
TimeTick: 1,
}
}
}
return &streamingpb.BroadcastTask{
Message: &messagespb.Message{
Payload: pb.Payload,
Properties: pb.Properties,
},
State: state,
AckedVchannelBitmap: bitmap,
AckedCheckpoints: acks,
}
}
func TestRecoverBroadcastTaskFromProto(t *testing.T) {
task := createNewBroadcastTask(8, []string{"v1", "v2", "v3"}, message.NewCollectionNameResourceKey("c1"))
b, err := proto.Marshal(task)
require.NoError(t, err)
task = unmarshalTask(t, b, 3)
assert.Equal(t, task.AckedVchannelBitmap, []byte{0x00, 0x00, 0x00})
assert.Len(t, task.AckedCheckpoints, 3)
assert.Nil(t, task.AckedCheckpoints[0])
assert.Nil(t, task.AckedCheckpoints[1])
assert.Nil(t, task.AckedCheckpoints[2])
cp := &streamingpb.AckedCheckpoint{
MessageId: walimplstest.NewTestMessageID(1).IntoProto(),
LastConfirmedMessageId: walimplstest.NewTestMessageID(1).IntoProto(),
TimeTick: 1,
}
task.AckedCheckpoints[2] = cp
task.AckedVchannelBitmap[2] = 0x01
b, err = proto.Marshal(task)
require.NoError(t, err)
task = unmarshalTask(t, b, 3)
assert.Equal(t, task.AckedVchannelBitmap, []byte{0x00, 0x00, 0x01})
assert.Len(t, task.AckedCheckpoints, 3)
assert.Nil(t, task.AckedCheckpoints[0])
assert.Nil(t, task.AckedCheckpoints[1])
assert.NotNil(t, task.AckedCheckpoints[2])
task.AckedCheckpoints[2] = nil
task.AckedVchannelBitmap[2] = 0x0
task.AckedCheckpoints[0] = cp
task.AckedVchannelBitmap[0] = 0x01
b, err = proto.Marshal(task)
require.NoError(t, err)
task = unmarshalTask(t, b, 3)
assert.Equal(t, task.AckedVchannelBitmap, []byte{0x01, 0x00, 0x00})
assert.Len(t, task.AckedCheckpoints, 3)
assert.NotNil(t, task.AckedCheckpoints[0])
assert.Nil(t, task.AckedCheckpoints[1])
assert.Nil(t, task.AckedCheckpoints[2])
task.AckedCheckpoints[0] = nil
task.AckedVchannelBitmap[0] = 0x0
task.AckedCheckpoints[1] = cp
task.AckedVchannelBitmap[1] = 0x01
b, err = proto.Marshal(task)
require.NoError(t, err)
task = unmarshalTask(t, b, 3)
assert.Equal(t, task.AckedVchannelBitmap, []byte{0x00, 0x01, 0x00})
assert.Len(t, task.AckedCheckpoints, 3)
assert.Nil(t, task.AckedCheckpoints[0])
assert.NotNil(t, task.AckedCheckpoints[1])
assert.Nil(t, task.AckedCheckpoints[2])
task.AckedVchannelBitmap = []byte{0x01, 0x01, 0x01}
task.AckedCheckpoints = []*streamingpb.AckedCheckpoint{
cp,
cp,
cp,
}
b, err = proto.Marshal(task)
require.NoError(t, err)
task = unmarshalTask(t, b, 3)
assert.Equal(t, task.AckedVchannelBitmap, []byte{0x01, 0x01, 0x01})
assert.Len(t, task.AckedCheckpoints, 3)
assert.NotNil(t, task.AckedCheckpoints[0])
assert.NotNil(t, task.AckedCheckpoints[1])
assert.NotNil(t, task.AckedCheckpoints[2])
}
func unmarshalTask(t *testing.T, b []byte, vchannelCount int) *streamingpb.BroadcastTask {
task := &streamingpb.BroadcastTask{}
err := proto.Unmarshal(b, task)
require.NoError(t, err)
fixAckInfoFromProto(task, vchannelCount)
return task
}
func TestGetIncompleteBroadcastTasks(t *testing.T) {
paramtable.Init()
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
// Task 1: PENDING state with pending (unacked) messages -> should be returned
pendingProto := createNewBroadcastTask(1, []string{"v1", "v2"})
pendingTask := newBroadcastTaskFromProto(pendingProto, metrics, ackScheduler)
// Task 2: REPLICATED state with pending (unacked) messages -> should be returned
replicatedProto := createNewWaitAckBroadcastTaskFromMessage(
createNewBroadcastMsg([]string{"v1", "v2", "v3"}).WithBroadcastID(2),
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_REPLICATED,
[]byte{0x00, 0x00, 0x00}, // none acked
)
replicatedTask := newBroadcastTaskFromProto(replicatedProto, metrics, ackScheduler)
// Task 3: PENDING state but ALL vchannels acked -> should NOT be returned (no pending messages)
allAckedProto := createNewWaitAckBroadcastTaskFromMessage(
createNewBroadcastMsg([]string{"v1", "v2", "v3"}).WithBroadcastID(3),
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01, 0x01, 0x01}, // all acked
)
allAckedTask := newBroadcastTaskFromProto(allAckedProto, metrics, ackScheduler)
// Task 4: TOMBSTONE state -> should NOT be returned
tombstoneProto := createNewWaitAckBroadcastTaskFromMessage(
createNewBroadcastMsg([]string{"v1", "v2"}).WithBroadcastID(4),
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_TOMBSTONE,
[]byte{0x01, 0x01}, // all acked
)
tombstoneTask := newBroadcastTaskFromProto(tombstoneProto, metrics, ackScheduler)
bm := &broadcastTaskManager{
mu: &sync.Mutex{},
tasks: make(map[uint64]*broadcastTask),
}
bm.tasks[1] = pendingTask
bm.tasks[2] = replicatedTask
bm.tasks[3] = allAckedTask
bm.tasks[4] = tombstoneTask
result := bm.getIncompleteBroadcastTasks()
// Should return exactly 2 tasks: the pending task (ID=1) and the replicated task (ID=2)
assert.Len(t, result, 2)
// Collect the broadcast IDs from the result
resultIDs := make(map[uint64]struct{})
for _, task := range result {
resultIDs[task.Header().BroadcastID] = struct{}{}
}
assert.Contains(t, resultIDs, uint64(1), "PENDING task with pending messages should be returned")
assert.Contains(t, resultIDs, uint64(2), "REPLICATED task with pending messages should be returned")
assert.NotContains(t, resultIDs, uint64(3), "PENDING task with all vchannels acked should not be returned")
assert.NotContains(t, resultIDs, uint64(4), "TOMBSTONE task should not be returned")
}
func TestGetPendingSchemaFileResources(t *testing.T) {
paramtable.Init()
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
createCollectionMsg := func(collectionID int64, fileResourceIDs []int64) message.BroadcastMutableMessage {
return message.NewCreateCollectionMessageBuilderV1().
WithHeader(&message.CreateCollectionMessageHeader{
CollectionId: collectionID,
}).
WithBody(&msgpb.CreateCollectionRequest{
CollectionSchema: &schemapb.CollectionSchema{
FileResourceIds: fileResourceIDs,
},
}).
WithBroadcast([]string{"v1"}).
MustBuildBroadcast()
}
alterCollectionMsg := func(collectionID int64, fileResourceIDs []int64) message.BroadcastMutableMessage {
return message.NewAlterCollectionMessageBuilderV2().
WithHeader(&message.AlterCollectionMessageHeader{
CollectionId: collectionID,
}).
WithBody(&message.AlterCollectionMessageBody{
Updates: &message.AlterCollectionMessageUpdates{
Schema: &schemapb.CollectionSchema{
FileResourceIds: fileResourceIDs,
},
},
}).
WithBroadcast([]string{"v1"}).
MustBuildBroadcast()
}
newTask := func(broadcastID uint64, msg message.BroadcastMutableMessage, state streamingpb.BroadcastTaskState) *broadcastTask {
proto := createNewWaitAckBroadcastTaskFromMessage(msg.WithBroadcastID(broadcastID), state, []byte{0x00})
return newBroadcastTaskFromProto(proto, metrics, ackScheduler)
}
bm := &broadcastTaskManager{
mu: &sync.Mutex{},
tasks: map[uint64]*broadcastTask{
1: newTask(1, createCollectionMsg(100, []int64{10, 20}), streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING),
2: newTask(2, alterCollectionMsg(100, []int64{20, 30}), streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING),
3: newTask(3, alterCollectionMsg(200, nil), streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING),
4: newTask(4, alterCollectionMsg(300, []int64{40}), streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_TOMBSTONE),
5: newTask(5, createNewBroadcastMsg([]string{"v1"}), streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING),
},
}
result := bm.GetPendingSchemaFileResources()
require.Len(t, result, 1)
assert.ElementsMatch(t, []int64{10, 20, 30}, result[100])
}
func TestWithSecondaryClusterResourceKey(t *testing.T) {
t.Run("success", func(t *testing.T) {
registry.ResetRegistration()
paramtable.Init()
balance.ResetBalancer()
mb := mock_balancer.NewMockBalancer(t)
mb.EXPECT().ReplicateRole().Return(replicateutil.RoleSecondary).Maybe()
mb.EXPECT().WatchChannelAssignments(mock.Anything, mock.Anything).RunAndReturn(func(ctx context.Context, cb balancer.WatchChannelAssignmentsCallback) error {
time.Sleep(100 * time.Second)
return nil
}).Maybe()
balance.Register(mb)
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
meta.EXPECT().ListBroadcastTask(mock.Anything).Return([]*streamingpb.BroadcastTask{}, nil).Times(1)
meta.EXPECT().SaveBroadcastTask(mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe()
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
mw := mock_streaming.NewMockWALAccesser(t)
streaming.SetWALForTest(mw)
bc, err := RecoverBroadcaster(context.Background())
assert.NoError(t, err)
// Should succeed on secondary cluster
api, err := bc.WithSecondaryClusterResourceKey(context.Background())
assert.NoError(t, err)
assert.NotNil(t, api)
api.Close()
bc.Close()
})
t.Run("not_secondary", func(t *testing.T) {
registry.ResetRegistration()
paramtable.Init()
balance.ResetBalancer()
mb := mock_balancer.NewMockBalancer(t)
mb.EXPECT().ReplicateRole().Return(replicateutil.RolePrimary).Maybe()
mb.EXPECT().WatchChannelAssignments(mock.Anything, mock.Anything).RunAndReturn(func(ctx context.Context, cb balancer.WatchChannelAssignmentsCallback) error {
time.Sleep(100 * time.Second)
return nil
}).Maybe()
balance.Register(mb)
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
meta.EXPECT().ListBroadcastTask(mock.Anything).Return([]*streamingpb.BroadcastTask{}, nil).Times(1)
meta.EXPECT().SaveBroadcastTask(mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe()
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
mw := mock_streaming.NewMockWALAccesser(t)
streaming.SetWALForTest(mw)
bc, err := RecoverBroadcaster(context.Background())
assert.NoError(t, err)
// Should fail on primary cluster
api, err := bc.WithSecondaryClusterResourceKey(context.Background())
assert.Error(t, err)
assert.True(t, errors.Is(err, ErrNotSecondary))
assert.Nil(t, api)
bc.Close()
})
t.Run("context_canceled", func(t *testing.T) {
registry.ResetRegistration()
paramtable.Init()
balance.ResetBalancer()
mb := mock_balancer.NewMockBalancer(t)
mb.EXPECT().ReplicateRole().Return(replicateutil.RoleSecondary).Maybe()
mb.EXPECT().WatchChannelAssignments(mock.Anything, mock.Anything).RunAndReturn(func(ctx context.Context, cb balancer.WatchChannelAssignmentsCallback) error {
time.Sleep(100 * time.Second)
return nil
}).Maybe()
balance.Register(mb)
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
meta.EXPECT().ListBroadcastTask(mock.Anything).Return([]*streamingpb.BroadcastTask{}, nil).Times(1)
meta.EXPECT().SaveBroadcastTask(mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe()
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
mw := mock_streaming.NewMockWALAccesser(t)
streaming.SetWALForTest(mw)
bc, err := RecoverBroadcaster(context.Background())
assert.NoError(t, err)
// Use canceled context
ctx, cancel := context.WithCancel(context.Background())
cancel()
api, err := bc.WithSecondaryClusterResourceKey(ctx)
assert.Error(t, err)
assert.Nil(t, api)
bc.Close()
})
}
func createAlterReplicateConfigBroadcastMsg(vchannels []string, forcePromote bool) message.BroadcastMutableMessage {
msg := message.NewAlterReplicateConfigMessageBuilderV2().
WithHeader(&message.AlterReplicateConfigMessageHeader{
ReplicateConfiguration: &commonpb.ReplicateConfiguration{},
ForcePromote: forcePromote,
}).
WithBody(&message.AlterReplicateConfigMessageBody{}).
WithBroadcast(vchannels).
MustBuildBroadcast()
return msg
}
func TestIsAlterReplicateConfigMessage(t *testing.T) {
paramtable.Init()
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
t.Run("alter_replicate_config_message", func(t *testing.T) {
msg := createAlterReplicateConfigBroadcastMsg([]string{"v1"}, false).WithBroadcastID(1)
proto := createNewWaitAckBroadcastTaskFromMessage(msg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x00})
task := newBroadcastTaskFromProto(proto, metrics, ackScheduler)
assert.True(t, task.IsAlterReplicateConfigMessage())
})
t.Run("non_alter_replicate_config_message", func(t *testing.T) {
proto := createNewBroadcastTask(1, []string{"v1"})
task := newBroadcastTaskFromProto(proto, metrics, ackScheduler)
assert.False(t, task.IsAlterReplicateConfigMessage())
})
}
func TestIsForcePromoteMessage(t *testing.T) {
paramtable.Init()
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
t.Run("force_promote_true", func(t *testing.T) {
msg := createAlterReplicateConfigBroadcastMsg([]string{"v1"}, true).WithBroadcastID(1)
proto := createNewWaitAckBroadcastTaskFromMessage(msg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x00})
task := newBroadcastTaskFromProto(proto, metrics, ackScheduler)
assert.True(t, task.IsForcePromoteMessage())
})
t.Run("force_promote_false", func(t *testing.T) {
msg := createAlterReplicateConfigBroadcastMsg([]string{"v1"}, false).WithBroadcastID(2)
proto := createNewWaitAckBroadcastTaskFromMessage(msg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x00})
task := newBroadcastTaskFromProto(proto, metrics, ackScheduler)
assert.False(t, task.IsForcePromoteMessage())
})
t.Run("non_alter_replicate_config", func(t *testing.T) {
proto := createNewBroadcastTask(3, []string{"v1"})
task := newBroadcastTaskFromProto(proto, metrics, ackScheduler)
assert.False(t, task.IsForcePromoteMessage())
})
}
func TestPendingBroadcastMessages(t *testing.T) {
paramtable.Init()
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
t.Run("all_pending", func(t *testing.T) {
msg := createNewBroadcastMsg([]string{"v1", "v2", "v3"}).WithBroadcastID(1)
proto := createNewWaitAckBroadcastTaskFromMessage(msg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x00, 0x00, 0x00})
task := newBroadcastTaskFromProto(proto, metrics, ackScheduler)
pending := task.PendingBroadcastMessages()
assert.Len(t, pending, 3)
})
t.Run("some_acked", func(t *testing.T) {
msg := createNewBroadcastMsg([]string{"v1", "v2", "v3"}).WithBroadcastID(2)
proto := createNewWaitAckBroadcastTaskFromMessage(msg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01, 0x00, 0x01})
task := newBroadcastTaskFromProto(proto, metrics, ackScheduler)
pending := task.PendingBroadcastMessages()
assert.Len(t, pending, 1)
})
t.Run("all_acked", func(t *testing.T) {
msg := createNewBroadcastMsg([]string{"v1", "v2"}).WithBroadcastID(3)
proto := createNewWaitAckBroadcastTaskFromMessage(msg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01, 0x01})
task := newBroadcastTaskFromProto(proto, metrics, ackScheduler)
pending := task.PendingBroadcastMessages()
assert.Len(t, pending, 0)
})
}
func TestMarkIgnore(t *testing.T) {
paramtable.Init()
t.Run("success", func(t *testing.T) {
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
msg := createAlterReplicateConfigBroadcastMsg([]string{"v1", "v2"}, false).WithBroadcastID(10)
proto := createNewWaitAckBroadcastTaskFromMessage(msg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x00, 0x00})
task := newBroadcastTaskFromProto(proto, metrics, ackScheduler)
task.SetLogger(mlog.With())
err := task.MarkIgnore()
assert.NoError(t, err)
// Verify the message now has ignore=true
alterMsg, err := message.AsMutableAlterReplicateConfigMessageV2(task.msg)
assert.NoError(t, err)
assert.True(t, alterMsg.Header().Ignore)
})
t.Run("non_alter_replicate_config", func(t *testing.T) {
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
proto := createNewBroadcastTask(11, []string{"v1"})
task := newBroadcastTaskFromProto(proto, metrics, ackScheduler)
task.SetLogger(mlog.With())
err := task.MarkIgnore()
assert.Error(t, err)
})
}
func TestSortByControlChannelTimeTick(t *testing.T) {
paramtable.Init()
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
// Use single-vchannel (control channel only) tasks to avoid proto round-trip ordering issues
makeTask := func(broadcastID uint64, vchannel string, timeTick uint64) *broadcastTask {
msg := createNewBroadcastMsg([]string{vchannel}).WithBroadcastID(broadcastID)
p := createNewWaitAckBroadcastTaskFromMessage(msg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01})
p.AckedCheckpoints[0] = &streamingpb.AckedCheckpoint{
MessageId: walimplstest.NewTestMessageID(int64(broadcastID)).IntoProto(),
LastConfirmedMessageId: walimplstest.NewTestMessageID(int64(broadcastID)).IntoProto(),
TimeTick: timeTick,
}
return newBroadcastTaskFromProto(p, metrics, ackScheduler)
}
task1 := makeTask(1, "by-dev-1_vcchan", 30)
task2 := makeTask(2, "by-dev-2_vcchan", 10)
task3 := makeTask(3, "by-dev-3_vcchan", 20)
tasks := []*broadcastTask{task1, task3, task2}
sortByControlChannelTimeTick(tasks)
// Should be sorted by control channel timetick: 10, 20, 30
assert.Equal(t, uint64(2), tasks[0].Header().BroadcastID)
assert.Equal(t, uint64(3), tasks[1].Header().BroadcastID)
assert.Equal(t, uint64(1), tasks[2].Header().BroadcastID)
}
func TestFixIncompleteBroadcastsForForcePromote(t *testing.T) {
t.Run("no_incomplete_tasks", func(t *testing.T) {
paramtable.Init()
registry.ResetRegistration()
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
ackScheduler := newAckCallbackScheduler(mlog.With())
bm := &broadcastTaskManager{
mu: &sync.Mutex{},
tasks: make(map[uint64]*broadcastTask),
}
ackScheduler.bm = bm
err := ackScheduler.fixIncompleteBroadcastsForForcePromote(context.Background())
assert.NoError(t, err)
})
t.Run("with_alter_replicate_config_tasks", func(t *testing.T) {
paramtable.Init()
registry.ResetRegistration()
registry.RegisterAlterReplicateConfigV2AckCallback(
func(ctx context.Context, result message.BroadcastResult[*message.AlterReplicateConfigMessageHeader, *message.AlterReplicateConfigMessageBody]) error {
return nil
})
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
meta.EXPECT().SaveBroadcastTask(mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe()
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
alterMsg := createAlterReplicateConfigBroadcastMsg([]string{"v1", "v2"}, false).WithBroadcastID(100)
alterProto := createNewWaitAckBroadcastTaskFromMessage(alterMsg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01, 0x00})
alterTask := newBroadcastTaskFromProto(alterProto, metrics, ackScheduler)
alterTask.SetLogger(mlog.With())
mw := mock_streaming.NewMockWALAccesser(t)
appendF := func(ctx context.Context, msgs ...message.MutableMessage) types.AppendResponses {
resps := types.AppendResponses{Responses: make([]types.AppendResponse, len(msgs))}
for i := range msgs {
resps.Responses[i] = types.AppendResponse{
AppendResult: &types.AppendResult{
MessageID: walimplstest.NewTestMessageID(int64(i + 1)),
TimeTick: uint64(100 + i),
},
}
}
return resps
}
mw.EXPECT().AppendMessages(mock.Anything, mock.Anything).RunAndReturn(appendF).Maybe()
mw.EXPECT().AppendMessages(mock.Anything, mock.Anything, mock.Anything).RunAndReturn(appendF).Maybe()
streaming.SetWALForTest(mw)
bm := &broadcastTaskManager{
lifetime: typeutil.NewLifetime(),
mu: &sync.Mutex{},
tasks: map[uint64]*broadcastTask{100: alterTask},
broadcastScheduler: newBroadcasterScheduler(nil, mlog.With()),
}
ackScheduler.bm = bm
ackScheduler.Initialize(nil, nil, bm)
defer ackScheduler.Close()
defer bm.broadcastScheduler.Close()
err := ackScheduler.fixIncompleteBroadcastsForForcePromote(context.Background())
assert.NoError(t, err)
parsedMsg, err := message.AsMutableAlterReplicateConfigMessageV2(alterTask.msg)
assert.NoError(t, err)
assert.True(t, parsedMsg.Header().Ignore)
})
t.Run("with_other_broadcast_tasks", func(t *testing.T) {
paramtable.Init()
registry.ResetRegistration()
registry.RegisterDropCollectionV1AckCallback(func(ctx context.Context, msg message.BroadcastResultDropCollectionMessageV1) error {
return nil
})
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
meta.EXPECT().SaveBroadcastTask(mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe()
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
dropMsg := createNewBroadcastMsg([]string{"v1", "v2", "v3"}).WithBroadcastID(200)
dropProto := createNewWaitAckBroadcastTaskFromMessage(dropMsg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01, 0x00, 0x00})
dropTask := newBroadcastTaskFromProto(dropProto, metrics, ackScheduler)
dropTask.SetLogger(mlog.With())
appendedCount := atomic.NewInt32(0)
mw := mock_streaming.NewMockWALAccesser(t)
appendF2 := func(ctx context.Context, msgs ...message.MutableMessage) types.AppendResponses {
resps := types.AppendResponses{Responses: make([]types.AppendResponse, len(msgs))}
for i := range msgs {
appendedCount.Inc()
resps.Responses[i] = types.AppendResponse{
AppendResult: &types.AppendResult{
MessageID: walimplstest.NewTestMessageID(int64(i + 1)),
TimeTick: uint64(100 + i),
},
}
}
return resps
}
mw.EXPECT().AppendMessages(mock.Anything, mock.Anything).RunAndReturn(appendF2).Maybe()
mw.EXPECT().AppendMessages(mock.Anything, mock.Anything, mock.Anything).RunAndReturn(appendF2).Maybe()
mw.EXPECT().AppendMessages(mock.Anything, mock.Anything, mock.Anything, mock.Anything).RunAndReturn(appendF2).Maybe()
streaming.SetWALForTest(mw)
bm := &broadcastTaskManager{
lifetime: typeutil.NewLifetime(),
mu: &sync.Mutex{},
tasks: map[uint64]*broadcastTask{200: dropTask},
broadcastScheduler: newBroadcasterScheduler(nil, mlog.With()),
}
ackScheduler.bm = bm
ackScheduler.Initialize(nil, nil, bm)
defer ackScheduler.Close()
defer bm.broadcastScheduler.Close()
err := ackScheduler.fixIncompleteBroadcastsForForcePromote(context.Background())
assert.NoError(t, err)
assert.Equal(t, int32(2), appendedCount.Load())
})
t.Run("append_failure_then_retry", func(t *testing.T) {
paramtable.Init()
registry.ResetRegistration()
registry.RegisterDropCollectionV1AckCallback(func(ctx context.Context, msg message.BroadcastResultDropCollectionMessageV1) error {
return nil
})
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
meta.EXPECT().SaveBroadcastTask(mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe()
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
dropMsg := createNewBroadcastMsg([]string{"v1", "v2"}).WithBroadcastID(300)
dropProto := createNewWaitAckBroadcastTaskFromMessage(dropMsg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01, 0x00})
dropTask := newBroadcastTaskFromProto(dropProto, metrics, ackScheduler)
dropTask.SetLogger(mlog.With())
// First call fails, subsequent calls succeed
callCount := atomic.NewInt32(0)
mw := mock_streaming.NewMockWALAccesser(t)
appendF := func(ctx context.Context, msgs ...message.MutableMessage) types.AppendResponses {
resps := types.AppendResponses{Responses: make([]types.AppendResponse, len(msgs))}
count := callCount.Inc()
for i := range msgs {
if count == 1 {
resps.Responses[i] = types.AppendResponse{Error: errors.New("append failed")}
} else {
resps.Responses[i] = types.AppendResponse{
AppendResult: &types.AppendResult{
MessageID: walimplstest.NewTestMessageID(int64(i + 1)),
TimeTick: uint64(100 + i),
},
}
}
}
return resps
}
mw.EXPECT().AppendMessages(mock.Anything, mock.Anything).RunAndReturn(appendF).Maybe()
streaming.SetWALForTest(mw)
bm := &broadcastTaskManager{
lifetime: typeutil.NewLifetime(),
mu: &sync.Mutex{},
tasks: map[uint64]*broadcastTask{300: dropTask},
broadcastScheduler: newBroadcasterScheduler(nil, mlog.With()),
}
ackScheduler.bm = bm
ackScheduler.Initialize(nil, nil, bm)
defer ackScheduler.Close()
defer bm.broadcastScheduler.Close()
err := ackScheduler.fixIncompleteBroadcastsForForcePromote(context.Background())
assert.NoError(t, err)
// broadcastScheduler retried after first failure
assert.GreaterOrEqual(t, callCount.Load(), int32(2))
})
t.Run("blocks_until_tombstone", func(t *testing.T) {
paramtable.Init()
registry.ResetRegistration()
registry.RegisterDropCollectionV1AckCallback(func(ctx context.Context, msg message.BroadcastResultDropCollectionMessageV1) error {
return nil
})
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
meta.EXPECT().SaveBroadcastTask(mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe()
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
// Create an incomplete task (v2 not acked)
dropMsg := createNewBroadcastMsg([]string{"v1", "v2"}).WithBroadcastID(500)
dropProto := createNewWaitAckBroadcastTaskFromMessage(dropMsg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01, 0x00})
dropTask := newBroadcastTaskFromProto(dropProto, metrics, ackScheduler)
dropTask.SetLogger(mlog.With())
mw := mock_streaming.NewMockWALAccesser(t)
appendF := func(ctx context.Context, msgs ...message.MutableMessage) types.AppendResponses {
resps := types.AppendResponses{Responses: make([]types.AppendResponse, len(msgs))}
for i := range msgs {
resps.Responses[i] = types.AppendResponse{
AppendResult: &types.AppendResult{
MessageID: walimplstest.NewTestMessageID(int64(i + 1)),
TimeTick: uint64(100 + i),
},
}
}
return resps
}
mw.EXPECT().AppendMessages(mock.Anything, mock.Anything).RunAndReturn(appendF).Maybe()
streaming.SetWALForTest(mw)
bm := &broadcastTaskManager{
lifetime: typeutil.NewLifetime(),
mu: &sync.Mutex{},
tasks: map[uint64]*broadcastTask{500: dropTask},
broadcastScheduler: newBroadcasterScheduler(nil, mlog.With()),
}
ackScheduler.bm = bm
ackScheduler.Initialize(nil, nil, bm)
defer ackScheduler.Close()
defer bm.broadcastScheduler.Close()
// AddTask blocks until tombstone; fixIncompleteBroadcastsForForcePromote
// should only return after task reaches TOMBSTONE via broadcastScheduler.
err := ackScheduler.fixIncompleteBroadcastsForForcePromote(context.Background())
assert.NoError(t, err)
assert.Equal(t, streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_TOMBSTONE, dropTask.State())
})
t.Run("context_canceled_during_supplement", func(t *testing.T) {
paramtable.Init()
registry.ResetRegistration()
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
meta.EXPECT().SaveBroadcastTask(mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe()
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
dropMsg := createNewBroadcastMsg([]string{"v1", "v2"}).WithBroadcastID(600)
dropProto := createNewWaitAckBroadcastTaskFromMessage(dropMsg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01, 0x00})
dropTask := newBroadcastTaskFromProto(dropProto, metrics, ackScheduler)
dropTask.SetLogger(mlog.With())
// WAL mock succeeds but never acks
mw := mock_streaming.NewMockWALAccesser(t)
mw.EXPECT().AppendMessages(mock.Anything, mock.Anything).RunAndReturn(
func(ctx context.Context, msgs ...message.MutableMessage) types.AppendResponses {
resps := types.AppendResponses{Responses: make([]types.AppendResponse, len(msgs))}
for i := range msgs {
resps.Responses[i] = types.AppendResponse{
AppendResult: &types.AppendResult{
MessageID: walimplstest.NewTestMessageID(int64(i + 1)),
TimeTick: uint64(100 + i),
},
}
}
return resps
}).Maybe()
streaming.SetWALForTest(mw)
bm := &broadcastTaskManager{
lifetime: typeutil.NewLifetime(),
mu: &sync.Mutex{},
tasks: map[uint64]*broadcastTask{600: dropTask},
broadcastScheduler: newBroadcasterScheduler(nil, mlog.With()),
}
ackScheduler.bm = bm
ackScheduler.Initialize(nil, nil, bm)
defer ackScheduler.Close()
defer bm.broadcastScheduler.Close()
ctx, cancel := context.WithCancel(context.Background())
done := make(chan error, 1)
go func() {
done <- ackScheduler.fixIncompleteBroadcastsForForcePromote(ctx)
}()
// Cancel context while AddTask is blocking
time.Sleep(100 * time.Millisecond)
cancel()
select {
case err := <-done:
assert.Error(t, err)
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for context cancellation")
}
})
}
func TestDoForcePromoteFixIncompleteBroadcasts(t *testing.T) {
t.Run("full_lifecycle_no_incomplete_tasks", func(t *testing.T) {
paramtable.Init()
registry.ResetRegistration()
// Register a no-op ack callback for AlterReplicateConfig so doAckCallback can proceed.
registry.RegisterAlterReplicateConfigV2AckCallback(
func(ctx context.Context, result message.BroadcastResult[*message.AlterReplicateConfigMessageHeader, *message.AlterReplicateConfigMessageBody]) error {
return nil
})
meta := mock_metastore.NewMockStreamingCoordCataLog(t)
meta.EXPECT().SaveBroadcastTask(mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe()
rc := idalloc.NewMockRootCoordClient(t)
f := syncutil.NewFuture[internaltypes.MixCoordClient]()
f.Set(rc)
resource.InitForTest(resource.OptStreamingCatalog(meta), resource.OptMixCoordClient(f))
mw := mock_streaming.NewMockWALAccesser(t)
streaming.SetWALForTest(mw)
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
// Create a force promote task that is already all acked
fpMsg := createAlterReplicateConfigBroadcastMsg([]string{"v1"}, true).WithBroadcastID(400)
fpProto := createNewWaitAckBroadcastTaskFromMessage(fpMsg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x01}) // already acked
fpTask := newBroadcastTaskFromProto(fpProto, metrics, ackScheduler)
fpTask.SetLogger(mlog.With())
// No incomplete tasks in the bm
bm := &broadcastTaskManager{
lifetime: typeutil.NewLifetime(),
mu: &sync.Mutex{},
tasks: map[uint64]*broadcastTask{400: fpTask},
}
ackScheduler.bm = bm
ackScheduler.Initialize(nil, nil, bm)
defer ackScheduler.Close()
// doForcePromoteFixIncompleteBroadcasts should complete the full lifecycle:
// BlockUntilAllAck → fix (no-op) → acquire lock → doAckCallback → close(done)
done := make(chan struct{})
go func() {
ackScheduler.doForcePromoteFixIncompleteBroadcasts(fpTask)
close(done)
}()
select {
case <-done:
// Verify task reached TOMBSTONE (ack callback completed)
assert.Equal(t, streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_TOMBSTONE, fpTask.State())
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for doForcePromoteFixIncompleteBroadcasts")
}
})
t.Run("context_canceled_before_ack", func(t *testing.T) {
paramtable.Init()
registry.ResetRegistration()
resource.InitForTest()
metrics := newBroadcasterMetrics()
ackScheduler := newAckCallbackScheduler(mlog.With())
// Create a force promote task that is NOT all acked
fpMsg := createAlterReplicateConfigBroadcastMsg([]string{"v1", "v2"}, true).WithBroadcastID(401)
fpProto := createNewWaitAckBroadcastTaskFromMessage(fpMsg,
streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING,
[]byte{0x00, 0x00}) // not acked
fpTask := newBroadcastTaskFromProto(fpProto, metrics, ackScheduler)
fpTask.SetLogger(mlog.With())
bm := &broadcastTaskManager{
mu: &sync.Mutex{},
tasks: make(map[uint64]*broadcastTask),
}
ackScheduler.bm = bm
done := make(chan struct{})
go func() {
ackScheduler.doForcePromoteFixIncompleteBroadcasts(fpTask)
close(done)
}()
// Cancel the scheduler context — should abort at BlockUntilAllAck
ackScheduler.notifier.Cancel()
select {
case <-done:
// Should return because context canceled, task NOT tombstoned
assert.NotEqual(t, streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_TOMBSTONE, fpTask.State())
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for doForcePromoteFixIncompleteBroadcasts to exit on cancel")
}
})
}