Files
HexaHost-GameCloud/apps/node-agent/internal/ws/client.go

499 lines
13 KiB
Go

package ws
import (
"context"
"encoding/json"
"fmt"
"log/slog"
"net/url"
"os"
"runtime"
"sync"
"time"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"github.com/hexahost/gamecloud/node-agent/internal/config"
"github.com/hexahost/gamecloud/node-agent/internal/protocol"
runtimepkg "github.com/hexahost/gamecloud/node-agent/internal/runtime"
)
const (
writeWait = 10 * time.Second
pongWait = 60 * time.Second
pingPeriod = (pongWait * 9) / 10
maxReconnectDelay = 60 * time.Second
)
// Client maintains an outbound WebSocket connection to the control plane.
type Client struct {
cfg *config.Config
log *slog.Logger
runtime *runtimepkg.Manager
agentVersion string
writeMu sync.Mutex
conn *websocket.Conn
}
// NewClient creates a WebSocket client for the control plane.
func NewClient(cfg *config.Config, log *slog.Logger, mgr *runtimepkg.Manager, agentVersion string) *Client {
return &Client{
cfg: cfg,
log: log,
runtime: mgr,
agentVersion: agentVersion,
}
}
// SetRuntime attaches the runtime manager after construction to break init cycles.
func (c *Client) SetRuntime(mgr *runtimepkg.Manager) {
c.runtime = mgr
}
// Run connects to the control plane and processes messages until ctx is cancelled.
func (c *Client) Run(ctx context.Context) error {
backoff := c.cfg.ReconnectBackoff
if backoff <= 0 {
backoff = 5 * time.Second
}
for {
if ctx.Err() != nil {
return ctx.Err()
}
err := c.session(ctx)
if ctx.Err() != nil {
return ctx.Err()
}
c.log.Warn("websocket disconnected", "error", err, "retry_in", backoff.String())
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(backoff):
}
backoff *= 2
if backoff > maxReconnectDelay {
backoff = maxReconnectDelay
}
}
}
func (c *Client) session(ctx context.Context) error {
conn, err := c.dial(ctx)
if err != nil {
return err
}
c.writeMu.Lock()
c.conn = conn
c.writeMu.Unlock()
defer func() {
c.writeMu.Lock()
c.conn = nil
c.writeMu.Unlock()
_ = conn.Close()
}()
conn.SetReadDeadline(time.Now().Add(pongWait))
conn.SetPongHandler(func(string) error {
return conn.SetReadDeadline(time.Now().Add(pongWait))
})
if err := c.sendHello(ctx); err != nil {
return fmt.Errorf("send hello: %w", err)
}
sessionCtx, cancel := context.WithCancel(ctx)
defer cancel()
errCh := make(chan error, 2)
go func() {
errCh <- c.heartbeatLoop(sessionCtx)
}()
go func() {
errCh <- c.readLoop(sessionCtx)
}()
select {
case <-ctx.Done():
cancel()
return ctx.Err()
case err := <-errCh:
cancel()
return err
}
}
func (c *Client) dial(ctx context.Context) (*websocket.Conn, error) {
endpoint, err := c.buildURL()
if err != nil {
return nil, err
}
dialer := websocket.DefaultDialer
conn, _, err := dialer.DialContext(ctx, endpoint, nil)
if err != nil {
return nil, fmt.Errorf("dial websocket: %w", err)
}
c.log.Info("connected to control plane", "url", redactToken(endpoint))
return conn, nil
}
func (c *Client) buildURL() (string, error) {
u, err := url.Parse(c.cfg.ControlPlaneURL)
if err != nil {
return "", fmt.Errorf("parse control plane url: %w", err)
}
q := u.Query()
q.Set("nodeId", c.cfg.NodeID)
if c.cfg.NodeToken != "" {
q.Set("token", c.cfg.NodeToken)
}
u.RawQuery = q.Encode()
return u.String(), nil
}
func (c *Client) sendHello(ctx context.Context) error {
hostname, _ := os.Hostname()
payload := protocol.AgentHelloPayload{
NodeID: c.cfg.NodeID,
AgentVersion: c.agentVersion,
ProtocolVersion: protocol.CurrentProtocolVersion,
Hostname: hostname,
Capabilities: []string{"minecraft", "docker"},
}
return c.send(ctx, protocol.TypeAgentHello, "", payload)
}
func (c *Client) heartbeatLoop(ctx context.Context) error {
ticker := time.NewTicker(c.cfg.HeartbeatInterval)
defer ticker.Stop()
pingTicker := time.NewTicker(pingPeriod)
defer pingTicker.Stop()
for {
select {
case <-ctx.Done():
return ctx.Err()
case <-ticker.C:
if err := c.sendHeartbeat(ctx); err != nil {
return err
}
case <-pingTicker.C:
if err := c.writePing(); err != nil {
return err
}
}
}
}
func (c *Client) sendHeartbeat(ctx context.Context) error {
var memStats runtime.MemStats
runtime.ReadMemStats(&memStats)
payload := protocol.AgentHeartbeatPayload{
NodeID: c.cfg.NodeID,
CPUUsagePercent: 0,
MemoryUsedBytes: int64(memStats.Alloc),
MemoryTotalBytes: int64(memStats.Sys),
RunningServers: c.runtime.RunningCount(),
}
return c.send(ctx, protocol.TypeAgentHeartbeat, "", payload)
}
func (c *Client) readLoop(ctx context.Context) error {
for {
if ctx.Err() != nil {
return ctx.Err()
}
c.writeMu.Lock()
conn := c.conn
c.writeMu.Unlock()
if conn == nil {
return fmt.Errorf("connection closed")
}
_, data, err := conn.ReadMessage()
if err != nil {
return fmt.Errorf("read message: %w", err)
}
env, err := protocol.DecodeEnvelope(data)
if err != nil {
c.log.Warn("invalid inbound envelope", "error", err)
continue
}
c.dispatch(ctx, env)
}
}
func (c *Client) dispatch(ctx context.Context, env *protocol.Envelope) {
correlationID := env.CorrelationID
if correlationID == "" {
correlationID = env.MessageID
}
switch env.Type {
case protocol.TypeServerProvision:
var payload protocol.ServerProvisionPayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.provision", "error", err)
return
}
c.runtime.HandleProvision(ctx, correlationID, payload)
case protocol.TypeServerStart:
var payload protocol.ServerStartPayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.start", "error", err)
return
}
c.runtime.HandleStart(ctx, correlationID, payload)
case protocol.TypeServerStop:
var payload protocol.ServerStopPayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.stop", "error", err)
return
}
c.runtime.HandleStop(ctx, correlationID, payload)
case protocol.TypeServerDelete:
var meta protocol.OperationMeta
if err := json.Unmarshal(env.Payload, &meta); err != nil {
c.log.Warn("decode server.delete", "error", err)
return
}
c.runtime.HandleDelete(ctx, correlationID, meta)
case protocol.TypeServerLogsSubscribe:
var payload protocol.ServerLogsSubscribePayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.logs.subscribe", "error", err)
return
}
c.runtime.HandleLogsSubscribe(ctx, env.MessageID, payload)
case protocol.TypeServerCommand:
var payload protocol.ServerCommandPayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.command", "error", err)
return
}
c.runtime.HandleCommand(ctx, env.MessageID, payload)
case protocol.TypeServerFilesList:
var payload protocol.ServerFilesListPayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.files.list", "error", err)
return
}
c.runtime.HandleFilesList(ctx, env.MessageID, payload)
case protocol.TypeServerFilesRead:
var payload protocol.ServerFilesReadPayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.files.read", "error", err)
return
}
c.runtime.HandleFilesRead(ctx, env.MessageID, payload)
case protocol.TypeServerFilesWrite:
var payload protocol.ServerFilesWritePayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.files.write", "error", err)
return
}
c.runtime.HandleFilesWrite(ctx, env.MessageID, payload)
case protocol.TypeServerFilesDelete:
var payload protocol.ServerFilesDeletePayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.files.delete", "error", err)
return
}
c.runtime.HandleFilesDelete(ctx, env.MessageID, payload)
case protocol.TypeServerFilesArchive:
var payload protocol.ServerFilesArchivePayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.files.archive", "error", err)
return
}
c.runtime.HandleFilesArchive(ctx, env.MessageID, payload)
case protocol.TypeServerFilesUnarchive:
var payload protocol.ServerFilesUnarchivePayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.files.unarchive", "error", err)
return
}
c.runtime.HandleFilesUnarchive(ctx, env.MessageID, payload)
case protocol.TypeServerBackupPrepare:
var meta protocol.OperationMeta
if err := json.Unmarshal(env.Payload, &meta); err != nil {
c.log.Warn("decode server.backup.prepare", "error", err)
return
}
c.runtime.HandleBackupPrepare(ctx, env.MessageID, meta)
case protocol.TypeServerBackupRelease:
var meta protocol.OperationMeta
if err := json.Unmarshal(env.Payload, &meta); err != nil {
c.log.Warn("decode server.backup.release", "error", err)
return
}
c.runtime.HandleBackupRelease(ctx, env.MessageID, meta)
case protocol.TypeServerWorldValidate:
var payload protocol.ServerWorldValidatePayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.world.validate", "error", err)
return
}
c.runtime.HandleWorldValidate(ctx, env.MessageID, payload)
case protocol.TypeServerWorldArchive:
var payload protocol.ServerWorldArchivePayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.world.archive", "error", err)
return
}
c.runtime.HandleWorldArchive(ctx, env.MessageID, payload)
case protocol.TypeServerStorageUpload:
var payload protocol.ServerStorageUploadPayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.storage.upload", "error", err)
return
}
c.runtime.HandleStorageUpload(ctx, env.MessageID, payload)
case protocol.TypeServerStorageDownload:
var payload protocol.ServerStorageDownloadPayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.storage.download", "error", err)
return
}
c.runtime.HandleStorageDownload(ctx, env.MessageID, payload)
case protocol.TypeServerWorldReplace:
var payload protocol.ServerWorldReplacePayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.world.replace", "error", err)
return
}
c.runtime.HandleWorldReplace(ctx, env.MessageID, payload)
case protocol.TypeServerAddonInstall:
var payload protocol.ServerAddonInstallPayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.addon.install", "error", err)
return
}
c.runtime.HandleAddonInstall(ctx, env.MessageID, payload)
case protocol.TypeServerAddonRemove:
var payload protocol.ServerAddonRemovePayload
if err := json.Unmarshal(env.Payload, &payload); err != nil {
c.log.Warn("decode server.addon.remove", "error", err)
return
}
c.runtime.HandleAddonRemove(ctx, env.MessageID, payload)
default:
c.log.Debug("ignored control message", "type", env.Type)
}
}
func (c *Client) send(ctx context.Context, msgType protocol.MessageType, correlationID string, payload any) error {
env, err := protocol.NewEnvelope(msgType, uuid.NewString(), payload)
if err != nil {
return err
}
if correlationID != "" {
env.CorrelationID = correlationID
}
return c.writeEnvelope(ctx, env)
}
func (c *Client) writeEnvelope(ctx context.Context, env *protocol.Envelope) error {
data, err := json.Marshal(env)
if err != nil {
return fmt.Errorf("marshal envelope: %w", err)
}
c.writeMu.Lock()
conn := c.conn
c.writeMu.Unlock()
if conn == nil {
return fmt.Errorf("not connected")
}
deadline, ok := ctx.Deadline()
if !ok {
deadline = time.Now().Add(writeWait)
}
if err := conn.SetWriteDeadline(deadline); err != nil {
return err
}
c.writeMu.Lock()
defer c.writeMu.Unlock()
if c.conn == nil {
return fmt.Errorf("not connected")
}
return c.conn.WriteMessage(websocket.TextMessage, data)
}
func (c *Client) writePing() error {
c.writeMu.Lock()
defer c.writeMu.Unlock()
if c.conn == nil {
return fmt.Errorf("not connected")
}
return c.conn.WriteMessage(websocket.PingMessage, nil)
}
// SendState implements runtime.Responder.
func (c *Client) SendState(ctx context.Context, correlationID string, payload protocol.ServerStatePayload) error {
return c.send(ctx, protocol.TypeServerState, correlationID, payload)
}
// SendLog implements runtime.Responder.
func (c *Client) SendLog(ctx context.Context, correlationID string, payload protocol.ServerLogPayload) error {
return c.send(ctx, protocol.TypeServerLog, correlationID, payload)
}
// SendOperationComplete implements runtime.Responder.
func (c *Client) SendOperationComplete(ctx context.Context, correlationID string, payload protocol.ServerOperationResultPayload) error {
return c.send(ctx, protocol.TypeServerOperationComplete, correlationID, payload)
}
// SendOperationFailed implements runtime.Responder.
func (c *Client) SendOperationFailed(ctx context.Context, correlationID string, payload protocol.ServerOperationResultPayload) error {
return c.send(ctx, protocol.TypeServerOperationFailed, correlationID, payload)
}
func redactToken(rawURL string) string {
u, err := url.Parse(rawURL)
if err != nil {
return rawURL
}
if u.Query().Get("token") != "" {
q := u.Query()
q.Set("token", "REDACTED")
u.RawQuery = q.Encode()
}
return u.String()
}