374 lines
12 KiB
Go
374 lines
12 KiB
Go
package control
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"net/http/httptest"
|
|
"net/netip"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/coder/websocket"
|
|
"github.com/coder/websocket/wsjson"
|
|
|
|
"remlink/internal/model"
|
|
"remlink/internal/protocol"
|
|
)
|
|
|
|
type memoryNodes struct {
|
|
mu sync.Mutex
|
|
nodes map[string]model.Node
|
|
tokens map[string]string
|
|
events []model.EventLog
|
|
}
|
|
|
|
type statusChangeRecorder struct {
|
|
mu sync.Mutex
|
|
changes []struct {
|
|
node model.Node
|
|
status model.NodeStatus
|
|
}
|
|
}
|
|
|
|
func (*statusChangeRecorder) HandleControl(context.Context, model.Node, protocol.ControlEnvelope) error {
|
|
return nil
|
|
}
|
|
|
|
func (r *statusChangeRecorder) HandleNodeStatusChange(_ context.Context, node model.Node, status model.NodeStatus) error {
|
|
r.mu.Lock()
|
|
r.changes = append(r.changes, struct {
|
|
node model.Node
|
|
status model.NodeStatus
|
|
}{node: node, status: status})
|
|
r.mu.Unlock()
|
|
return nil
|
|
}
|
|
|
|
func (m *memoryNodes) AppendEvent(_ context.Context, event model.EventLog) error {
|
|
m.mu.Lock()
|
|
m.events = append(m.events, event)
|
|
m.mu.Unlock()
|
|
return nil
|
|
}
|
|
|
|
func (m *memoryNodes) AuthenticateNode(_ context.Context, nodeID, token string) (model.Node, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
if m.tokens[nodeID] != token {
|
|
return model.Node{}, errors.New("authentication failed")
|
|
}
|
|
return m.nodes[nodeID], nil
|
|
}
|
|
|
|
func (m *memoryNodes) ListNodes(context.Context) ([]model.Node, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
result := make([]model.Node, 0, len(m.nodes))
|
|
for _, node := range m.nodes {
|
|
result = append(result, node)
|
|
}
|
|
return result, nil
|
|
}
|
|
|
|
func (m *memoryNodes) UpdateNodeHeartbeat(_ context.Context, nodeID string, status model.NodeStatus, at time.Time, version, osVersion string) error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
node := m.nodes[nodeID]
|
|
node.Status = status
|
|
node.LastSeen = &at
|
|
node.Version = version
|
|
node.OSVersion = osVersion
|
|
m.nodes[nodeID] = node
|
|
return nil
|
|
}
|
|
|
|
func (m *memoryNodes) UpdateNodeStatus(_ context.Context, nodeID string, status model.NodeStatus) error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
node := m.nodes[nodeID]
|
|
node.Status = status
|
|
m.nodes[nodeID] = node
|
|
return nil
|
|
}
|
|
|
|
func TestHubHandshakeHeartbeatAndNodeList(t *testing.T) {
|
|
store := testNodes()
|
|
hub, err := NewHub(store, store, nil, HubConfig{NetworkConfigVersion: 3, HandshakeTimeout: time.Second})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
server := httptest.NewServer(hub)
|
|
defer server.Close()
|
|
controlURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
|
|
|
siteSocket := connectNode(t, controlURL, protocol.HelloPayload{
|
|
NodeID: "site", NodeToken: "site-token", ConfigVersion: 3,
|
|
Capabilities: protocol.NodeCapabilities{RemoteSubnet: true, NetstackStatus: "READY", TCPCapacity: 2048, UDPCapacity: 4096},
|
|
})
|
|
defer siteSocket.Close(websocket.StatusNormalClosure, "test done")
|
|
engineerSocket := connectNode(t, controlURL, protocol.HelloPayload{
|
|
NodeID: "engineer", NodeToken: "engineer-token", ConfigVersion: 3,
|
|
})
|
|
defer engineerSocket.Close(websocket.StatusNormalClosure, "test done")
|
|
|
|
var nodeListEnvelope protocol.ControlEnvelope
|
|
if err := wsjson.Read(context.Background(), engineerSocket, &nodeListEnvelope); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if nodeListEnvelope.Type != protocol.ControlNodeList {
|
|
t.Fatalf("message type = %s, want NODE_LIST", nodeListEnvelope.Type)
|
|
}
|
|
var nodeList protocol.NodeListPayload
|
|
if err := nodeListEnvelope.DecodePayload(&nodeList); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if len(nodeList.Sites) != 1 || !nodeList.Sites[0].Online || !nodeList.Sites[0].RemoteSubnetCapability {
|
|
t.Fatalf("unexpected Node list: %+v", nodeList)
|
|
}
|
|
store.mu.Lock()
|
|
eventCount := len(store.events)
|
|
store.mu.Unlock()
|
|
if eventCount < 2 {
|
|
t.Fatalf("Control connection events = %d, want at least 2", eventCount)
|
|
}
|
|
|
|
heartbeat, _ := protocol.NewControlEnvelope(protocol.ControlHeartbeat, "hb-1", protocol.HeartbeatPayload{
|
|
Timestamp: time.Now().UTC(), Status: "OK",
|
|
})
|
|
if err := wsjson.Write(context.Background(), engineerSocket, heartbeat); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var heartbeatReply protocol.ControlEnvelope
|
|
if err := wsjson.Read(context.Background(), engineerSocket, &heartbeatReply); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if heartbeatReply.Type != protocol.ControlHeartbeat || heartbeatReply.RequestID != "hb-1" {
|
|
t.Fatalf("heartbeat reply = %+v", heartbeatReply)
|
|
}
|
|
|
|
conflict, _ := protocol.NewControlEnvelope(protocol.ControlHeartbeat, "overlay-conflict-1", protocol.HeartbeatPayload{
|
|
Timestamp: time.Now().UTC(), Status: string(protocol.ErrorOverlayLocalConflict),
|
|
})
|
|
if err := wsjson.Write(context.Background(), engineerSocket, conflict); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := wsjson.Read(context.Background(), engineerSocket, &heartbeatReply); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
store.mu.Lock()
|
|
var conflictEvent *model.EventLog
|
|
for index := range store.events {
|
|
if store.events[index].Message == "节点拒绝了 Overlay 网络配置" {
|
|
value := store.events[index]
|
|
conflictEvent = &value
|
|
}
|
|
}
|
|
store.mu.Unlock()
|
|
if conflictEvent == nil || conflictEvent.Level != "ERROR" || conflictEvent.NodeID != "engineer" {
|
|
t.Fatalf("Overlay conflict event = %+v", conflictEvent)
|
|
}
|
|
var fields map[string]any
|
|
if err := json.Unmarshal(conflictEvent.FieldsJSON, &fields); err != nil || fields["error_code"] != string(protocol.ErrorOverlayLocalConflict) {
|
|
t.Fatalf("Overlay conflict fields = %s, error=%v", conflictEvent.FieldsJSON, err)
|
|
}
|
|
}
|
|
|
|
func TestHeartbeatThresholds(t *testing.T) {
|
|
now := time.Date(2026, 8, 25, 12, 0, 0, 0, time.UTC)
|
|
for _, test := range []struct {
|
|
age time.Duration
|
|
want model.NodeStatus
|
|
}{
|
|
{15 * time.Second, model.NodeOnline},
|
|
{15*time.Second + time.Nanosecond, model.NodeUnstable},
|
|
{30 * time.Second, model.NodeUnstable},
|
|
{30*time.Second + time.Nanosecond, model.NodeOffline},
|
|
} {
|
|
lastSeen := now.Add(-test.age)
|
|
if got := statusAt(&lastSeen, now); got != test.want {
|
|
t.Errorf("status at age %s = %s, want %s", test.age, got, test.want)
|
|
}
|
|
}
|
|
if got := statusAt(nil, now); got != model.NodeOffline {
|
|
t.Fatalf("nil last seen status = %s", got)
|
|
}
|
|
}
|
|
|
|
func TestSweepNotifiesHandlerWhenSiteBecomesOffline(t *testing.T) {
|
|
store := testNodes()
|
|
recorder := &statusChangeRecorder{}
|
|
hub, err := NewHub(store, store, recorder, HubConfig{NetworkConfigVersion: 1})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
now := time.Date(2026, 8, 27, 12, 0, 0, 0, time.UTC)
|
|
if err := store.UpdateNodeHeartbeat(context.Background(), "site", model.NodeOnline, now.Add(-31*time.Second), "1.0", "test"); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := hub.Sweep(context.Background(), now); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
recorder.mu.Lock()
|
|
defer recorder.mu.Unlock()
|
|
if len(recorder.changes) != 1 || recorder.changes[0].node.ID != "site" || recorder.changes[0].status != model.NodeOffline {
|
|
t.Fatalf("OFFLINE callbacks = %+v", recorder.changes)
|
|
}
|
|
}
|
|
|
|
func TestHubRequiresRebootstrapOnConfigVersionMismatch(t *testing.T) {
|
|
store := testNodes()
|
|
hub, err := NewHub(store, store, nil, HubConfig{NetworkConfigVersion: 2})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
server := httptest.NewServer(hub)
|
|
defer server.Close()
|
|
socket := connectNode(t, "ws"+strings.TrimPrefix(server.URL, "http"), protocol.HelloPayload{
|
|
NodeID: "engineer", NodeToken: "engineer-token", ConfigVersion: 1,
|
|
})
|
|
defer socket.Close(websocket.StatusNormalClosure, "test done")
|
|
var envelope protocol.ControlEnvelope
|
|
if err := wsjson.Read(context.Background(), socket, &envelope); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if envelope.Type != protocol.ControlRebootstrapRequired {
|
|
t.Fatalf("message type = %s, want REBOOTSTRAP_REQUIRED", envelope.Type)
|
|
}
|
|
var payload protocol.RebootstrapRequiredPayload
|
|
if err := envelope.DecodePayload(&payload); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if payload.ConfigVersion != 2 || payload.Reason != "CONFIG_VERSION_MISMATCH" {
|
|
t.Fatalf("rebootstrap payload = %+v", payload)
|
|
}
|
|
}
|
|
|
|
func TestClientCompletesHandshakeAndReceivesNodeList(t *testing.T) {
|
|
store := testNodes()
|
|
hub, err := NewHub(store, store, nil, HubConfig{NetworkConfigVersion: 1})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
server := httptest.NewServer(hub)
|
|
defer server.Close()
|
|
received := make(chan protocol.ControlMessageType, 1)
|
|
client, err := NewClient(ClientConfig{
|
|
URL: "ws" + strings.TrimPrefix(server.URL, "http"),
|
|
Hello: protocol.HelloPayload{NodeID: "engineer", NodeToken: "engineer-token", ConfigVersion: 1},
|
|
HeartbeatInterval: 20 * time.Millisecond,
|
|
}, func(_ context.Context, envelope protocol.ControlEnvelope) error {
|
|
received <- envelope.Type
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
done := make(chan error, 1)
|
|
go func() {
|
|
_, err := client.runOnce(ctx)
|
|
done <- err
|
|
}()
|
|
select {
|
|
case messageType := <-received:
|
|
if messageType != protocol.ControlNodeList {
|
|
t.Fatalf("received %s, want NODE_LIST", messageType)
|
|
}
|
|
cancel()
|
|
case <-time.After(2 * time.Second):
|
|
cancel()
|
|
t.Fatal("timed out waiting for Node list")
|
|
}
|
|
<-done
|
|
}
|
|
|
|
func TestClientReportsHeartbeatRTT(t *testing.T) {
|
|
store := testNodes()
|
|
hub, err := NewHub(store, store, nil, HubConfig{NetworkConfigVersion: 1})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
server := httptest.NewServer(hub)
|
|
defer server.Close()
|
|
rtt := make(chan time.Duration, 1)
|
|
client, err := NewClient(ClientConfig{
|
|
URL: "ws" + strings.TrimPrefix(server.URL, "http"),
|
|
Hello: protocol.HelloPayload{NodeID: "engineer", NodeToken: "engineer-token", ConfigVersion: 1},
|
|
HeartbeatInterval: 10 * time.Millisecond,
|
|
OnHeartbeatRTT: func(delay time.Duration) { rtt <- delay },
|
|
}, nil)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
done := make(chan error, 1)
|
|
go func() {
|
|
_, err := client.runOnce(ctx)
|
|
done <- err
|
|
}()
|
|
select {
|
|
case delay := <-rtt:
|
|
if delay < 0 || delay > time.Second {
|
|
t.Fatalf("unexpected heartbeat RTT: %s", delay)
|
|
}
|
|
cancel()
|
|
case <-time.After(2 * time.Second):
|
|
cancel()
|
|
t.Fatal("timed out waiting for heartbeat RTT")
|
|
}
|
|
<-done
|
|
}
|
|
|
|
func TestClientRefreshesBootstrapAfterContinuousDisconnect(t *testing.T) {
|
|
client, err := NewClient(ClientConfig{
|
|
URL: "ws://127.0.0.1:1",
|
|
Hello: protocol.HelloPayload{NodeID: "engineer", NodeToken: "token"},
|
|
BootstrapRefreshAfter: 25 * time.Millisecond,
|
|
}, nil)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
|
defer cancel()
|
|
err = client.Run(ctx)
|
|
if !errors.Is(err, protocol.ErrRebootstrapRequired) {
|
|
t.Fatalf("Run error = %v, want ErrRebootstrapRequired", err)
|
|
}
|
|
}
|
|
|
|
func connectNode(t *testing.T, controlURL string, hello protocol.HelloPayload) *websocket.Conn {
|
|
t.Helper()
|
|
socket, _, err := websocket.Dial(context.Background(), controlURL, nil)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
envelope, err := protocol.NewControlEnvelope(protocol.ControlHello, "hello", hello)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := wsjson.Write(context.Background(), socket, envelope); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var welcomeEnvelope protocol.ControlEnvelope
|
|
if err := wsjson.Read(context.Background(), socket, &welcomeEnvelope); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if welcomeEnvelope.Type != protocol.ControlWelcome {
|
|
t.Fatalf("first message = %s, want WELCOME", welcomeEnvelope.Type)
|
|
}
|
|
return socket
|
|
}
|
|
|
|
func testNodes() *memoryNodes {
|
|
return &memoryNodes{
|
|
nodes: map[string]model.Node{
|
|
"engineer": {ID: "engineer", Type: model.NodeTypeEngineer, Name: "Engineer", OverlayIP: netip.MustParseAddr("10.88.0.2"), Status: model.NodeOffline},
|
|
"site": {ID: "site", Type: model.NodeTypeSite, Name: "Site", OverlayIP: netip.MustParseAddr("10.88.0.3"), Status: model.NodeOffline},
|
|
},
|
|
tokens: map[string]string{"engineer": "engineer-token", "site": "site-token"},
|
|
}
|
|
}
|