diff --git a/engine/storage/raft_storage/cluster_metadata_test.go b/engine/storage/raft_storage/cluster_metadata_test.go index fc3564f..742f1d7 100644 --- a/engine/storage/raft_storage/cluster_metadata_test.go +++ b/engine/storage/raft_storage/cluster_metadata_test.go @@ -55,6 +55,12 @@ func TestClusterMetadataSurvivesSnapshotAndRestart(t *testing.T) { } t.Cleanup(func() { _ = second.Stop() }) assertClusterMetadata(t, second.state.cluster, 42, map[uint64]string{1: "127.0.0.1:1"}) + second.transport.mu.RLock() + routedAddress := second.transport.members[1] + second.transport.mu.RUnlock() + if routedAddress != "127.0.0.1:1" { + t.Fatalf("transport used bootstrap address %q instead of persisted metadata", routedAddress) + } } func TestApplyMembershipPersistsConfStateAndMetadata(t *testing.T) { diff --git a/engine/storage/raft_storage/membership.go b/engine/storage/raft_storage/membership.go index 5b23814..c3cbde3 100644 --- a/engine/storage/raft_storage/membership.go +++ b/engine/storage/raft_storage/membership.go @@ -32,6 +32,19 @@ const ( memberRoleLearner ) +func clusterAddresses(metadata *clusterpb.ClusterMetadata) map[uint64]string { + addresses := make(map[uint64]string) + if metadata == nil { + return addresses + } + for _, member := range metadata.Members { + if member != nil { + addresses[member.Id] = member.RaftAddress + } + } + return addresses +} + func findMember(metadata *clusterpb.ClusterMetadata, id uint64) (*clusterpb.Member, int) { if metadata == nil { return nil, -1 diff --git a/engine/storage/raft_storage/raft.go b/engine/storage/raft_storage/raft.go index d081d87..cbc4808 100644 --- a/engine/storage/raft_storage/raft.go +++ b/engine/storage/raft_storage/raft.go @@ -122,6 +122,9 @@ func (rs *RaftStorage) handleReady() error { if err := rs.state.persist(&ready); err != nil { return fmt.Errorf("raft storage: persist ready: %w", err) } + if !raft.IsEmptySnap(ready.Snapshot) { + rs.transport.ReplaceMembers(clusterAddresses(rs.state.cluster)) + } for _, entry := range ready.CommittedEntries { switch entry.EntryType { @@ -154,6 +157,7 @@ func (rs *RaftStorage) handleReady() error { if err := rs.state.applyMembership(entry.Index, state, nextCluster); err != nil { return fmt.Errorf("raft storage: apply conf change at %d: %w", entry.Index, err) } + rs.transport.ReplaceMembers(clusterAddresses(nextCluster)) if context.ProposerId == rs.config.StoreID { rs.completePending(context.Sequence, nil) } diff --git a/engine/storage/raft_storage/raft_storage.go b/engine/storage/raft_storage/raft_storage.go index f573f81..6f6c66f 100644 --- a/engine/storage/raft_storage/raft_storage.go +++ b/engine/storage/raft_storage/raft_storage.go @@ -38,14 +38,16 @@ type requestID struct { sequence uint64 } -// RaftStorage runs one local replica of a fixed Raft-backed KV storage. +// RaftStorage runs one local replica of the Raft-backed KV and membership +// state machines. type RaftStorage struct { rspb.UnimplementedRaftServiceServer - config *config.Config - db *badger.DB - node *raft.RawNode - state *raftStatePersistence + config *config.Config + clusterID uint64 + db *badger.DB + node *raft.RawNode + state *raftStatePersistence transport *ServerTransport cancel context.CancelFunc @@ -163,8 +165,9 @@ func (rs *RaftStorage) Start() error { } rs.state = state rs.node = rawNode + rs.clusterID = state.cluster.ClusterId - rs.transport = NewServerTransport(rs.config) + rs.transport = NewServerTransport(rs.clusterID, clusterAddresses(state.cluster)) rs.inbox = make(chan raftEvent, 256) rs.done = make(chan struct{}) rs.runErr = nil @@ -231,13 +234,16 @@ func (rs *RaftStorage) Raft(stream rspb.RaftService_RaftServer) error { if err != nil { return err } - if message == nil { + if message == nil || message.Message == nil { continue } - if message.To != 0 && message.To != rs.config.StoreID { + if message.ClusterId != rs.clusterID { + return fmt.Errorf("raft storage: cluster ID mismatch: got %d, expected %d", message.ClusterId, rs.clusterID) + } + if message.Message.To != 0 && message.Message.To != rs.config.StoreID { continue } - if err := rs.step(stream.Context(), message); err != nil { + if err := rs.step(stream.Context(), message.Message); err != nil { return err } } @@ -320,6 +326,7 @@ func (rs *RaftStorage) onApplied(data []byte) { func (rs *RaftStorage) sendMessages(messages []*raftpb.Message) { for _, message := range messages { if err := rs.transport.Send(message); err != nil { + rs.node.ReportSnapshot(message.To, false) log.Printf("send raft message %s from %d to %d: %v", message.MsgType, message.From, message.To, err) } } diff --git a/engine/storage/raft_storage/transport.go b/engine/storage/raft_storage/transport.go index c5bca43..7c3e029 100644 --- a/engine/storage/raft_storage/transport.go +++ b/engine/storage/raft_storage/transport.go @@ -8,7 +8,6 @@ import ( "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" - "github.com/Aetherance/kv/engine/config" rspb "github.com/Aetherance/kv/proto/pkg/raft_serverpb" "github.com/Aetherance/kv/proto/pkg/raftpb" ) @@ -16,7 +15,7 @@ import ( // raftConn is a pooled gRPC client stream to one store. type raftConn struct { streamMu sync.Mutex - stream grpc.ClientStreamingClient[raftpb.Message, rspb.Done] + stream grpc.ClientStreamingClient[rspb.RaftEnvelope, rspb.Done] cancel context.CancelFunc client *grpc.ClientConn } @@ -36,7 +35,7 @@ func newRaftConn(addr string) (*raftConn, error) { return &raftConn{stream: stream, cancel: cancel, client: cc}, nil } -func (c *raftConn) Send(msg *raftpb.Message) error { +func (c *raftConn) Send(msg *rspb.RaftEnvelope) error { c.streamMu.Lock() defer c.streamMu.Unlock() return c.stream.Send(msg) @@ -47,19 +46,22 @@ func (c *raftConn) Stop() { _ = c.client.Close() } -// ServerTransport sends messages for the fixed-topology Raft storage. -// Connections are resolved by the raw Raft destination ID and pooled. +// ServerTransport resolves destinations from the applied cluster metadata. +// Replacing the member map also drops connections whose address changed or +// whose member was removed. type ServerTransport struct { - cfg *config.Config - mu sync.RWMutex + clusterID uint64 + mu sync.RWMutex + members map[uint64]string // store id -> connection conns map[uint64]*raftConn } -func NewServerTransport(cfg *config.Config) *ServerTransport { +func NewServerTransport(clusterID uint64, members map[uint64]string) *ServerTransport { return &ServerTransport{ - cfg: cfg, - conns: make(map[uint64]*raftConn), + clusterID: clusterID, + members: cloneAddresses(members), + conns: make(map[uint64]*raftConn), } } @@ -68,7 +70,9 @@ func (t *ServerTransport) Send(message *raftpb.Message) error { return nil } storeID := message.To - addr, ok := t.cfg.Peers[storeID] + t.mu.RLock() + addr, ok := t.members[storeID] + t.mu.RUnlock() if !ok { return fmt.Errorf("no address for store %d", storeID) } @@ -76,7 +80,7 @@ func (t *ServerTransport) Send(message *raftpb.Message) error { if err != nil { return err } - if err := conn.Send(message); err != nil { + if err := conn.Send(&rspb.RaftEnvelope{ClusterId: t.clusterID, Message: message}); err != nil { // Drop the broken connection so the next send reconnects. t.mu.Lock() if t.conns[storeID] == conn { @@ -89,10 +93,30 @@ func (t *ServerTransport) Send(message *raftpb.Message) error { return nil } +func (t *ServerTransport) ReplaceMembers(members map[uint64]string) { + members = cloneAddresses(members) + t.mu.Lock() + defer t.mu.Unlock() + + for id, conn := range t.conns { + oldAddress := t.members[id] + newAddress, exists := members[id] + if !exists || newAddress != oldAddress { + conn.Stop() + delete(t.conns, id) + } + } + t.members = members +} + func (t *ServerTransport) getConn(storeID uint64, addr string) (*raftConn, error) { t.mu.RLock() + currentAddress, memberExists := t.members[storeID] conn, ok := t.conns[storeID] t.mu.RUnlock() + if !memberExists || currentAddress != addr { + return nil, fmt.Errorf("address for store %d changed while connecting", storeID) + } if ok { return conn, nil } @@ -102,6 +126,11 @@ func (t *ServerTransport) getConn(storeID uint64, addr string) (*raftConn, error } t.mu.Lock() defer t.mu.Unlock() + currentAddress, memberExists = t.members[storeID] + if !memberExists || currentAddress != addr { + newConn.Stop() + return nil, fmt.Errorf("address for store %d changed while connecting", storeID) + } if conn, ok := t.conns[storeID]; ok { newConn.Stop() return conn, nil @@ -118,3 +147,11 @@ func (t *ServerTransport) Stop() { } t.conns = make(map[uint64]*raftConn) } + +func cloneAddresses(addresses map[uint64]string) map[uint64]string { + cloned := make(map[uint64]string, len(addresses)) + for id, address := range addresses { + cloned[id] = address + } + return cloned +} diff --git a/engine/storage/raft_storage/transport_test.go b/engine/storage/raft_storage/transport_test.go new file mode 100644 index 0000000..23304ee --- /dev/null +++ b/engine/storage/raft_storage/transport_test.go @@ -0,0 +1,170 @@ +package raft_storage + +import ( + "context" + "errors" + "io" + "net" + "strings" + "testing" + "time" + + "google.golang.org/grpc" + "google.golang.org/protobuf/proto" + + "github.com/Aetherance/kv/engine/config" + "github.com/Aetherance/kv/proto/pkg/clusterpb" + rspb "github.com/Aetherance/kv/proto/pkg/raft_serverpb" + "github.com/Aetherance/kv/proto/pkg/raftpb" +) + +func TestServerTransportRoutesEnvelopeAndRefreshesAddress(t *testing.T) { + firstAddress, firstReceived := startRecordingRaftServer(t) + secondAddress, secondReceived := startRecordingRaftServer(t) + members := map[uint64]string{2: firstAddress} + transport := NewServerTransport(42, members) + t.Cleanup(transport.Stop) + + // The transport owns its routing snapshot; caller mutation must not change it. + members[2] = secondAddress + first := &raftpb.Message{MsgType: raftpb.MessageType_MsgAppend, From: 1, To: 2, Term: 3} + if err := transport.Send(first); err != nil { + t.Fatalf("send to first address: %v", err) + } + assertEnvelope(t, receiveEnvelope(t, firstReceived), 42, first) + + transport.ReplaceMembers(map[uint64]string{2: secondAddress}) + second := &raftpb.Message{MsgType: raftpb.MessageType_MsgHeartbeat, From: 1, To: 2, Term: 4} + if err := transport.Send(second); err != nil { + t.Fatalf("send to replacement address: %v", err) + } + assertEnvelope(t, receiveEnvelope(t, secondReceived), 42, second) + + transport.ReplaceMembers(nil) + if err := transport.Send(second); err == nil || !strings.Contains(err.Error(), "no address for store 2") { + t.Fatalf("send to removed member error = %v", err) + } +} + +func TestAppliedMembershipRefreshesTransport(t *testing.T) { + store := startMembershipStore(t) + member := &clusterpb.Member{Id: 2, RaftAddress: "127.0.0.1:2"} + if err := store.proposeMembership(context.Background(), raftpb.ConfChangeType_AddLearnerNode, member); err != nil { + t.Fatalf("add learner: %v", err) + } + + store.transport.mu.RLock() + address, exists := store.transport.members[2] + store.transport.mu.RUnlock() + if !exists || address != member.RaftAddress { + t.Fatalf("transport route for member 2 = %q, %v", address, exists) + } +} + +func TestRaftRejectsEnvelopeFromAnotherCluster(t *testing.T) { + store := &RaftStorage{ + config: &config.Config{StoreID: 1}, + clusterID: 42, + } + stream := &fakeRaftServerStream{ + ctx: context.Background(), + envelopes: []*rspb.RaftEnvelope{{ + ClusterId: 99, + Message: &raftpb.Message{From: 2, To: 1}, + }}, + } + err := store.Raft(stream) + if err == nil || !strings.Contains(err.Error(), "cluster ID mismatch: got 99, expected 42") { + t.Fatalf("cluster mismatch error = %v", err) + } +} + +type recordingRaftService struct { + rspb.UnimplementedRaftServiceServer + received chan *rspb.RaftEnvelope +} + +func (s *recordingRaftService) Raft(stream rspb.RaftService_RaftServer) error { + for { + envelope, err := stream.Recv() + if err == io.EOF { + return stream.SendAndClose(&rspb.Done{}) + } + if err != nil { + return err + } + cloned := proto.Clone(envelope).(*rspb.RaftEnvelope) + select { + case s.received <- cloned: + case <-stream.Context().Done(): + return stream.Context().Err() + } + } +} + +func startRecordingRaftServer(t *testing.T) (string, <-chan *rspb.RaftEnvelope) { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + received := make(chan *rspb.RaftEnvelope, 4) + server := grpc.NewServer() + rspb.RegisterRaftServiceServer(server, &recordingRaftService{received: received}) + serveErr := make(chan error, 1) + go func() { serveErr <- server.Serve(listener) }() + t.Cleanup(func() { + server.Stop() + _ = listener.Close() + select { + case err := <-serveErr: + if err != nil && !errors.Is(err, grpc.ErrServerStopped) { + t.Errorf("serve: %v", err) + } + case <-time.After(time.Second): + t.Error("gRPC server did not stop") + } + }) + return listener.Addr().String(), received +} + +func receiveEnvelope(t *testing.T, received <-chan *rspb.RaftEnvelope) *rspb.RaftEnvelope { + t.Helper() + select { + case envelope := <-received: + return envelope + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for raft envelope") + return nil + } +} + +func assertEnvelope(t *testing.T, envelope *rspb.RaftEnvelope, clusterID uint64, message *raftpb.Message) { + t.Helper() + if envelope.ClusterId != clusterID || !proto.Equal(envelope.Message, message) { + t.Fatalf("envelope = %v, want cluster %d message %v", envelope, clusterID, message) + } +} + +type fakeRaftServerStream struct { + grpc.ServerStream + ctx context.Context + envelopes []*rspb.RaftEnvelope + response *rspb.Done +} + +func (s *fakeRaftServerStream) Context() context.Context { return s.ctx } + +func (s *fakeRaftServerStream) Recv() (*rspb.RaftEnvelope, error) { + if len(s.envelopes) == 0 { + return nil, io.EOF + } + envelope := s.envelopes[0] + s.envelopes = s.envelopes[1:] + return envelope, nil +} + +func (s *fakeRaftServerStream) SendAndClose(response *rspb.Done) error { + s.response = response + return nil +} diff --git a/proto/pkg/raft_serverpb/raft_serverpb.pb.go b/proto/pkg/raft_serverpb/raft_serverpb.pb.go index 97ea1b8..8353fb7 100644 --- a/proto/pkg/raft_serverpb/raft_serverpb.pb.go +++ b/proto/pkg/raft_serverpb/raft_serverpb.pb.go @@ -163,6 +163,60 @@ func (*Done) Descriptor() ([]byte, []int) { return file_raft_serverpb_proto_rawDescGZIP(), []int{2} } +// RaftEnvelope prevents messages from a different logical cluster from being +// accepted by a store that happens to reuse the same member ID. +type RaftEnvelope struct { + state protoimpl.MessageState `protogen:"open.v1"` + ClusterId uint64 `protobuf:"varint,1,opt,name=cluster_id,json=clusterId,proto3" json:"cluster_id,omitempty"` + Message *raftpb.Message `protobuf:"bytes,2,opt,name=message,proto3" json:"message,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *RaftEnvelope) Reset() { + *x = RaftEnvelope{} + mi := &file_raft_serverpb_proto_msgTypes[3] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *RaftEnvelope) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*RaftEnvelope) ProtoMessage() {} + +func (x *RaftEnvelope) ProtoReflect() protoreflect.Message { + mi := &file_raft_serverpb_proto_msgTypes[3] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use RaftEnvelope.ProtoReflect.Descriptor instead. +func (*RaftEnvelope) Descriptor() ([]byte, []int) { + return file_raft_serverpb_proto_rawDescGZIP(), []int{3} +} + +func (x *RaftEnvelope) GetClusterId() uint64 { + if x != nil { + return x.ClusterId + } + return 0 +} + +func (x *RaftEnvelope) GetMessage() *raftpb.Message { + if x != nil { + return x.Message + } + return nil +} + var File_raft_serverpb_proto protoreflect.FileDescriptor const file_raft_serverpb_proto_rawDesc = "" + @@ -174,9 +228,13 @@ const file_raft_serverpb_proto_rawDesc = "" + "\x10RaftSnapshotData\x12+\n" + "\x04data\x18\x03 \x03(\v2\x17.raft_serverpb.KeyValueR\x04data\x124\n" + "\acluster\x18\x04 \x01(\v2\x1a.clusterpb.ClusterMetadataR\aclusterJ\x04\b\x01\x10\x02J\x04\b\x02\x10\x03J\x04\b\x05\x10\x06R\x06regionR\tfile_sizeR\x04meta\"\x06\n" + - "\x04Done2?\n" + - "\vRaftService\x120\n" + - "\x04Raft\x12\x0f.raftpb.Message\x1a\x13.raft_serverpb.Done\"\x00(\x01B@Z>github.com/Aetherance/kv/proto/pkg/raft_serverpb;raft_serverpbb\x06proto3" + "\x04Done\"X\n" + + "\fRaftEnvelope\x12\x1d\n" + + "\n" + + "cluster_id\x18\x01 \x01(\x04R\tclusterId\x12)\n" + + "\amessage\x18\x02 \x01(\v2\x0f.raftpb.MessageR\amessage2K\n" + + "\vRaftService\x12<\n" + + "\x04Raft\x12\x1b.raft_serverpb.RaftEnvelope\x1a\x13.raft_serverpb.Done\"\x00(\x01B@Z>github.com/Aetherance/kv/proto/pkg/raft_serverpb;raft_serverpbb\x06proto3" var ( file_raft_serverpb_proto_rawDescOnce sync.Once @@ -190,24 +248,26 @@ func file_raft_serverpb_proto_rawDescGZIP() []byte { return file_raft_serverpb_proto_rawDescData } -var file_raft_serverpb_proto_msgTypes = make([]protoimpl.MessageInfo, 3) +var file_raft_serverpb_proto_msgTypes = make([]protoimpl.MessageInfo, 4) var file_raft_serverpb_proto_goTypes = []any{ (*KeyValue)(nil), // 0: raft_serverpb.KeyValue (*RaftSnapshotData)(nil), // 1: raft_serverpb.RaftSnapshotData (*Done)(nil), // 2: raft_serverpb.Done - (*clusterpb.ClusterMetadata)(nil), // 3: clusterpb.ClusterMetadata - (*raftpb.Message)(nil), // 4: raftpb.Message + (*RaftEnvelope)(nil), // 3: raft_serverpb.RaftEnvelope + (*clusterpb.ClusterMetadata)(nil), // 4: clusterpb.ClusterMetadata + (*raftpb.Message)(nil), // 5: raftpb.Message } var file_raft_serverpb_proto_depIdxs = []int32{ 0, // 0: raft_serverpb.RaftSnapshotData.data:type_name -> raft_serverpb.KeyValue - 3, // 1: raft_serverpb.RaftSnapshotData.cluster:type_name -> clusterpb.ClusterMetadata - 4, // 2: raft_serverpb.RaftService.Raft:input_type -> raftpb.Message - 2, // 3: raft_serverpb.RaftService.Raft:output_type -> raft_serverpb.Done - 3, // [3:4] is the sub-list for method output_type - 2, // [2:3] is the sub-list for method input_type - 2, // [2:2] is the sub-list for extension type_name - 2, // [2:2] is the sub-list for extension extendee - 0, // [0:2] is the sub-list for field type_name + 4, // 1: raft_serverpb.RaftSnapshotData.cluster:type_name -> clusterpb.ClusterMetadata + 5, // 2: raft_serverpb.RaftEnvelope.message:type_name -> raftpb.Message + 3, // 3: raft_serverpb.RaftService.Raft:input_type -> raft_serverpb.RaftEnvelope + 2, // 4: raft_serverpb.RaftService.Raft:output_type -> raft_serverpb.Done + 4, // [4:5] is the sub-list for method output_type + 3, // [3:4] is the sub-list for method input_type + 3, // [3:3] is the sub-list for extension type_name + 3, // [3:3] is the sub-list for extension extendee + 0, // [0:3] is the sub-list for field type_name } func init() { file_raft_serverpb_proto_init() } @@ -221,7 +281,7 @@ func file_raft_serverpb_proto_init() { GoPackagePath: reflect.TypeOf(x{}).PkgPath(), RawDescriptor: unsafe.Slice(unsafe.StringData(file_raft_serverpb_proto_rawDesc), len(file_raft_serverpb_proto_rawDesc)), NumEnums: 0, - NumMessages: 3, + NumMessages: 4, NumExtensions: 0, NumServices: 1, }, diff --git a/proto/pkg/raft_serverpb/raft_serverpb_grpc.pb.go b/proto/pkg/raft_serverpb/raft_serverpb_grpc.pb.go index f8a08c6..db9a793 100644 --- a/proto/pkg/raft_serverpb/raft_serverpb_grpc.pb.go +++ b/proto/pkg/raft_serverpb/raft_serverpb_grpc.pb.go @@ -8,7 +8,6 @@ package raft_serverpb import ( context "context" - raftpb "github.com/Aetherance/kv/proto/pkg/raftpb" grpc "google.golang.org/grpc" codes "google.golang.org/grpc/codes" status "google.golang.org/grpc/status" @@ -27,7 +26,7 @@ const ( // // For semantics around ctx use and closing/ending streaming RPCs, please refer to https://pkg.go.dev/google.golang.org/grpc/?tab=doc#ClientConn.NewStream. type RaftServiceClient interface { - Raft(ctx context.Context, opts ...grpc.CallOption) (grpc.ClientStreamingClient[raftpb.Message, Done], error) + Raft(ctx context.Context, opts ...grpc.CallOption) (grpc.ClientStreamingClient[RaftEnvelope, Done], error) } type raftServiceClient struct { @@ -38,24 +37,24 @@ func NewRaftServiceClient(cc grpc.ClientConnInterface) RaftServiceClient { return &raftServiceClient{cc} } -func (c *raftServiceClient) Raft(ctx context.Context, opts ...grpc.CallOption) (grpc.ClientStreamingClient[raftpb.Message, Done], error) { +func (c *raftServiceClient) Raft(ctx context.Context, opts ...grpc.CallOption) (grpc.ClientStreamingClient[RaftEnvelope, Done], error) { cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) stream, err := c.cc.NewStream(ctx, &RaftService_ServiceDesc.Streams[0], RaftService_Raft_FullMethodName, cOpts...) if err != nil { return nil, err } - x := &grpc.GenericClientStream[raftpb.Message, Done]{ClientStream: stream} + x := &grpc.GenericClientStream[RaftEnvelope, Done]{ClientStream: stream} return x, nil } // This type alias is provided for backwards compatibility with existing code that references the prior non-generic stream type by name. -type RaftService_RaftClient = grpc.ClientStreamingClient[raftpb.Message, Done] +type RaftService_RaftClient = grpc.ClientStreamingClient[RaftEnvelope, Done] // RaftServiceServer is the server API for RaftService service. // All implementations must embed UnimplementedRaftServiceServer // for forward compatibility. type RaftServiceServer interface { - Raft(grpc.ClientStreamingServer[raftpb.Message, Done]) error + Raft(grpc.ClientStreamingServer[RaftEnvelope, Done]) error mustEmbedUnimplementedRaftServiceServer() } @@ -66,7 +65,7 @@ type RaftServiceServer interface { // pointer dereference when methods are called. type UnimplementedRaftServiceServer struct{} -func (UnimplementedRaftServiceServer) Raft(grpc.ClientStreamingServer[raftpb.Message, Done]) error { +func (UnimplementedRaftServiceServer) Raft(grpc.ClientStreamingServer[RaftEnvelope, Done]) error { return status.Error(codes.Unimplemented, "method Raft not implemented") } func (UnimplementedRaftServiceServer) mustEmbedUnimplementedRaftServiceServer() {} @@ -91,11 +90,11 @@ func RegisterRaftServiceServer(s grpc.ServiceRegistrar, srv RaftServiceServer) { } func _RaftService_Raft_Handler(srv interface{}, stream grpc.ServerStream) error { - return srv.(RaftServiceServer).Raft(&grpc.GenericServerStream[raftpb.Message, Done]{ServerStream: stream}) + return srv.(RaftServiceServer).Raft(&grpc.GenericServerStream[RaftEnvelope, Done]{ServerStream: stream}) } // This type alias is provided for backwards compatibility with existing code that references the prior non-generic stream type by name. -type RaftService_RaftServer = grpc.ClientStreamingServer[raftpb.Message, Done] +type RaftService_RaftServer = grpc.ClientStreamingServer[RaftEnvelope, Done] // RaftService_ServiceDesc is the grpc.ServiceDesc for RaftService service. // It's only intended for direct use with grpc.RegisterService, diff --git a/proto/proto/raft_serverpb.proto b/proto/proto/raft_serverpb.proto index b52d35d..73e193b 100644 --- a/proto/proto/raft_serverpb.proto +++ b/proto/proto/raft_serverpb.proto @@ -21,6 +21,13 @@ message RaftSnapshotData { message Done {} +// RaftEnvelope prevents messages from a different logical cluster from being +// accepted by a store that happens to reuse the same member ID. +message RaftEnvelope { + uint64 cluster_id = 1; + raftpb.Message message = 2; +} + service RaftService { - rpc Raft(stream raftpb.Message) returns (Done) {} + rpc Raft(stream RaftEnvelope) returns (Done) {} }