Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions internal/coordinator/snmanager/streaming_node_manager.go
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,11 @@ type StreamingNodeManager struct {
nodeChangedNotifier *syncutil.VersionedNotifier // used to notify that node in streaming node manager has been changed.
}

// GetBalancer returns the balancer of the streaming node manager.
func (s *StreamingNodeManager) GetBalancer() balancer.Balancer {
return s.balancer.Get()
}

// GetLatestWALLocated returns the server id of the node that the wal of the vChannel is located.
// Return -1 and error if the vchannel is not found or context is canceled.
func (s *StreamingNodeManager) GetLatestWALLocated(ctx context.Context, vchannel string) (int64, error) {
Expand Down
138 changes: 119 additions & 19 deletions internal/distributed/streaming/balancer.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,16 @@ package streaming
import (
"context"

"github.com/cockroachdb/errors"
"google.golang.org/protobuf/types/known/fieldmaskpb"

"github.com/milvus-io/milvus/internal/coordinator/snmanager"
"github.com/milvus-io/milvus/internal/streamingcoord/server/balancer"
"github.com/milvus-io/milvus/pkg/v2/proto/streamingpb"
"github.com/milvus-io/milvus/pkg/v2/streaming/util/types"
"github.com/milvus-io/milvus/pkg/v2/util/merr"
"github.com/milvus-io/milvus/pkg/v2/util/paramtable"
"github.com/milvus-io/milvus/pkg/v2/util/typeutil"
)

type balancerImpl struct {
Expand All @@ -16,49 +21,94 @@ type balancerImpl struct {

// GetWALDistribution returns the wal distribution of the streaming node.
func (b balancerImpl) ListStreamingNode(ctx context.Context) ([]types.StreamingNodeInfo, error) {
assignments, err := b.streamingCoordClient.Assignment().GetLatestAssignments(ctx)
ready, err := b.checkIfStreamingServiceReady(ctx)
if err != nil {
return nil, err
}
if !ready {
return nil, nil
}

nodes := make([]types.StreamingNodeInfo, 0, len(assignments.Assignments))
for _, assignment := range assignments.Assignments {
nodes = append(nodes, assignment.NodeInfo)
nodes, err := snmanager.StaticStreamingNodeManager.GetBalancer().GetAllStreamingNodes(ctx)
if err != nil {
return nil, err
}
return nodes, nil
nodeInfos := make([]types.StreamingNodeInfo, 0, len(nodes))
for _, node := range nodes {
nodeInfos = append(nodeInfos, *node)
}
return nodeInfos, nil
}

// GetWALDistribution returns the wal distribution of the streaming node.
func (b balancerImpl) GetWALDistribution(ctx context.Context, nodeID int64) (*types.StreamingNodeAssignment, error) {
assignments, err := b.streamingCoordClient.Assignment().GetLatestAssignments(ctx)
ready, err := b.checkIfStreamingServiceReady(ctx)
if err != nil {
return nil, err
}
for _, assignment := range assignments.Assignments {
if assignment.NodeInfo.ServerID == nodeID {
return &assignment, nil
if !ready {
return nil, nil
}

sbalancer := snmanager.StaticStreamingNodeManager.GetBalancer()
var result *types.StreamingNodeAssignment
stopErr := errors.New("stop watching")
err = sbalancer.WatchChannelAssignments(ctx, func(param balancer.WatchChannelAssignmentsCallbackParam) error {
for _, assignment := range param.Relations {
if assignment.Node.ServerID == nodeID {
if result == nil {
result = &types.StreamingNodeAssignment{
NodeInfo: assignment.Node,
Channels: make(map[string]types.PChannelInfo),
}
}
result.Channels[assignment.Channel.Name] = assignment.Channel
}
}
return errors.New("stop watching")
})
if errors.Is(err, stopErr) {
if result == nil {
return nil, merr.ErrNodeNotFound
}
return result, nil
}
return nil, merr.WrapErrNodeNotFound(nodeID, "streaming node not found")
return nil, err
}

// GetFrozenNodeIDs returns the frozen node ids.
func (b balancerImpl) GetFrozenNodeIDs(ctx context.Context) ([]int64, error) {
// Update nothing, just fetch the current resp back.
resp, err := b.streamingCoordClient.Assignment().UpdateWALBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
ready, err := b.checkIfStreamingServiceReady(ctx)
if err != nil {
return nil, err
}
if !ready {
return nil, nil
}

sbalancer := snmanager.StaticStreamingNodeManager.GetBalancer()
resp, err := sbalancer.UpdateBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
Config: &streamingpb.WALBalancePolicyConfig{},
UpdateMask: &fieldmaskpb.FieldMask{},
})
if err != nil {
return nil, err
}
return resp.GetFreezeNodeIds(), nil
return resp.FreezeNodeIds, nil
}

// IsRebalanceSuspended returns whether the rebalance of the wal is suspended.
func (b balancerImpl) IsRebalanceSuspended(ctx context.Context) (bool, error) {
// Update nothing, just fetch the current resp back.
resp, err := b.streamingCoordClient.Assignment().UpdateWALBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
ready, err := b.checkIfStreamingServiceReady(ctx)
if err != nil {
return false, err
}
if !ready {
return false, nil
}

sbalancer := snmanager.StaticStreamingNodeManager.GetBalancer()
resp, err := sbalancer.UpdateBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
Config: &streamingpb.WALBalancePolicyConfig{},
UpdateMask: &fieldmaskpb.FieldMask{},
})
Expand All @@ -69,7 +119,16 @@ func (b balancerImpl) IsRebalanceSuspended(ctx context.Context) (bool, error) {
}

func (b balancerImpl) SuspendRebalance(ctx context.Context) error {
_, err := b.streamingCoordClient.Assignment().UpdateWALBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
ready, err := b.checkIfStreamingServiceReady(ctx)
if err != nil {
return err
}
if !ready {
return nil
}

sbalancer := snmanager.StaticStreamingNodeManager.GetBalancer()
_, err = sbalancer.UpdateBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
Config: &streamingpb.WALBalancePolicyConfig{
AllowRebalance: false,
},
Expand All @@ -81,7 +140,16 @@ func (b balancerImpl) SuspendRebalance(ctx context.Context) error {
}

func (b balancerImpl) ResumeRebalance(ctx context.Context) error {
_, err := b.streamingCoordClient.Assignment().UpdateWALBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
ready, err := b.checkIfStreamingServiceReady(ctx)
if err != nil {
return err
}
if !ready {
return nil
}

sbalancer := snmanager.StaticStreamingNodeManager.GetBalancer()
_, err = sbalancer.UpdateBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
Config: &streamingpb.WALBalancePolicyConfig{
AllowRebalance: true,
},
Expand All @@ -93,7 +161,16 @@ func (b balancerImpl) ResumeRebalance(ctx context.Context) error {
}

func (b balancerImpl) FreezeNodeIDs(ctx context.Context, nodeIDs []int64) error {
_, err := b.streamingCoordClient.Assignment().UpdateWALBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
ready, err := b.checkIfStreamingServiceReady(ctx)
if err != nil {
return err
}
if !ready {
return nil
}

sbalancer := snmanager.StaticStreamingNodeManager.GetBalancer()
_, err = sbalancer.UpdateBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{}},
Nodes: &streamingpb.WALBalancePolicyNodes{
FreezeNodeIds: nodeIDs,
Expand All @@ -103,11 +180,34 @@ func (b balancerImpl) FreezeNodeIDs(ctx context.Context, nodeIDs []int64) error
}

func (b balancerImpl) DefreezeNodeIDs(ctx context.Context, nodeIDs []int64) error {
_, err := b.streamingCoordClient.Assignment().UpdateWALBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
ready, err := b.checkIfStreamingServiceReady(ctx)
if err != nil {
return err
}
if !ready {
return nil
}

sbalancer := snmanager.StaticStreamingNodeManager.GetBalancer()
_, err = sbalancer.UpdateBalancePolicy(ctx, &types.UpdateWALBalancePolicyRequest{
UpdateMask: &fieldmaskpb.FieldMask{Paths: []string{}},
Nodes: &streamingpb.WALBalancePolicyNodes{
DefreezeNodeIds: nodeIDs,
},
})
return err
}

func (b balancerImpl) checkIfStreamingServiceReady(ctx context.Context) (bool, error) {
if !paramtable.IsLocalComponentEnabled(typeutil.MixCoordRole) {
panic("should be only called at mix coord")
}
if err := snmanager.StaticStreamingNodeManager.CheckIfStreamingServiceReady(ctx); err != nil {
if errors.Is(err, snmanager.ErrStreamingServiceNotReady) {
// for 2.5.x compatibility, return empty result when streaming service is not ready.
return false, nil
}
return false, err
}
return true, nil
}
101 changes: 79 additions & 22 deletions internal/distributed/streaming/balancer_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,59 +3,87 @@ package streaming
import (
"context"
"testing"
"time"

"github.com/cockroachdb/errors"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"

"github.com/milvus-io/milvus/internal/mocks/streamingcoord/mock_client"
"github.com/milvus-io/milvus/internal/coordinator/snmanager"
"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/pkg/v2/proto/streamingpb"
"github.com/milvus-io/milvus/pkg/v2/streaming/util/types"
"github.com/milvus-io/milvus/pkg/v2/util/merr"
"github.com/milvus-io/milvus/pkg/v2/util/paramtable"
"github.com/milvus-io/milvus/pkg/v2/util/syncutil"
"github.com/milvus-io/milvus/pkg/v2/util/typeutil"
)

func TestBalancer(t *testing.T) {
scClient := mock_client.NewMockClient(t)
assignmentService := mock_client.NewMockAssignmentService(t)
scClient.EXPECT().Assignment().Return(assignmentService)
assignmentService.EXPECT().GetLatestAssignments(mock.Anything).Return(&types.VersionedStreamingNodeAssignments{
Assignments: map[int64]types.StreamingNodeAssignment{
1: {
NodeInfo: types.StreamingNodeInfo{ServerID: 1},
Channels: map[string]types.PChannelInfo{
"v1": {},
paramtable.SetLocalComponentEnabled(typeutil.MixCoordRole)
sbalancer := mock_balancer.NewMockBalancer(t)
sbalancer.EXPECT().GetAllStreamingNodes(mock.Anything).Return(map[int64]*types.StreamingNodeInfo{
1: {ServerID: 1},
2: {ServerID: 2},
}, nil)
sbalancer.EXPECT().WatchChannelAssignments(mock.Anything, mock.Anything).RunAndReturn(func(ctx context.Context, cb balancer.WatchChannelAssignmentsCallback) error {
if err := cb(balancer.WatchChannelAssignmentsCallbackParam{
Version: typeutil.VersionInt64Pair{
Global: 1,
Local: 1,
},
Relations: []types.PChannelInfoAssigned{
{
Channel: types.PChannelInfo{Name: "v1"},
Node: types.StreamingNodeInfo{ServerID: 1},
},
{
Channel: types.PChannelInfo{Name: "v2"},
Node: types.StreamingNodeInfo{ServerID: 1},
},
},
},
}, nil)
}); err != nil {
return err
}
time.Sleep(100 * time.Millisecond)
return nil
})
sbalancer.EXPECT().RegisterStreamingEnabledNotifier(mock.Anything).RunAndReturn(func(notifier *syncutil.AsyncTaskNotifier[struct{}]) {
notifier.Cancel()
})

snmanager.ResetStreamingNodeManager()
snmanager.StaticStreamingNodeManager.SetBalancerReady(sbalancer)

balancer := balancerImpl{
walAccesserImpl: &walAccesserImpl{
streamingCoordClient: scClient,
},
walAccesserImpl: &walAccesserImpl{},
}

nodes, err := balancer.ListStreamingNode(context.Background())
assert.NoError(t, err)
assert.Equal(t, 1, len(nodes))
assert.Equal(t, 2, len(nodes))
assignment, err := balancer.GetWALDistribution(context.Background(), 1)
assert.NoError(t, err)
assert.Equal(t, 1, len(assignment.Channels))
assert.Equal(t, 2, len(assignment.Channels))

assignment, err = balancer.GetWALDistribution(context.Background(), 2)
assert.True(t, errors.Is(err, merr.ErrNodeNotFound))
assert.Nil(t, assignment)

assignmentService.EXPECT().GetLatestAssignments(mock.Anything).Unset()
assignmentService.EXPECT().GetLatestAssignments(mock.Anything).Return(nil, errors.New("test"))
sbalancer.EXPECT().GetAllStreamingNodes(mock.Anything).Unset()
sbalancer.EXPECT().GetAllStreamingNodes(mock.Anything).Return(nil, errors.New("test"))
nodes, err = balancer.ListStreamingNode(context.Background())
assert.Error(t, err)
assert.Nil(t, nodes)

sbalancer.EXPECT().WatchChannelAssignments(mock.Anything, mock.Anything).Unset()
sbalancer.EXPECT().WatchChannelAssignments(mock.Anything, mock.Anything).Return(errors.New("test"))
assignment, err = balancer.GetWALDistribution(context.Background(), 1)
assert.Error(t, err)
assert.Nil(t, assignment)

assignmentService.EXPECT().UpdateWALBalancePolicy(mock.Anything, mock.Anything).Return(&types.UpdateWALBalancePolicyResponse{}, nil)
sbalancer.EXPECT().UpdateBalancePolicy(mock.Anything, mock.Anything).Return(&streamingpb.UpdateWALBalancePolicyResponse{}, nil)
err = balancer.SuspendRebalance(context.Background())
assert.NoError(t, err)
err = balancer.ResumeRebalance(context.Background())
Expand All @@ -64,9 +92,13 @@ func TestBalancer(t *testing.T) {
assert.NoError(t, err)
err = balancer.DefreezeNodeIDs(context.Background(), []int64{1})
assert.NoError(t, err)
_, err = balancer.GetFrozenNodeIDs(context.Background())
assert.NoError(t, err)
_, err = balancer.IsRebalanceSuspended(context.Background())
assert.NoError(t, err)

assignmentService.EXPECT().UpdateWALBalancePolicy(mock.Anything, mock.Anything).Unset()
assignmentService.EXPECT().UpdateWALBalancePolicy(mock.Anything, mock.Anything).Return(nil, errors.New("test"))
sbalancer.EXPECT().UpdateBalancePolicy(mock.Anything, mock.Anything).Unset()
sbalancer.EXPECT().UpdateBalancePolicy(mock.Anything, mock.Anything).Return(nil, errors.New("test"))
err = balancer.SuspendRebalance(context.Background())
assert.Error(t, err)
err = balancer.ResumeRebalance(context.Background())
Expand All @@ -75,4 +107,29 @@ func TestBalancer(t *testing.T) {
assert.Error(t, err)
err = balancer.DefreezeNodeIDs(context.Background(), []int64{1})
assert.Error(t, err)
_, err = balancer.GetFrozenNodeIDs(context.Background())
assert.Error(t, err)
_, err = balancer.IsRebalanceSuspended(context.Background())
assert.Error(t, err)

sbalancer.EXPECT().RegisterStreamingEnabledNotifier(mock.Anything).Unset()
sbalancer.EXPECT().RegisterStreamingEnabledNotifier(mock.Anything).RunAndReturn(func(notifier *syncutil.AsyncTaskNotifier[struct{}]) {
})

_, err = balancer.ListStreamingNode(context.Background())
assert.NoError(t, err)
_, err = balancer.GetWALDistribution(context.Background(), 1)
assert.NoError(t, err)
err = balancer.SuspendRebalance(context.Background())
assert.NoError(t, err)
err = balancer.ResumeRebalance(context.Background())
assert.NoError(t, err)
err = balancer.FreezeNodeIDs(context.Background(), []int64{1})
assert.NoError(t, err)
err = balancer.DefreezeNodeIDs(context.Background(), []int64{1})
assert.NoError(t, err)
_, err = balancer.GetFrozenNodeIDs(context.Background())
assert.NoError(t, err)
_, err = balancer.IsRebalanceSuspended(context.Background())
assert.NoError(t, err)
}
1 change: 1 addition & 0 deletions internal/streamingcoord/client/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ type AssignmentService interface {

// UpdateWALBalancePolicy is used to update the WAL balance policy.
// Return the WAL balance policy after the update.
// Deprecated: This function is deprecated and will be removed in the future.
UpdateWALBalancePolicy(ctx context.Context, req *types.UpdateWALBalancePolicyRequest) (*types.UpdateWALBalancePolicyResponse, error)
}

Expand Down
Loading