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"
// 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 (
// Default values for configurable parameters

View file

@ -282,7 +282,7 @@ func (ds *DiscoveryService) gossipLoop() {
// GossipMessage represents a gossip message.
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"`
SeqNum uint64 `json:"seq_num"`
Timestamp int64 `json:"timestamp"`
@ -312,15 +312,15 @@ func (ds *DiscoveryService) handleGossip(data []byte, addr *net.UDPAddr) {
}
switch msg.Type {
case "ping":
case GossipTypePing:
ds.handlePing(msg, addr)
case "pong":
case GossipTypePong:
ds.handlePong(msg)
case "join":
case GossipTypeJoin:
ds.handleJoin(msg, addr)
case "update":
case GossipTypeUpdate:
ds.handleUpdate(msg)
case "sync":
case GossipTypeSync:
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) {
// Respond with pong
pong := GossipMessage{
Type: "pong",
Type: GossipTypePong,
FromNode: ds.localNode.ID,
Timestamp: time.Now().UnixNano(),
}
@ -372,7 +372,7 @@ func (ds *DiscoveryService) handleJoin(msg GossipMessage, addr *net.UDPAddr) {
}
response := GossipMessage{
Type: "sync",
Type: GossipTypeSync,
FromNode: ds.localNode.ID,
Timestamp: time.Now().UnixNano(),
Nodes: nodes,
@ -432,7 +432,7 @@ func (ds *DiscoveryService) broadcastUpdate() {
}
msg := GossipMessage{
Type: "update",
Type: GossipTypeUpdate,
FromNode: ds.localNode.ID,
SeqNum: ds.seqNum,
Timestamp: time.Now().UnixNano(),
@ -479,7 +479,7 @@ func (ds *DiscoveryService) sendJoin(ctx context.Context, addr string) error {
}
msg := GossipMessage{
Type: "join",
Type: GossipTypeJoin,
FromNode: ds.localNode.ID,
Timestamp: time.Now().UnixNano(),
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 {
// Handoff message type
msg := map[string]any{
"type": "handoff_request",
"payload": req,
MsgFieldType: RPCTypeHandoffRequest,
MsgFieldPayload: req,
}
data, err := json.Marshal(msg)
@ -369,12 +369,15 @@ func (hc *HandoffCoordinator) handleMessage(data []byte, addr *net.UDPAddr) {
return
}
msgType, _ := msg["type"].(string)
msgType, ok := msg[MsgFieldType].(string)
if !ok {
return
}
switch msgType {
case "handoff_request":
switch RPCMessageType(msgType) {
case RPCTypeHandoffRequest:
hc.handleHandoffRequest(data, addr)
case "handoff_response":
case RPCTypeHandoffResponse:
hc.handleHandoffResponse(data)
}
}
@ -386,7 +389,7 @@ func (hc *HandoffCoordinator) handleHandoffRequest(data []byte, addr *net.UDPAdd
return
}
payloadData, _ := json.Marshal(msg["payload"])
payloadData, _ := json.Marshal(msg[MsgFieldPayload])
var req HandoffRequest
if err := json.Unmarshal(payloadData, &req); err != nil {
return
@ -414,8 +417,8 @@ func (hc *HandoffCoordinator) handleHandoffRequest(data []byte, addr *net.UDPAdd
// Send response
respMsg := map[string]any{
"type": "handoff_response",
"payload": response,
MsgFieldType: RPCTypeHandoffResponse,
MsgFieldPayload: response,
}
respData, _ := json.Marshal(respMsg)
@ -442,7 +445,7 @@ func (hc *HandoffCoordinator) handleHandoffResponse(data []byte) {
return
}
payloadData, _ := json.Marshal(msg["payload"])
payloadData, _ := json.Marshal(msg[MsgFieldPayload])
var resp HandoffResponse
if err := json.Unmarshal(payloadData, &resp); err != nil {
return

View file

@ -18,6 +18,7 @@ import (
type LeaderElection struct {
localNodeID string
membership *MembershipManager
config LeaderElectionConfig
mu sync.RWMutex
currentLeader string
@ -28,10 +29,11 @@ type LeaderElection struct {
}
// 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{
localNodeID: nodeID,
membership: membership,
config: config,
leaderChangeCh: make(chan string, 10),
stopCh: make(chan struct{}),
}
@ -72,7 +74,11 @@ func (le *LeaderElection) LeaderChanges() <-chan string {
// electionChecker periodically checks if we should become leader.
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()
for {
@ -91,17 +97,24 @@ func (le *LeaderElection) checkElection() {
defer le.mu.Unlock()
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()
return
}
// Find the node with the lowest ID (simple deterministic leader selection)
var candidateID string
candidateID = le.localNodeID
for _, m := range members {
candidateID := le.localNodeID
for _, m := range aliveMembers {
if m.Node.ID < candidateID {
candidateID = m.Node.ID
}
@ -109,7 +122,6 @@ func (le *LeaderElection) checkElection() {
// Update current leader
if le.currentLeader != candidateID {
oldLeader := le.currentLeader
le.currentLeader = candidateID
if candidateID == le.localNodeID {
@ -117,14 +129,6 @@ func (le *LeaderElection) checkElection() {
} else {
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
logger.InfoCF("swarm", "This node is now the leader", map[string]any{"node_id": le.localNodeID})
// Notify listeners
// Notify listeners (non-blocking)
select {
case le.leaderChangeCh <- le.localNodeID:
default:
logger.WarnC("swarm", "Leader change notification dropped, channel full")
}
}
}
// becomeFollower marks this node as a follower.
func (le *LeaderElection) becomeFollower() {
if le.isLeader {
le.isLeader = false
logger.InfoCF("swarm", "This node is now a follower", map[string]any{"node_id": le.localNodeID})
wasLeader := le.isLeader
le.isLeader = false
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.
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()
for {
@ -176,11 +197,20 @@ func (le *LeaderElection) monitorLeader() {
return
}
// Check if leader is still in the membership
if _, exists := le.membership.GetNode(leaderID); !exists {
// Check if leader is still in the membership and healthy
needReelection := false
node, exists := le.membership.GetNode(leaderID)
if !exists {
logger.WarnCF("swarm", "Leader no longer in membership, triggering reelection",
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.currentLeader = ""
le.mu.Unlock()

View file

@ -234,13 +234,13 @@ func (lm *LoadMonitor) OnThreshold(callback func(float64)) {
lm.onThreshold = append(lm.onThreshold, callback)
}
// GetTrend returns the load trend: "increasing", "decreasing", or "stable".
func (lm *LoadMonitor) GetTrend() string {
// GetTrend returns the load trend.
func (lm *LoadMonitor) GetTrend() LoadTrend {
lm.mu.RLock()
defer lm.mu.RUnlock()
if len(lm.samples) < 3 {
return "stable"
return LoadTrendStable
}
// 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)
if slope > TrendIncreasingThreshold {
return "increasing"
return LoadTrendIncreasing
} else if slope < TrendDecreasingThreshold {
return "decreasing"
return LoadTrendDecreasing
}
return "stable"
return LoadTrendStable
}
// 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.
func (st *SessionTransfer) sendTransfer(targetNode *NodeInfo, payload *TransferPayload) error {
msg := map[string]any{
"type": "session_transfer",
"payload": payload,
MsgFieldType: RPCTypeSessionTransfer,
MsgFieldPayload: payload,
}
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.
func (st *SessionTransfer) SendAck(targetNode *NodeInfo, transferID string, accepted bool) error {
msg := map[string]any{
"type": "session_transfer_ack",
"payload": map[string]any{
MsgFieldType: RPCTypeSessionTransferAck,
MsgFieldPayload: map[string]any{
"transfer_id": transferID,
"accepted": accepted,
"node_id": st.localNode.ID,
@ -207,12 +207,15 @@ func (st *SessionTransfer) handleMessage(data []byte, addr *net.UDPAddr) {
return
}
msgType, _ := msg["type"].(string)
msgType, ok := msg[MsgFieldType].(string)
if !ok {
return
}
switch msgType {
case "session_transfer":
switch RPCMessageType(msgType) {
case RPCTypeSessionTransfer:
st.handleTransfer(data, addr)
case "session_transfer_ack":
case RPCTypeSessionTransferAck:
st.handleTransferAck(data)
}
}
@ -224,7 +227,7 @@ func (st *SessionTransfer) handleTransfer(data []byte, addr *net.UDPAddr) {
return
}
payloadData, _ := json.Marshal(msg["payload"])
payloadData, _ := json.Marshal(msg[MsgFieldPayload])
var payload TransferPayload
if err := json.Unmarshal(payloadData, &payload); err != nil {
return
@ -270,7 +273,10 @@ func (st *SessionTransfer) handleTransferAck(data []byte) {
return
}
payload, _ := msg["payload"].(map[string]any)
payload, ok := msg[MsgFieldPayload].(map[string]any)
if !ok {
return
}
transferID, _ := payload["transfer_id"].(string)
accepted, _ := payload["accepted"].(bool)

View file

@ -422,3 +422,188 @@ func TestSessionTransfer(t *testing.T) {
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)
})
}