feat: enhance swarm protocol with gossip and RPC message types, improve leader election logic, and add tests
This commit is contained in:
parent
4a76c91e5d
commit
8a441888f4
7 changed files with 321 additions and 61 deletions
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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},
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue