feat: enhance swarm protocol with gossip and RPC message types, improve leader election logic, and add tests

This commit is contained in:
Zhaoyikaiii 2026-02-24 10:17:23 +08:00
parent 4a76c91e5d
commit 8a441888f4
No known key found for this signature in database
GPG key ID: 14C8E2DD713BB8D6
7 changed files with 321 additions and 61 deletions

View file

@ -8,6 +8,42 @@ package swarm
import "time" import "time"
// GossipMessageType represents the type of a gossip protocol message.
type GossipMessageType string
const (
GossipTypePing GossipMessageType = "ping"
GossipTypePong GossipMessageType = "pong"
GossipTypeJoin GossipMessageType = "join"
GossipTypeUpdate GossipMessageType = "update"
GossipTypeSync GossipMessageType = "sync"
)
// RPCMessageType represents the type of an RPC message between nodes.
type RPCMessageType string
const (
RPCTypeHandoffRequest RPCMessageType = "handoff_request"
RPCTypeHandoffResponse RPCMessageType = "handoff_response"
RPCTypeSessionTransfer RPCMessageType = "session_transfer"
RPCTypeSessionTransferAck RPCMessageType = "session_transfer_ack"
)
// LoadTrend represents the direction of load change over time.
type LoadTrend string
const (
LoadTrendIncreasing LoadTrend = "increasing"
LoadTrendDecreasing LoadTrend = "decreasing"
LoadTrendStable LoadTrend = "stable"
)
// RPC message field keys used in map-based message encoding.
const (
MsgFieldType = "type"
MsgFieldPayload = "payload"
)
const ( const (
// Default values for configurable parameters // Default values for configurable parameters

View file

@ -282,7 +282,7 @@ func (ds *DiscoveryService) gossipLoop() {
// GossipMessage represents a gossip message. // GossipMessage represents a gossip message.
type GossipMessage struct { type GossipMessage struct {
Type string `json:"type"` // "ping", "pong", "join", "update" Type GossipMessageType `json:"type"` // ping, pong, join, update, sync
FromNode string `json:"from_node"` FromNode string `json:"from_node"`
SeqNum uint64 `json:"seq_num"` SeqNum uint64 `json:"seq_num"`
Timestamp int64 `json:"timestamp"` Timestamp int64 `json:"timestamp"`
@ -312,15 +312,15 @@ func (ds *DiscoveryService) handleGossip(data []byte, addr *net.UDPAddr) {
} }
switch msg.Type { switch msg.Type {
case "ping": case GossipTypePing:
ds.handlePing(msg, addr) ds.handlePing(msg, addr)
case "pong": case GossipTypePong:
ds.handlePong(msg) ds.handlePong(msg)
case "join": case GossipTypeJoin:
ds.handleJoin(msg, addr) ds.handleJoin(msg, addr)
case "update": case GossipTypeUpdate:
ds.handleUpdate(msg) ds.handleUpdate(msg)
case "sync": case GossipTypeSync:
ds.handleSync(msg, addr) ds.handleSync(msg, addr)
} }
} }
@ -329,7 +329,7 @@ func (ds *DiscoveryService) handleGossip(data []byte, addr *net.UDPAddr) {
func (ds *DiscoveryService) handlePing(msg GossipMessage, addr *net.UDPAddr) { func (ds *DiscoveryService) handlePing(msg GossipMessage, addr *net.UDPAddr) {
// Respond with pong // Respond with pong
pong := GossipMessage{ pong := GossipMessage{
Type: "pong", Type: GossipTypePong,
FromNode: ds.localNode.ID, FromNode: ds.localNode.ID,
Timestamp: time.Now().UnixNano(), Timestamp: time.Now().UnixNano(),
} }
@ -372,7 +372,7 @@ func (ds *DiscoveryService) handleJoin(msg GossipMessage, addr *net.UDPAddr) {
} }
response := GossipMessage{ response := GossipMessage{
Type: "sync", Type: GossipTypeSync,
FromNode: ds.localNode.ID, FromNode: ds.localNode.ID,
Timestamp: time.Now().UnixNano(), Timestamp: time.Now().UnixNano(),
Nodes: nodes, Nodes: nodes,
@ -432,7 +432,7 @@ func (ds *DiscoveryService) broadcastUpdate() {
} }
msg := GossipMessage{ msg := GossipMessage{
Type: "update", Type: GossipTypeUpdate,
FromNode: ds.localNode.ID, FromNode: ds.localNode.ID,
SeqNum: ds.seqNum, SeqNum: ds.seqNum,
Timestamp: time.Now().UnixNano(), Timestamp: time.Now().UnixNano(),
@ -479,7 +479,7 @@ func (ds *DiscoveryService) sendJoin(ctx context.Context, addr string) error {
} }
msg := GossipMessage{ msg := GossipMessage{
Type: "join", Type: GossipTypeJoin,
FromNode: ds.localNode.ID, FromNode: ds.localNode.ID,
Timestamp: time.Now().UnixNano(), Timestamp: time.Now().UnixNano(),
Nodes: []*NodeInfo{ds.localNode}, Nodes: []*NodeInfo{ds.localNode},

View file

@ -291,8 +291,8 @@ func (hc *HandoffCoordinator) findTargetNode(req *HandoffRequest) (*NodeWithStat
func (hc *HandoffCoordinator) sendHandoffRequest(req *HandoffRequest, target *NodeWithState) error { func (hc *HandoffCoordinator) sendHandoffRequest(req *HandoffRequest, target *NodeWithState) error {
// Handoff message type // Handoff message type
msg := map[string]any{ msg := map[string]any{
"type": "handoff_request", MsgFieldType: RPCTypeHandoffRequest,
"payload": req, MsgFieldPayload: req,
} }
data, err := json.Marshal(msg) data, err := json.Marshal(msg)
@ -369,12 +369,15 @@ func (hc *HandoffCoordinator) handleMessage(data []byte, addr *net.UDPAddr) {
return return
} }
msgType, _ := msg["type"].(string) msgType, ok := msg[MsgFieldType].(string)
if !ok {
return
}
switch msgType { switch RPCMessageType(msgType) {
case "handoff_request": case RPCTypeHandoffRequest:
hc.handleHandoffRequest(data, addr) hc.handleHandoffRequest(data, addr)
case "handoff_response": case RPCTypeHandoffResponse:
hc.handleHandoffResponse(data) hc.handleHandoffResponse(data)
} }
} }
@ -386,7 +389,7 @@ func (hc *HandoffCoordinator) handleHandoffRequest(data []byte, addr *net.UDPAdd
return return
} }
payloadData, _ := json.Marshal(msg["payload"]) payloadData, _ := json.Marshal(msg[MsgFieldPayload])
var req HandoffRequest var req HandoffRequest
if err := json.Unmarshal(payloadData, &req); err != nil { if err := json.Unmarshal(payloadData, &req); err != nil {
return return
@ -414,8 +417,8 @@ func (hc *HandoffCoordinator) handleHandoffRequest(data []byte, addr *net.UDPAdd
// Send response // Send response
respMsg := map[string]any{ respMsg := map[string]any{
"type": "handoff_response", MsgFieldType: RPCTypeHandoffResponse,
"payload": response, MsgFieldPayload: response,
} }
respData, _ := json.Marshal(respMsg) respData, _ := json.Marshal(respMsg)
@ -442,7 +445,7 @@ func (hc *HandoffCoordinator) handleHandoffResponse(data []byte) {
return return
} }
payloadData, _ := json.Marshal(msg["payload"]) payloadData, _ := json.Marshal(msg[MsgFieldPayload])
var resp HandoffResponse var resp HandoffResponse
if err := json.Unmarshal(payloadData, &resp); err != nil { if err := json.Unmarshal(payloadData, &resp); err != nil {
return return

View file

@ -18,6 +18,7 @@ import (
type LeaderElection struct { type LeaderElection struct {
localNodeID string localNodeID string
membership *MembershipManager membership *MembershipManager
config LeaderElectionConfig
mu sync.RWMutex mu sync.RWMutex
currentLeader string currentLeader string
@ -28,10 +29,11 @@ type LeaderElection struct {
} }
// NewLeaderElection creates a new leader election instance. // NewLeaderElection creates a new leader election instance.
func NewLeaderElection(nodeID string, membership *MembershipManager) *LeaderElection { func NewLeaderElection(nodeID string, membership *MembershipManager, config LeaderElectionConfig) *LeaderElection {
return &LeaderElection{ return &LeaderElection{
localNodeID: nodeID, localNodeID: nodeID,
membership: membership, membership: membership,
config: config,
leaderChangeCh: make(chan string, 10), leaderChangeCh: make(chan string, 10),
stopCh: make(chan struct{}), stopCh: make(chan struct{}),
} }
@ -72,7 +74,11 @@ func (le *LeaderElection) LeaderChanges() <-chan string {
// electionChecker periodically checks if we should become leader. // electionChecker periodically checks if we should become leader.
func (le *LeaderElection) electionChecker() { func (le *LeaderElection) electionChecker() {
ticker := time.NewTicker(time.Second * 5) interval := le.config.ElectionInterval.Duration
if interval <= 0 {
interval = 5 * time.Second
}
ticker := time.NewTicker(interval)
defer ticker.Stop() defer ticker.Stop()
for { for {
@ -91,17 +97,24 @@ func (le *LeaderElection) checkElection() {
defer le.mu.Unlock() defer le.mu.Unlock()
members := le.membership.GetMembers() members := le.membership.GetMembers()
if len(members) == 0 {
// No other members, we become leader // Filter to only alive nodes for leader candidacy
aliveMembers := make([]*NodeWithState, 0, len(members))
for _, m := range members {
if m.State != nil && m.State.Status == NodeStatusAlive {
aliveMembers = append(aliveMembers, m)
}
}
if len(aliveMembers) == 0 {
// No alive members in view (including self not yet registered), become leader as fallback
le.becomeLeader() le.becomeLeader()
return return
} }
// Find the node with the lowest ID (simple deterministic leader selection) // Find the node with the lowest ID (simple deterministic leader selection)
var candidateID string candidateID := le.localNodeID
candidateID = le.localNodeID for _, m := range aliveMembers {
for _, m := range members {
if m.Node.ID < candidateID { if m.Node.ID < candidateID {
candidateID = m.Node.ID candidateID = m.Node.ID
} }
@ -109,7 +122,6 @@ func (le *LeaderElection) checkElection() {
// Update current leader // Update current leader
if le.currentLeader != candidateID { if le.currentLeader != candidateID {
oldLeader := le.currentLeader
le.currentLeader = candidateID le.currentLeader = candidateID
if candidateID == le.localNodeID { if candidateID == le.localNodeID {
@ -117,14 +129,6 @@ func (le *LeaderElection) checkElection() {
} else { } else {
le.becomeFollower() le.becomeFollower()
} }
logger.InfoCF("swarm", "Leader changed", map[string]any{"old_leader": oldLeader, "new_leader": candidateID})
// Notify followers of leader change
select {
case le.leaderChangeCh <- candidateID:
default:
}
} }
} }
@ -134,25 +138,42 @@ func (le *LeaderElection) becomeLeader() {
le.isLeader = true le.isLeader = true
logger.InfoCF("swarm", "This node is now the leader", map[string]any{"node_id": le.localNodeID}) logger.InfoCF("swarm", "This node is now the leader", map[string]any{"node_id": le.localNodeID})
// Notify listeners // Notify listeners (non-blocking)
select { select {
case le.leaderChangeCh <- le.localNodeID: case le.leaderChangeCh <- le.localNodeID:
default: default:
logger.WarnC("swarm", "Leader change notification dropped, channel full")
} }
} }
} }
// becomeFollower marks this node as a follower. // becomeFollower marks this node as a follower.
func (le *LeaderElection) becomeFollower() { func (le *LeaderElection) becomeFollower() {
if le.isLeader { wasLeader := le.isLeader
le.isLeader = false le.isLeader = false
logger.InfoCF("swarm", "This node is now a follower", map[string]any{"node_id": le.localNodeID})
if wasLeader {
logger.InfoCF("swarm", "This node is now a follower", map[string]any{
"node_id": le.localNodeID,
"new_leader": le.currentLeader,
})
// Notify listeners of leader change (non-blocking)
select {
case le.leaderChangeCh <- le.currentLeader:
default:
logger.WarnC("swarm", "Leader change notification dropped, channel full")
}
} }
} }
// leaderMonitor monitors if the current leader is still alive. // leaderMonitor monitors if the current leader is still alive.
func (le *LeaderElection) leaderMonitor() { func (le *LeaderElection) leaderMonitor() {
ticker := time.NewTicker(time.Second * 10) interval := le.config.LeaderHeartbeatTimeout.Duration
if interval <= 0 {
interval = 10 * time.Second
}
ticker := time.NewTicker(interval)
defer ticker.Stop() defer ticker.Stop()
for { for {
@ -176,11 +197,20 @@ func (le *LeaderElection) monitorLeader() {
return return
} }
// Check if leader is still in the membership // Check if leader is still in the membership and healthy
if _, exists := le.membership.GetNode(leaderID); !exists { needReelection := false
node, exists := le.membership.GetNode(leaderID)
if !exists {
logger.WarnCF("swarm", "Leader no longer in membership, triggering reelection", logger.WarnCF("swarm", "Leader no longer in membership, triggering reelection",
map[string]any{"leader_id": leaderID}) map[string]any{"leader_id": leaderID})
// Trigger reelection by clearing current leader needReelection = true
} else if node.State != nil && node.State.Status != NodeStatusAlive {
logger.WarnCF("swarm", "Leader is no longer alive, triggering reelection",
map[string]any{"leader_id": leaderID, "status": node.State.Status})
needReelection = true
}
if needReelection {
le.mu.Lock() le.mu.Lock()
le.currentLeader = "" le.currentLeader = ""
le.mu.Unlock() le.mu.Unlock()

View file

@ -234,13 +234,13 @@ func (lm *LoadMonitor) OnThreshold(callback func(float64)) {
lm.onThreshold = append(lm.onThreshold, callback) lm.onThreshold = append(lm.onThreshold, callback)
} }
// GetTrend returns the load trend: "increasing", "decreasing", or "stable". // GetTrend returns the load trend.
func (lm *LoadMonitor) GetTrend() string { func (lm *LoadMonitor) GetTrend() LoadTrend {
lm.mu.RLock() lm.mu.RLock()
defer lm.mu.RUnlock() defer lm.mu.RUnlock()
if len(lm.samples) < 3 { if len(lm.samples) < 3 {
return "stable" return LoadTrendStable
} }
// Simple linear regression to detect trend // Simple linear regression to detect trend
@ -258,11 +258,11 @@ func (lm *LoadMonitor) GetTrend() string {
slope := (n*sumXY - sumX*sumY) / (n * (n - 1) * (2*n - 1) / 6) slope := (n*sumXY - sumX*sumY) / (n * (n - 1) * (2*n - 1) / 6)
if slope > TrendIncreasingThreshold { if slope > TrendIncreasingThreshold {
return "increasing" return LoadTrendIncreasing
} else if slope < TrendDecreasingThreshold { } else if slope < TrendDecreasingThreshold {
return "decreasing" return LoadTrendDecreasing
} }
return "stable" return LoadTrendStable
} }
// Helper functions for normalization // Helper functions for normalization

View file

@ -134,8 +134,8 @@ func (st *SessionTransfer) TransferSession(ctx context.Context, targetNode *Node
// sendTransfer sends a transfer message to a target node. // sendTransfer sends a transfer message to a target node.
func (st *SessionTransfer) sendTransfer(targetNode *NodeInfo, payload *TransferPayload) error { func (st *SessionTransfer) sendTransfer(targetNode *NodeInfo, payload *TransferPayload) error {
msg := map[string]any{ msg := map[string]any{
"type": "session_transfer", MsgFieldType: RPCTypeSessionTransfer,
"payload": payload, MsgFieldPayload: payload,
} }
data, err := json.Marshal(msg) data, err := json.Marshal(msg)
@ -157,8 +157,8 @@ func (st *SessionTransfer) sendTransfer(targetNode *NodeInfo, payload *TransferP
// SendAck sends an acknowledgment for a received transfer. // SendAck sends an acknowledgment for a received transfer.
func (st *SessionTransfer) SendAck(targetNode *NodeInfo, transferID string, accepted bool) error { func (st *SessionTransfer) SendAck(targetNode *NodeInfo, transferID string, accepted bool) error {
msg := map[string]any{ msg := map[string]any{
"type": "session_transfer_ack", MsgFieldType: RPCTypeSessionTransferAck,
"payload": map[string]any{ MsgFieldPayload: map[string]any{
"transfer_id": transferID, "transfer_id": transferID,
"accepted": accepted, "accepted": accepted,
"node_id": st.localNode.ID, "node_id": st.localNode.ID,
@ -207,12 +207,15 @@ func (st *SessionTransfer) handleMessage(data []byte, addr *net.UDPAddr) {
return return
} }
msgType, _ := msg["type"].(string) msgType, ok := msg[MsgFieldType].(string)
if !ok {
return
}
switch msgType { switch RPCMessageType(msgType) {
case "session_transfer": case RPCTypeSessionTransfer:
st.handleTransfer(data, addr) st.handleTransfer(data, addr)
case "session_transfer_ack": case RPCTypeSessionTransferAck:
st.handleTransferAck(data) st.handleTransferAck(data)
} }
} }
@ -224,7 +227,7 @@ func (st *SessionTransfer) handleTransfer(data []byte, addr *net.UDPAddr) {
return return
} }
payloadData, _ := json.Marshal(msg["payload"]) payloadData, _ := json.Marshal(msg[MsgFieldPayload])
var payload TransferPayload var payload TransferPayload
if err := json.Unmarshal(payloadData, &payload); err != nil { if err := json.Unmarshal(payloadData, &payload); err != nil {
return return
@ -270,7 +273,10 @@ func (st *SessionTransfer) handleTransferAck(data []byte) {
return return
} }
payload, _ := msg["payload"].(map[string]any) payload, ok := msg[MsgFieldPayload].(map[string]any)
if !ok {
return
}
transferID, _ := payload["transfer_id"].(string) transferID, _ := payload["transfer_id"].(string)
accepted, _ := payload["accepted"].(bool) accepted, _ := payload["accepted"].(bool)

View file

@ -422,3 +422,188 @@ func TestSessionTransfer(t *testing.T) {
assert.Equal(t, 0, len(transfers)) assert.Equal(t, 0, len(transfers))
}) })
} }
func TestLeaderElection(t *testing.T) {
// Helper to create a discovery service with membership for testing
newTestDiscovery := func(t *testing.T, nodeID string, port int) *DiscoveryService {
t.Helper()
cfg := &Config{
NodeID: nodeID,
BindAddr: "127.0.0.1",
BindPort: port,
RPC: RPCConfig{Port: port + 1},
Discovery: DiscoveryConfig{
GossipInterval: Duration{100 * time.Millisecond},
NodeTimeout: Duration{500 * time.Millisecond},
DeadNodeTimeout: Duration{2 * time.Second},
},
}
ds, err := NewDiscoveryService(cfg)
require.NoError(t, err)
// Register the local node into membership so checkElection sees it
ds.membership.UpdateNode(ds.localNode)
return ds
}
defaultConfig := LeaderElectionConfig{
Enabled: true,
ElectionInterval: Duration{100 * time.Millisecond},
LeaderHeartbeatTimeout: Duration{200 * time.Millisecond},
}
t.Run("SingleNodeBecomesLeader", func(t *testing.T) {
ds := newTestDiscovery(t, "node-a", 18100)
le := NewLeaderElection(ds.localNode.ID, ds.membership, defaultConfig)
le.checkElection()
assert.True(t, le.IsLeader())
assert.Equal(t, "node-a", le.GetLeader())
})
t.Run("LowestIDBecomesLeader", func(t *testing.T) {
ds := newTestDiscovery(t, "node-c", 18110)
// Add two remote alive nodes
ds.membership.UpdateNode(&NodeInfo{
ID: "node-a",
Addr: "192.168.1.1",
Port: 7947,
LoadScore: 0.5,
Timestamp: time.Now().UnixNano(),
})
ds.membership.UpdateNode(&NodeInfo{
ID: "node-b",
Addr: "192.168.1.2",
Port: 7947,
LoadScore: 0.3,
Timestamp: time.Now().UnixNano(),
})
le := NewLeaderElection("node-c", ds.membership, defaultConfig)
le.checkElection()
// node-a has the lowest ID
assert.False(t, le.IsLeader())
assert.Equal(t, "node-a", le.GetLeader())
})
t.Run("DeadNodeNotElected", func(t *testing.T) {
ds := newTestDiscovery(t, "node-c", 18120)
// Add node-a (lowest ID) but mark it dead
ds.membership.UpdateNode(&NodeInfo{
ID: "node-a",
Addr: "192.168.1.1",
Port: 7947,
LoadScore: 0.2,
Timestamp: time.Now().UnixNano(),
})
ds.membership.MarkDead("node-a")
// Add node-b as alive
ds.membership.UpdateNode(&NodeInfo{
ID: "node-b",
Addr: "192.168.1.2",
Port: 7947,
LoadScore: 0.3,
Timestamp: time.Now().UnixNano(),
})
le := NewLeaderElection("node-c", ds.membership, defaultConfig)
le.checkElection()
// node-a is dead, so node-b (next lowest alive ID) should be leader
assert.Equal(t, "node-b", le.GetLeader())
assert.False(t, le.IsLeader())
})
t.Run("SuspectNodeNotElected", func(t *testing.T) {
ds := newTestDiscovery(t, "node-c", 18130)
// Add node-a (lowest ID) but mark it suspect
ds.membership.UpdateNode(&NodeInfo{
ID: "node-a",
Addr: "192.168.1.1",
Port: 7947,
LoadScore: 0.2,
Timestamp: time.Now().UnixNano(),
})
ds.membership.MarkSuspect("node-a")
le := NewLeaderElection("node-c", ds.membership, defaultConfig)
le.checkElection()
// node-a is suspect, local node-c should be leader
assert.Equal(t, "node-c", le.GetLeader())
assert.True(t, le.IsLeader())
})
t.Run("LeaderReelectionOnLeaderDeath", func(t *testing.T) {
ds := newTestDiscovery(t, "node-b", 18140)
// Add node-a as alive leader
ds.membership.UpdateNode(&NodeInfo{
ID: "node-a",
Addr: "192.168.1.1",
Port: 7947,
LoadScore: 0.2,
Timestamp: time.Now().UnixNano(),
})
le := NewLeaderElection("node-b", ds.membership, defaultConfig)
le.checkElection()
assert.Equal(t, "node-a", le.GetLeader())
assert.False(t, le.IsLeader())
// Now mark node-a as dead
ds.membership.MarkDead("node-a")
// monitorLeader should detect and trigger reelection
le.monitorLeader()
// node-b should now be leader
assert.Equal(t, "node-b", le.GetLeader())
assert.True(t, le.IsLeader())
})
t.Run("NoDoubleNotification", func(t *testing.T) {
ds := newTestDiscovery(t, "node-a", 18150)
le := NewLeaderElection("node-a", ds.membership, defaultConfig)
// First election — should produce exactly one notification
le.checkElection()
assert.True(t, le.IsLeader())
// Drain the channel
count := 0
for {
select {
case <-le.LeaderChanges():
count++
default:
goto done
}
}
done:
assert.Equal(t, 1, count, "should receive exactly one leader change notification")
})
t.Run("GetState", func(t *testing.T) {
ds := newTestDiscovery(t, "node-a", 18160)
le := NewLeaderElection("node-a", ds.membership, defaultConfig)
le.checkElection()
state := le.GetState()
assert.Equal(t, "node-a", state.LeaderID)
assert.True(t, state.IsLeader)
assert.GreaterOrEqual(t, state.MemberCount, 1)
})
}