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