Skip to content
Draft
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
6 changes: 6 additions & 0 deletions engine/storage/raft_storage/cluster_metadata_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
13 changes: 13 additions & 0 deletions engine/storage/raft_storage/membership.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
4 changes: 4 additions & 0 deletions engine/storage/raft_storage/raft.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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)
}
Expand Down
25 changes: 16 additions & 9 deletions engine/storage/raft_storage/raft_storage.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
}
}
Expand Down Expand Up @@ -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)
}
}
Expand Down
61 changes: 49 additions & 12 deletions engine/storage/raft_storage/transport.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,15 +8,14 @@ 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"
)

// 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
}
Expand All @@ -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)
Expand All @@ -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),
}
}

Expand All @@ -68,15 +70,17 @@ 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)
}
conn, err := t.getConn(storeID, addr)
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 {
Expand All @@ -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
}
Expand All @@ -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
Expand All @@ -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
}
Loading