初版功能完成
ci / Go checks (ubuntu-latest) (push) Has been cancelled
ci / Go checks (windows-latest) (push) Has been cancelled

This commit is contained in:
qsc
2026-08-29 13:12:17 +08:00
commit 142e5dc7d6
217 changed files with 21313 additions and 0 deletions
+253
View File
@@ -0,0 +1,253 @@
package control
import (
"context"
"errors"
"fmt"
"net/http"
"sync"
"time"
"github.com/coder/websocket"
"github.com/coder/websocket/wsjson"
"remlink/internal/protocol"
)
var ErrControlDisconnected = errors.New("Control WebSocket is not connected")
var DefaultReconnectBackoff = [...]time.Duration{
1 * time.Second, 2 * time.Second, 5 * time.Second, 10 * time.Second, 30 * time.Second,
}
const DefaultBootstrapRefreshAfter = 30 * time.Second
type ClientConfig struct {
URL string
Hello protocol.HelloPayload
HeartbeatInterval time.Duration
HandshakeTimeout time.Duration
HTTPClient *http.Client
OnConnectionState func(bool)
OnHeartbeatRTT func(time.Duration)
BootstrapRefreshAfter time.Duration
}
type EnvelopeHandler func(context.Context, protocol.ControlEnvelope) error
// Client maintains one authenticated WebSocket with the specified backoff.
type Client struct {
config ClientConfig
handler EnvelopeHandler
mu sync.RWMutex
active *clientConnection
}
type clientConnection struct {
socket *websocket.Conn
sendMu sync.Mutex
}
func NewClient(config ClientConfig, handler EnvelopeHandler) (*Client, error) {
if config.URL == "" || config.Hello.NodeID == "" || config.Hello.NodeToken == "" {
return nil, errors.New("Control URL, NodeID, and NodeToken are required")
}
if config.HeartbeatInterval <= 0 {
config.HeartbeatInterval = DefaultHeartbeatInterval
}
if config.HandshakeTimeout <= 0 {
config.HandshakeTimeout = 10 * time.Second
}
if config.BootstrapRefreshAfter <= 0 {
config.BootstrapRefreshAfter = DefaultBootstrapRefreshAfter
}
return &Client{config: config, handler: handler}, nil
}
// Run reconnects until ctx is canceled. Successful handshakes reset backoff.
func (c *Client) Run(ctx context.Context) error {
backoffIndex := 0
disconnectedSince := time.Now()
for {
if time.Since(disconnectedSince) >= c.config.BootstrapRefreshAfter {
return fmt.Errorf("Control unavailable for %s: %w", c.config.BootstrapRefreshAfter, protocol.ErrRebootstrapRequired)
}
connected, err := c.runOnce(ctx)
if ctx.Err() != nil {
return ctx.Err()
}
if errors.Is(err, protocol.ErrRebootstrapRequired) {
return err
}
if connected {
backoffIndex = 0
disconnectedSince = time.Now()
}
delay := DefaultReconnectBackoff[backoffIndex]
remaining := c.config.BootstrapRefreshAfter - time.Since(disconnectedSince)
if delay > remaining {
delay = remaining
}
if backoffIndex < len(DefaultReconnectBackoff)-1 {
backoffIndex++
}
timer := time.NewTimer(delay)
select {
case <-ctx.Done():
timer.Stop()
return ctx.Err()
case <-timer.C:
_ = err
}
}
}
// Send writes one typed Node-to-Server message on the current authenticated
// connection. Callers may retry after ErrControlDisconnected.
func (c *Client) Send(ctx context.Context, messageType protocol.ControlMessageType, requestID string, payload any) error {
c.mu.RLock()
active := c.active
c.mu.RUnlock()
if active == nil {
return ErrControlDisconnected
}
envelope, err := protocol.NewControlEnvelope(messageType, requestID, payload)
if err != nil {
return err
}
return active.write(ctx, envelope)
}
func (c *Client) runOnce(ctx context.Context) (bool, error) {
dialOptions := &websocket.DialOptions{HTTPClient: c.config.HTTPClient, CompressionMode: websocket.CompressionDisabled}
socket, response, err := websocket.Dial(ctx, c.config.URL, dialOptions)
if err != nil {
if response != nil {
return false, fmt.Errorf("dial Control WebSocket: HTTP %d: %w", response.StatusCode, err)
}
return false, fmt.Errorf("dial Control WebSocket: %w", err)
}
defer socket.Close(websocket.StatusNormalClosure, "Node stopping")
socket.SetReadLimit(maxControlMessage)
handshakeContext, cancel := context.WithTimeout(ctx, c.config.HandshakeTimeout)
helloEnvelope, err := protocol.NewControlEnvelope(protocol.ControlHello, "", c.config.Hello)
if err != nil {
cancel()
return false, err
}
if err := wsjson.Write(handshakeContext, socket, helloEnvelope); err != nil {
cancel()
return false, fmt.Errorf("send HELLO: %w", err)
}
var welcomeEnvelope protocol.ControlEnvelope
if err := wsjson.Read(handshakeContext, socket, &welcomeEnvelope); err != nil {
cancel()
return false, fmt.Errorf("read WELCOME: %w", err)
}
cancel()
if welcomeEnvelope.Type != protocol.ControlWelcome {
return false, fmt.Errorf("first Server message is %s, want WELCOME", welcomeEnvelope.Type)
}
var welcome protocol.WelcomePayload
if err := welcomeEnvelope.DecodePayload(&welcome); err != nil {
return false, err
}
active := &clientConnection{socket: socket}
c.setActive(active)
if c.config.OnConnectionState != nil {
c.config.OnConnectionState(true)
}
defer func() {
c.clearActive(active)
if c.config.OnConnectionState != nil {
c.config.OnConnectionState(false)
}
}()
connectionContext, cancelConnection := context.WithCancel(ctx)
defer cancelConnection()
readErrors := make(chan error, 1)
var heartbeatMu sync.Mutex
heartbeats := make(map[string]time.Time)
go func() {
for {
var envelope protocol.ControlEnvelope
if err := wsjson.Read(connectionContext, socket, &envelope); err != nil {
readErrors <- err
return
}
if envelope.Type == protocol.ControlHeartbeat {
heartbeatMu.Lock()
sentAt, found := heartbeats[envelope.RequestID]
if found {
delete(heartbeats, envelope.RequestID)
}
heartbeatMu.Unlock()
if found && c.config.OnHeartbeatRTT != nil {
c.config.OnHeartbeatRTT(time.Since(sentAt))
}
continue
}
if !envelope.Type.Valid() {
readErrors <- fmt.Errorf("invalid Server Control type %q", envelope.Type)
return
}
if c.handler != nil {
if err := c.handler(connectionContext, envelope); err != nil {
readErrors <- err
return
}
}
}
}()
ticker := time.NewTicker(c.config.HeartbeatInterval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return true, ctx.Err()
case err := <-readErrors:
return true, fmt.Errorf("read Control message: %w", err)
case now := <-ticker.C:
requestID := fmt.Sprintf("hb-%d", now.UnixNano())
envelope, err := protocol.NewControlEnvelope(protocol.ControlHeartbeat, requestID, protocol.HeartbeatPayload{
Timestamp: now.UTC(), Status: "OK",
})
if err != nil {
return true, err
}
heartbeatMu.Lock()
heartbeats[requestID] = time.Now()
heartbeatMu.Unlock()
err = active.write(ctx, envelope)
if err != nil {
heartbeatMu.Lock()
delete(heartbeats, requestID)
heartbeatMu.Unlock()
return true, fmt.Errorf("send HEARTBEAT: %w", err)
}
}
}
}
func (c *Client) setActive(active *clientConnection) {
c.mu.Lock()
c.active = active
c.mu.Unlock()
}
func (c *Client) clearActive(expected *clientConnection) {
c.mu.Lock()
if c.active == expected {
c.active = nil
}
c.mu.Unlock()
}
func (c *clientConnection) write(ctx context.Context, envelope protocol.ControlEnvelope) error {
c.sendMu.Lock()
defer c.sendMu.Unlock()
return wsjson.Write(ctx, c.socket, envelope)
}
+480
View File
@@ -0,0 +1,480 @@
// Package control implements the Overlay-only Control WebSocket plane.
package control
import (
"context"
"encoding/json"
"errors"
"fmt"
"net"
"net/http"
"net/netip"
"sync"
"time"
"github.com/coder/websocket"
"github.com/coder/websocket/wsjson"
"remlink/internal/localization"
"remlink/internal/logging"
"remlink/internal/model"
"remlink/internal/protocol"
)
const (
DefaultHeartbeatInterval = 5 * time.Second
OnlineThreshold = 15 * time.Second
UnstableThreshold = 30 * time.Second
maxControlMessage = 1 << 20
defaultSendTimeout = 10 * time.Second
)
// Authenticator verifies NodeID+NodeToken without exposing token hashes.
type Authenticator interface {
AuthenticateNode(context.Context, string, string) (model.Node, error)
}
// NodeStore persists heartbeat-derived Node state.
type NodeStore interface {
ListNodes(context.Context) ([]model.Node, error)
UpdateNodeHeartbeat(context.Context, string, model.NodeStatus, time.Time, string, string) error
UpdateNodeStatus(context.Context, string, model.NodeStatus) error
}
type eventAppender interface {
AppendEvent(context.Context, model.EventLog) error
}
// MessageHandler receives authenticated non-heartbeat messages.
type MessageHandler interface {
HandleControl(context.Context, model.Node, protocol.ControlEnvelope) error
}
// NodeStatusChangeHandler receives authoritative heartbeat state transitions.
// Session orchestration uses OFFLINE transitions to close stale live Sessions.
type NodeStatusChangeHandler interface {
HandleNodeStatusChange(context.Context, model.Node, model.NodeStatus) error
}
type HubConfig struct {
NetworkConfigVersion uint64
EnforceRemoteIP bool
HeartbeatInterval time.Duration
HandshakeTimeout time.Duration
}
// Hub owns at most one active Control socket per NodeID.
type Hub struct {
mu sync.RWMutex
authenticator Authenticator
store NodeStore
handler MessageHandler
config HubConfig
connections map[string]*connection
capabilities map[string]protocol.NodeCapabilities
}
type connection struct {
node model.Node
socket *websocket.Conn
sendMu sync.Mutex
}
func NewHub(authenticator Authenticator, store NodeStore, handler MessageHandler, config HubConfig) (*Hub, error) {
if authenticator == nil || store == nil {
return nil, errors.New("Control authenticator and NodeStore are required")
}
if config.NetworkConfigVersion == 0 {
return nil, errors.New("Control network config version must be positive")
}
if config.HeartbeatInterval <= 0 {
config.HeartbeatInterval = DefaultHeartbeatInterval
}
if config.HandshakeTimeout <= 0 {
config.HandshakeTimeout = 10 * time.Second
}
return &Hub{
authenticator: authenticator, store: store, handler: handler, config: config,
connections: make(map[string]*connection), capabilities: make(map[string]protocol.NodeCapabilities),
}, nil
}
// SetMessageHandler installs the post-handshake protocol handler. It is safe
// to call during startup before accepting Control connections.
func (h *Hub) SetMessageHandler(handler MessageHandler) {
h.mu.Lock()
h.handler = handler
h.mu.Unlock()
}
// ServeHTTP upgrades only /control requests and requires HELLO as message one.
func (h *Hub) ServeHTTP(writer http.ResponseWriter, request *http.Request) {
remoteIP, remoteErr := remoteAddress(request.RemoteAddr)
socket, err := websocket.Accept(writer, request, &websocket.AcceptOptions{
CompressionMode: websocket.CompressionDisabled,
})
if err != nil {
return
}
socket.SetReadLimit(maxControlMessage)
defer socket.Close(websocket.StatusNormalClosure, "Control connection closed")
if remoteErr != nil {
_ = socket.Close(websocket.StatusPolicyViolation, "invalid Overlay source")
return
}
handshakeContext, cancel := context.WithTimeout(context.Background(), h.config.HandshakeTimeout)
var helloEnvelope protocol.ControlEnvelope
if err := wsjson.Read(handshakeContext, socket, &helloEnvelope); err != nil {
cancel()
_ = socket.Close(websocket.StatusPolicyViolation, "HELLO required")
return
}
cancel()
if helloEnvelope.Type != protocol.ControlHello {
_ = socket.Close(websocket.StatusPolicyViolation, "HELLO must be first")
return
}
var hello protocol.HelloPayload
if err := helloEnvelope.DecodePayload(&hello); err != nil {
_ = socket.Close(websocket.StatusPolicyViolation, "invalid HELLO")
return
}
node, err := h.authenticator.AuthenticateNode(context.Background(), hello.NodeID, hello.NodeToken)
if err != nil {
_ = socket.Close(websocket.StatusPolicyViolation, "Node authentication failed")
return
}
if h.config.EnforceRemoteIP && node.OverlayIP != remoteIP {
_ = socket.Close(websocket.StatusPolicyViolation, "Overlay source mismatch")
return
}
connected := &connection{node: node, socket: socket}
previous := h.register(connected, hello.Capabilities)
if previous != nil {
_ = previous.socket.Close(websocket.StatusPolicyViolation, "replaced by newer Node connection")
}
defer h.unregister(node.ID, connected)
now := time.Now().UTC()
if err := h.store.UpdateNodeHeartbeat(context.Background(), node.ID, model.NodeOnline, now, hello.Version, hello.OSVersion); err != nil {
_ = socket.Close(websocket.StatusInternalError, "persist HELLO failed")
return
}
connected.node.Status = model.NodeOnline
connected.node.LastSeen = &now
connected.node.Version = hello.Version
connected.node.OSVersion = hello.OSVersion
h.recordEvent(context.Background(), node.ID, "INFO", "节点 Control 通道已连接", map[string]any{
"overlay_ip": node.OverlayIP.String(), "version": hello.Version, "os_version": hello.OSVersion,
})
h.mu.RLock()
networkConfigVersion := h.config.NetworkConfigVersion
h.mu.RUnlock()
if err := h.send(connected, protocol.ControlWelcome, helloEnvelope.RequestID, protocol.WelcomePayload{
ServerTime: now, NetworkConfigVersion: networkConfigVersion,
}); err != nil {
return
}
if hello.ConfigVersion != networkConfigVersion {
_ = h.send(connected, protocol.ControlRebootstrapRequired, "", protocol.RebootstrapRequiredPayload{
ConfigVersion: networkConfigVersion, Reason: "CONFIG_VERSION_MISMATCH",
})
return
}
if node.Type == model.NodeTypeEngineer {
if err := h.sendNodeList(context.Background(), connected); err != nil {
return
}
}
if node.Type == model.NodeTypeSite {
h.broadcastNodeLists(context.Background())
}
for {
var envelope protocol.ControlEnvelope
if err := wsjson.Read(context.Background(), socket, &envelope); err != nil {
return
}
if envelope.Type == protocol.ControlHello || !envelope.Type.Valid() {
_ = socket.Close(websocket.StatusPolicyViolation, "invalid Control message type")
return
}
if envelope.Type == protocol.ControlHeartbeat {
if err := h.handleHeartbeat(connected, envelope, hello.Version, hello.OSVersion); err != nil {
return
}
continue
}
h.mu.RLock()
handler := h.handler
h.mu.RUnlock()
if handler == nil {
_ = socket.Close(websocket.StatusUnsupportedData, "message is not available in this phase")
return
}
if err := handler.HandleControl(context.Background(), connected.node, envelope); err != nil {
_ = socket.Close(websocket.StatusPolicyViolation, "Control message rejected")
return
}
}
}
// SetNetworkConfigVersion updates WELCOME after a completed network migration.
func (h *Hub) SetNetworkConfigVersion(version uint64) error {
if version == 0 {
return errors.New("network config version must be positive")
}
h.mu.Lock()
h.config.NetworkConfigVersion = version
h.mu.Unlock()
return nil
}
// ResetNodeConnection forces one Node to bootstrap/reconnect after an
// authoritative address or credential change.
func (h *Hub) ResetNodeConnection(nodeID, reason string) {
h.mu.RLock()
connected := h.connections[nodeID]
h.mu.RUnlock()
if connected != nil {
_ = connected.socket.Close(websocket.StatusGoingAway, reason)
}
}
// ResetConnections closes all current sockets after migration notifications.
func (h *Hub) ResetConnections(reason string) {
h.mu.RLock()
connections := make([]*connection, 0, len(h.connections))
for _, connected := range h.connections {
connections = append(connections, connected)
}
h.mu.RUnlock()
for _, connected := range connections {
_ = connected.socket.Close(websocket.StatusGoingAway, reason)
}
}
// Run classifies persisted heartbeat age until ctx is canceled.
func (h *Hub) Run(ctx context.Context) error {
ticker := time.NewTicker(h.config.HeartbeatInterval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
h.closeAll()
return ctx.Err()
case now := <-ticker.C:
if err := h.Sweep(ctx, now.UTC()); err != nil {
return err
}
}
}
}
// Sweep applies ONLINE/UNSTABLE/OFFLINE thresholds and refreshes Engineer lists.
func (h *Hub) Sweep(ctx context.Context, now time.Time) error {
nodes, err := h.store.ListNodes(ctx)
if err != nil {
return err
}
changed := false
for _, node := range nodes {
status := statusAt(node.LastSeen, now)
if node.Status == status {
continue
}
if err := h.store.UpdateNodeStatus(ctx, node.ID, status); err != nil {
return err
}
node.Status = status
level := "WARN"
if status == model.NodeOffline {
level = "ERROR"
}
h.recordEvent(ctx, node.ID, level, "节点心跳状态变更为 "+localization.NodeStatus(string(status)), nil)
if status == model.NodeOffline {
h.mu.RLock()
handler := h.handler
h.mu.RUnlock()
if listener, ok := handler.(NodeStatusChangeHandler); ok {
if err := listener.HandleNodeStatusChange(ctx, node, status); err != nil {
return fmt.Errorf("handle Node %s status %s: %w", node.ID, status, err)
}
}
}
changed = true
}
if changed {
h.broadcastNodeLists(ctx)
}
return nil
}
func (h *Hub) recordEvent(ctx context.Context, nodeID, level, message string, fields map[string]any) {
appender, ok := h.store.(eventAppender)
if !ok {
return
}
if fields == nil {
fields = map[string]any{}
}
raw, _ := json.Marshal(fields)
_ = appender.AppendEvent(ctx, model.EventLog{
Level: level, Module: string(logging.ModuleControl), NodeID: nodeID, Message: message, FieldsJSON: raw,
})
}
// Send routes a typed Server message to one connected Node.
func (h *Hub) Send(ctx context.Context, nodeID string, messageType protocol.ControlMessageType, payload any) error {
return h.SendRequest(ctx, nodeID, messageType, "", payload)
}
// SendRequest preserves a request ID while routing a typed Server message.
func (h *Hub) SendRequest(ctx context.Context, nodeID string, messageType protocol.ControlMessageType, requestID string, payload any) error {
h.mu.RLock()
connected := h.connections[nodeID]
h.mu.RUnlock()
if connected == nil {
return fmt.Errorf("Node %s has no Control connection", nodeID)
}
return h.sendContext(ctx, connected, messageType, requestID, payload)
}
func (h *Hub) handleHeartbeat(connected *connection, envelope protocol.ControlEnvelope, version, osVersion string) error {
var heartbeat protocol.HeartbeatPayload
if err := envelope.DecodePayload(&heartbeat); err != nil {
return err
}
if heartbeat.Status == string(protocol.ErrorOverlayLocalConflict) {
h.recordEvent(context.Background(), connected.node.ID, "ERROR", "节点拒绝了 Overlay 网络配置", map[string]any{
"error_code": string(protocol.ErrorOverlayLocalConflict), "reported_at": heartbeat.Timestamp.UTC(),
})
}
now := time.Now().UTC()
if err := h.store.UpdateNodeHeartbeat(context.Background(), connected.node.ID, model.NodeOnline, now, version, osVersion); err != nil {
return err
}
connected.node.Status = model.NodeOnline
connected.node.LastSeen = &now
return h.send(connected, protocol.ControlHeartbeat, envelope.RequestID,
protocol.HeartbeatPayload{Timestamp: now, Status: string(model.NodeOnline)})
}
func (h *Hub) register(connected *connection, capabilities protocol.NodeCapabilities) *connection {
h.mu.Lock()
defer h.mu.Unlock()
previous := h.connections[connected.node.ID]
h.connections[connected.node.ID] = connected
h.capabilities[connected.node.ID] = capabilities
return previous
}
func (h *Hub) unregister(nodeID string, expected *connection) {
h.mu.Lock()
if h.connections[nodeID] == expected {
delete(h.connections, nodeID)
}
h.mu.Unlock()
}
func (h *Hub) send(connected *connection, messageType protocol.ControlMessageType, requestID string, payload any) error {
return h.sendContext(context.Background(), connected, messageType, requestID, payload)
}
func (h *Hub) sendContext(ctx context.Context, connected *connection, messageType protocol.ControlMessageType, requestID string, payload any) error {
if _, hasDeadline := ctx.Deadline(); !hasDeadline {
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, defaultSendTimeout)
defer cancel()
}
envelope, err := protocol.NewControlEnvelope(messageType, requestID, payload)
if err != nil {
return err
}
connected.sendMu.Lock()
defer connected.sendMu.Unlock()
return wsjson.Write(ctx, connected.socket, envelope)
}
func (h *Hub) sendNodeList(ctx context.Context, connected *connection) error {
payload, err := h.nodeList(ctx)
if err != nil {
return err
}
return h.sendContext(ctx, connected, protocol.ControlNodeList, "", payload)
}
func (h *Hub) broadcastNodeLists(ctx context.Context) {
h.mu.RLock()
engineers := make([]*connection, 0)
for _, connected := range h.connections {
if connected.node.Type == model.NodeTypeEngineer {
engineers = append(engineers, connected)
}
}
h.mu.RUnlock()
for _, engineer := range engineers {
_ = h.sendNodeList(ctx, engineer)
}
}
func (h *Hub) nodeList(ctx context.Context) (protocol.NodeListPayload, error) {
nodes, err := h.store.ListNodes(ctx)
if err != nil {
return protocol.NodeListPayload{}, err
}
h.mu.RLock()
defer h.mu.RUnlock()
payload := protocol.NodeListPayload{Sites: make([]protocol.SiteSummary, 0)}
for _, node := range nodes {
if node.Type != model.NodeTypeSite {
continue
}
site := protocol.SiteSummary{
NodeID: node.ID, Name: node.Name, OverlayIP: node.OverlayIP.String(),
Online: node.Status == model.NodeOnline,
RemoteSubnetCapability: h.capabilities[node.ID].RemoteSubnet,
}
if node.LastSeen != nil {
site.LastSeen = *node.LastSeen
}
payload.Sites = append(payload.Sites, site)
}
return payload, nil
}
func (h *Hub) closeAll() {
h.mu.Lock()
connections := make([]*connection, 0, len(h.connections))
for _, connected := range h.connections {
connections = append(connections, connected)
}
h.connections = make(map[string]*connection)
h.mu.Unlock()
for _, connected := range connections {
_ = connected.socket.Close(websocket.StatusGoingAway, "Server stopping")
}
}
func statusAt(lastSeen *time.Time, now time.Time) model.NodeStatus {
if lastSeen == nil {
return model.NodeOffline
}
age := now.Sub(*lastSeen)
if age <= OnlineThreshold {
return model.NodeOnline
}
if age <= UnstableThreshold {
return model.NodeUnstable
}
return model.NodeOffline
}
func remoteAddress(remote string) (netip.Addr, error) {
host, _, err := net.SplitHostPort(remote)
if err != nil {
return netip.Addr{}, err
}
return netip.ParseAddr(host)
}
+373
View File
@@ -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"},
}
}
+134
View File
@@ -0,0 +1,134 @@
package control
import (
"context"
"errors"
"fmt"
"net"
"net/http"
"net/netip"
"strconv"
"sync"
"time"
)
// Supervisor supports the required live Control-listener move during an
// Overlay network migration.
type Supervisor struct {
mu sync.Mutex
handler http.Handler
port int
ctx context.Context
server *http.Server
listener net.Listener
errors chan error
closed bool
}
func NewSupervisor(handler http.Handler, port int) (*Supervisor, error) {
if handler == nil || port < 1 || port > 65535 {
return nil, errors.New("Control Supervisor requires Handler and valid port")
}
return &Supervisor{handler: handler, port: port, errors: make(chan error, 1)}, nil
}
func (s *Supervisor) Start(ctx context.Context, address netip.Addr) error {
if !address.Is4() {
return errors.New("Control listener address must be IPv4")
}
s.mu.Lock()
defer s.mu.Unlock()
if s.server != nil {
return errors.New("Control Supervisor already started")
}
s.ctx = ctx
server, listener, err := s.open(address)
if err != nil {
return err
}
s.server, s.listener = server, listener
s.serve(server, listener)
return nil
}
func (s *Supervisor) Rebind(address netip.Addr) error {
if !address.Is4() {
return errors.New("Control listener address must be IPv4")
}
s.mu.Lock()
defer s.mu.Unlock()
if s.closed || s.server == nil {
return errors.New("Control Supervisor is not running")
}
newServer, newListener, err := s.open(address)
if err != nil {
return err
}
oldServer := s.server
oldListener := s.listener
s.server, s.listener = newServer, newListener
s.serve(newServer, newListener)
// From this point the rebind is committed and callers may safely switch the
// rest of the Overlay. An old-server close failure must not be reported as
// if the new listener were absent; that would cause the caller to roll back
// wg0 while this Supervisor remained bound to the new address.
_ = oldServer.Close()
_ = oldListener.Close()
return nil
}
func (s *Supervisor) Wait(ctx context.Context) error {
select {
case <-ctx.Done():
return ctx.Err()
case err := <-s.errors:
return err
}
}
func (s *Supervisor) Close() error {
s.mu.Lock()
if s.closed {
s.mu.Unlock()
return nil
}
s.closed = true
server := s.server
s.mu.Unlock()
if server == nil {
return nil
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := server.Shutdown(ctx); err != nil {
_ = server.Close()
return err
}
return nil
}
func (s *Supervisor) open(address netip.Addr) (*http.Server, net.Listener, error) {
listenAddress := net.JoinHostPort(address.String(), strconv.Itoa(s.port))
listener, err := net.Listen("tcp4", listenAddress)
if err != nil {
return nil, nil, fmt.Errorf("listen for Overlay Control on %s: %w", listenAddress, err)
}
server := &http.Server{
Addr: listenAddress, Handler: s.handler, ReadHeaderTimeout: 10 * time.Second,
ReadTimeout: 0, WriteTimeout: 0, IdleTimeout: 0, MaxHeaderBytes: 1 << 20,
}
return server, listener, nil
}
func (s *Supervisor) serve(server *http.Server, listener net.Listener) {
go func() {
err := server.Serve(listener)
if errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
return
}
select {
case s.errors <- err:
default:
}
}()
}
+79
View File
@@ -0,0 +1,79 @@
package control
import (
"context"
"fmt"
"io"
"net"
"net/http"
"net/netip"
"strconv"
"testing"
"time"
)
func TestSupervisorRebindMovesListenerAndPreservesFailedRebind(t *testing.T) {
oldAddress := netip.MustParseAddr("127.0.0.1")
newAddress := netip.MustParseAddr("127.0.0.2")
port := availablePort(t, oldAddress)
handler := http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { _, _ = io.WriteString(writer, "ok") })
supervisor, err := NewSupervisor(handler, port)
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
if err := supervisor.Start(ctx, oldAddress); err != nil {
t.Fatal(err)
}
defer supervisor.Close()
client := &http.Client{Transport: &http.Transport{Proxy: nil}, Timeout: time.Second}
assertHTTPBody(t, client, oldAddress, port, "ok")
occupied, err := net.Listen("tcp4", net.JoinHostPort(newAddress.String(), strconv.Itoa(port)))
if err != nil {
t.Fatal(err)
}
if err := supervisor.Rebind(newAddress); err == nil {
occupied.Close()
t.Fatal("Rebind succeeded while the target address was occupied")
}
assertHTTPBody(t, client, oldAddress, port, "ok")
if err := occupied.Close(); err != nil {
t.Fatal(err)
}
if err := supervisor.Rebind(newAddress); err != nil {
t.Fatal(err)
}
assertHTTPBody(t, client, newAddress, port, "ok")
}
func availablePort(t *testing.T, address netip.Addr) int {
t.Helper()
listener, err := net.Listen("tcp4", net.JoinHostPort(address.String(), "0"))
if err != nil {
t.Fatal(err)
}
port := listener.Addr().(*net.TCPAddr).Port
if err := listener.Close(); err != nil {
t.Fatal(err)
}
return port
}
func assertHTTPBody(t *testing.T, client *http.Client, address netip.Addr, port int, want string) {
t.Helper()
response, err := client.Get(fmt.Sprintf("http://%s/", net.JoinHostPort(address.String(), strconv.Itoa(port))))
if err != nil {
t.Fatal(err)
}
defer response.Body.Close()
body, err := io.ReadAll(response.Body)
if err != nil {
t.Fatal(err)
}
if string(body) != want {
t.Fatalf("body=%q, want %q", body, want)
}
}