初版功能完成
This commit is contained in:
@@ -0,0 +1,373 @@
|
||||
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"},
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user