diff --git a/pkg/swarm/constants.go b/pkg/swarm/constants.go index f7a670a43..1346d5f9e 100644 --- a/pkg/swarm/constants.go +++ b/pkg/swarm/constants.go @@ -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 diff --git a/pkg/swarm/discovery.go b/pkg/swarm/discovery.go index 6ca6913e7..bc7e40726 100644 --- a/pkg/swarm/discovery.go +++ b/pkg/swarm/discovery.go @@ -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}, diff --git a/pkg/swarm/handoff.go b/pkg/swarm/handoff.go index 3cb35ffca..08519942b 100644 --- a/pkg/swarm/handoff.go +++ b/pkg/swarm/handoff.go @@ -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 diff --git a/pkg/swarm/leader_election.go b/pkg/swarm/leader_election.go index 66a384c70..58dab483d 100644 --- a/pkg/swarm/leader_election.go +++ b/pkg/swarm/leader_election.go @@ -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() diff --git a/pkg/swarm/load_monitor.go b/pkg/swarm/load_monitor.go index fdf130faf..250060d03 100644 --- a/pkg/swarm/load_monitor.go +++ b/pkg/swarm/load_monitor.go @@ -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 diff --git a/pkg/swarm/session_transfer.go b/pkg/swarm/session_transfer.go index f89417320..9aab994f0 100644 --- a/pkg/swarm/session_transfer.go +++ b/pkg/swarm/session_transfer.go @@ -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) diff --git a/pkg/swarm/swarm_test.go b/pkg/swarm/swarm_test.go index 3708458d1..90b4862db 100644 --- a/pkg/swarm/swarm_test.go +++ b/pkg/swarm/swarm_test.go @@ -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) + }) +}