When the WebSocket breaks (broken pipe), the main loop can get stuck in handleMessage -> dispatchJob -> submit if the worker pool is full. submit blocks on ctx.Done() but ctx is the parent context which only cancels on SIGINT/SIGTERM, not on connection errors. Create a session-scoped context that cancels when read/write errors occur. Pass it to handleMessage so blocked submit calls unblock immediately, allowing the main loop to receive the write error and reconnect.
542 lines
15 KiB
Go
542 lines
15 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"log/slog"
|
|
"math/rand/v2"
|
|
"runtime"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"sync/atomic"
|
|
"time"
|
|
|
|
"github.com/towerops-app/towerops-agent/pb"
|
|
"google.golang.org/protobuf/proto"
|
|
)
|
|
|
|
var doSelfUpdate = selfUpdate
|
|
|
|
var errRestartRequested = fmt.Errorf("restart requested")
|
|
var joinTimeout = 10 * time.Second
|
|
var heartbeatInterval = 60 * time.Second
|
|
var channelHeartbeatInterval = 25 * time.Second
|
|
|
|
const maxJobPayloadBytes = 4 << 20 // 4 MB — well above any legitimate job list
|
|
|
|
// channelMsg is the WebSocket channel message format (JSON wrapper around binary protobuf).
|
|
type channelMsg struct {
|
|
Topic string `json:"topic"`
|
|
Event string `json:"event"`
|
|
Payload json.RawMessage `json:"payload"`
|
|
Ref *string `json:"ref"`
|
|
}
|
|
|
|
// runAgent connects to the server and runs the event loop with reconnect.
|
|
func runAgent(ctx context.Context, wsURL, token string) {
|
|
baseURL := strings.TrimRight(wsURL, "/")
|
|
retryDelay := time.Second
|
|
maxRetry := 10 * time.Second
|
|
const successfulConnectionThreshold = 30 * time.Second
|
|
|
|
for {
|
|
select {
|
|
case <-ctx.Done():
|
|
return
|
|
default:
|
|
}
|
|
|
|
sessionStart := time.Now()
|
|
err := runSession(ctx, baseURL, token)
|
|
sessionDuration := time.Since(sessionStart)
|
|
|
|
// Reset backoff if session ran successfully for a while (indicates stable connection)
|
|
if sessionDuration >= successfulConnectionThreshold {
|
|
slog.Debug("resetting reconnect backoff after successful session",
|
|
"duration", sessionDuration,
|
|
"previous_delay", retryDelay)
|
|
retryDelay = time.Second
|
|
}
|
|
|
|
if ctx.Err() != nil {
|
|
return
|
|
}
|
|
if errors.Is(err, errRestartRequested) {
|
|
slog.Info("restart requested, reconnecting immediately")
|
|
retryDelay = time.Second
|
|
continue
|
|
}
|
|
if err != nil {
|
|
slog.Error("agent disconnected", "error", err)
|
|
}
|
|
|
|
slog.Info("reconnecting", "delay", retryDelay)
|
|
select {
|
|
case <-ctx.Done():
|
|
return
|
|
case <-time.After(retryDelay):
|
|
}
|
|
retryDelay = nextBackoff(retryDelay, maxRetry)
|
|
}
|
|
}
|
|
|
|
// runSession runs a single WebSocket session. Returns when disconnected or ctx cancelled.
|
|
func runSession(ctx context.Context, baseURL, token string) error {
|
|
endpoint := baseURL + "/socket/agent/websocket"
|
|
slog.Info("connecting", "url", sanitizeURL(endpoint))
|
|
|
|
ws, err := WSDial(endpoint)
|
|
if err != nil {
|
|
return fmt.Errorf("connect: %w", err)
|
|
}
|
|
defer func() { _ = ws.Close() }()
|
|
|
|
agentID := fmt.Sprintf("agent-%d", time.Now().Unix())
|
|
topic := "agent:" + agentID
|
|
|
|
slog.Info("connected", "agent_id", agentID)
|
|
|
|
// Session-scoped context: cancelled on write/read errors so blocked
|
|
// pool submits unblock immediately instead of waiting for workers to finish.
|
|
sessionCtx, sessionCancel := context.WithCancel(ctx)
|
|
defer sessionCancel()
|
|
|
|
// Channel for serializing WebSocket writes
|
|
writeCh := make(chan []byte, 10000)
|
|
|
|
// Worker pools — bounded concurrency for each job type
|
|
pools := &jobPools{
|
|
snmp: newWorkerPool(100),
|
|
mikrotik: newWorkerPool(20),
|
|
ping: newWorkerPool(50),
|
|
checks: newWorkerPool(50),
|
|
}
|
|
|
|
// Result channels
|
|
snmpResultCh := make(chan *pb.SnmpResult, 10000)
|
|
mikrotikResultCh := make(chan *pb.MikrotikResult, 5000)
|
|
credTestResultCh := make(chan *pb.CredentialTestResult, 5000)
|
|
monitoringCheckCh := make(chan *pb.MonitoringCheck, 10000)
|
|
checkResultCh := make(chan *pb.CheckResult, 10000)
|
|
lldpTopologyResultCh := make(chan *pb.LldpTopologyResult, 1000)
|
|
|
|
// Ref counter for outbound messages
|
|
var refCounter atomic.Uint64
|
|
refCounter.Store(1)
|
|
|
|
nextRef := func() string {
|
|
return strconv.FormatUint(refCounter.Add(1), 10)
|
|
}
|
|
|
|
sendMsg := func(event string, payload json.RawMessage) {
|
|
msg := channelMsg{
|
|
Topic: topic,
|
|
Event: event,
|
|
Payload: payload,
|
|
}
|
|
data, err := json.Marshal(msg)
|
|
if err != nil {
|
|
slog.Error("marshal message", "error", err)
|
|
return
|
|
}
|
|
select {
|
|
case writeCh <- data:
|
|
default:
|
|
slog.Warn("write channel full, dropping message", "event", event)
|
|
}
|
|
}
|
|
|
|
bufPool := &sync.Pool{
|
|
New: func() any {
|
|
b := make([]byte, 0, 4096)
|
|
return &b
|
|
},
|
|
}
|
|
|
|
sendBinaryResult := func(event string, msg proto.Message) {
|
|
bp := bufPool.Get().(*[]byte)
|
|
buf := (*bp)[:0]
|
|
bin, err := proto.MarshalOptions{}.MarshalAppend(buf, msg)
|
|
if err != nil {
|
|
slog.Error("marshal protobuf", "error", err)
|
|
return
|
|
}
|
|
encoded := base64.StdEncoding.EncodeToString(bin)
|
|
// Zero the serialized protobuf data (may contain credentials).
|
|
// MarshalAppend may return a new backing array, so zero both.
|
|
zeroBytes(bin)
|
|
full := (*bp)[:cap(*bp)]
|
|
zeroBytes(full)
|
|
*bp = full[:0]
|
|
bufPool.Put(bp)
|
|
payload, _ := json.Marshal(map[string]string{"binary": encoded})
|
|
sendMsg(event, payload)
|
|
}
|
|
|
|
// Reader goroutine — must start before join so we can receive the reply
|
|
msgCh := make(chan []byte, 100)
|
|
errCh := make(chan error, 1)
|
|
go func() {
|
|
for {
|
|
data, _, err := ws.ReadMessage()
|
|
if err != nil {
|
|
errCh <- err
|
|
sessionCancel() // Unblock any stuck pool submits
|
|
return
|
|
}
|
|
msgCh <- data
|
|
}
|
|
}()
|
|
|
|
// Join channel
|
|
joinPayload, _ := json.Marshal(map[string]string{"token": token})
|
|
joinMsg := channelMsg{
|
|
Topic: topic,
|
|
Event: "phx_join",
|
|
Payload: joinPayload,
|
|
Ref: strPtr("1"),
|
|
}
|
|
joinData, _ := json.Marshal(joinMsg)
|
|
if err := ws.WriteText(joinData); err != nil {
|
|
return fmt.Errorf("send join: %w", err)
|
|
}
|
|
slog.Debug("sent channel join request")
|
|
|
|
// Wait for join reply before entering main loop
|
|
select {
|
|
case data := <-msgCh:
|
|
var reply channelMsg
|
|
if err := json.Unmarshal(data, &reply); err != nil {
|
|
return fmt.Errorf("join reply unmarshal: %w", err)
|
|
}
|
|
if reply.Event == "phx_reply" {
|
|
var status struct {
|
|
Status string `json:"status"`
|
|
Response any `json:"response"`
|
|
}
|
|
if err := json.Unmarshal(reply.Payload, &status); err == nil && status.Status != "ok" {
|
|
return fmt.Errorf("join rejected: %s", status.Status)
|
|
}
|
|
}
|
|
slog.Info("channel joined")
|
|
case err := <-errCh:
|
|
return fmt.Errorf("read during join: %w", err)
|
|
case <-time.After(joinTimeout):
|
|
return fmt.Errorf("join timeout")
|
|
}
|
|
|
|
// Writer goroutine - serializes all writes to the WebSocket
|
|
var writerWg sync.WaitGroup
|
|
writeErrCh := make(chan error, 1)
|
|
writerWg.Add(1)
|
|
go func() {
|
|
defer writerWg.Done()
|
|
for data := range writeCh {
|
|
if err := ws.WriteText(data); err != nil {
|
|
slog.Error("websocket write", "error", err)
|
|
select {
|
|
case writeErrCh <- err:
|
|
default:
|
|
}
|
|
sessionCancel() // Unblock any stuck pool submits
|
|
return
|
|
}
|
|
}
|
|
}()
|
|
|
|
heartbeatTicker := time.NewTicker(heartbeatInterval)
|
|
defer heartbeatTicker.Stop()
|
|
channelHeartbeatTicker := time.NewTicker(channelHeartbeatInterval)
|
|
defer channelHeartbeatTicker.Stop()
|
|
flushTicker := time.NewTicker(100 * time.Millisecond)
|
|
defer flushTicker.Stop()
|
|
startTime := time.Now()
|
|
|
|
defer func() {
|
|
const poolShutdownTimeout = 5 * time.Second
|
|
if !pools.snmp.stopWithTimeout(poolShutdownTimeout) {
|
|
slog.Warn("snmp pool shutdown timed out, abandoning in-flight jobs")
|
|
}
|
|
if !pools.mikrotik.stopWithTimeout(poolShutdownTimeout) {
|
|
slog.Warn("mikrotik pool shutdown timed out, abandoning in-flight jobs")
|
|
}
|
|
if !pools.ping.stopWithTimeout(poolShutdownTimeout) {
|
|
slog.Warn("ping pool shutdown timed out, abandoning in-flight jobs")
|
|
}
|
|
if !pools.checks.stopWithTimeout(poolShutdownTimeout) {
|
|
slog.Warn("checks pool shutdown timed out, abandoning in-flight jobs")
|
|
}
|
|
close(writeCh)
|
|
writerWg.Wait()
|
|
}()
|
|
|
|
var snmpBatch []*pb.SnmpResult
|
|
|
|
flushSnmpBatch := func() {
|
|
if len(snmpBatch) == 0 {
|
|
return
|
|
}
|
|
for _, r := range snmpBatch {
|
|
sendBinaryResult("result", r)
|
|
}
|
|
slog.Info("flushed snmp results", "count", len(snmpBatch))
|
|
snmpBatch = snmpBatch[:0]
|
|
}
|
|
|
|
for {
|
|
select {
|
|
case <-ctx.Done():
|
|
slog.Info("shutdown signal, closing connection")
|
|
flushSnmpBatch()
|
|
return nil
|
|
|
|
case err := <-errCh:
|
|
flushSnmpBatch()
|
|
return fmt.Errorf("read: %w", err)
|
|
|
|
case err := <-writeErrCh:
|
|
flushSnmpBatch()
|
|
return fmt.Errorf("write: %w", err)
|
|
|
|
case data := <-msgCh:
|
|
var msg channelMsg
|
|
if err := json.Unmarshal(data, &msg); err != nil {
|
|
slog.Warn("invalid message", "error", err)
|
|
continue
|
|
}
|
|
if handleMessage(sessionCtx, msg, pools, snmpResultCh, mikrotikResultCh, credTestResultCh, monitoringCheckCh, checkResultCh, lldpTopologyResultCh) {
|
|
flushSnmpBatch()
|
|
return errRestartRequested
|
|
}
|
|
|
|
case result := <-snmpResultCh:
|
|
snmpBatch = append(snmpBatch, result)
|
|
if len(snmpBatch) >= 50 {
|
|
flushSnmpBatch()
|
|
}
|
|
|
|
case result := <-mikrotikResultCh:
|
|
sendBinaryResult("mikrotik_result", result)
|
|
slog.Info("sent mikrotik result", "device", result.DeviceId, "job", result.JobId)
|
|
|
|
case result := <-credTestResultCh:
|
|
sendBinaryResult("credential_test_result", result)
|
|
slog.Info("sent credential test result", "test_id", result.TestId, "success", result.Success)
|
|
|
|
case result := <-monitoringCheckCh:
|
|
sendBinaryResult("monitoring_check", result)
|
|
slog.Info("sent monitoring check", "device", result.DeviceId, "status", result.Status)
|
|
|
|
case result := <-checkResultCh:
|
|
sendBinaryResult("check_result", result)
|
|
slog.Info("sent check result", "check", result.CheckId, "status", result.Status)
|
|
|
|
case result := <-lldpTopologyResultCh:
|
|
sendBinaryResult("lldp_topology_result", result)
|
|
slog.Info("sent LLDP topology result", "device", result.DeviceId, "neighbors", len(result.Neighbors))
|
|
|
|
case <-flushTicker.C:
|
|
flushSnmpBatch()
|
|
|
|
case <-heartbeatTicker.C:
|
|
hb := &pb.AgentHeartbeat{
|
|
Version: version,
|
|
UptimeSeconds: uint64(time.Since(startTime).Seconds()),
|
|
Arch: runtime.GOARCH,
|
|
}
|
|
sendBinaryResult("heartbeat", hb)
|
|
slog.Debug("sent heartbeat")
|
|
|
|
case <-channelHeartbeatTicker.C:
|
|
ref := nextRef()
|
|
msg := channelMsg{
|
|
Topic: "phoenix",
|
|
Event: "heartbeat",
|
|
Payload: json.RawMessage(`{}`),
|
|
Ref: &ref,
|
|
}
|
|
data, _ := json.Marshal(msg)
|
|
select {
|
|
case writeCh <- data:
|
|
default:
|
|
}
|
|
slog.Debug("sent channel heartbeat", "ref", ref)
|
|
}
|
|
}
|
|
}
|
|
|
|
// handleMessage dispatches incoming channel messages.
|
|
// Returns true if the session should end (e.g. restart requested).
|
|
func handleMessage(
|
|
ctx context.Context,
|
|
msg channelMsg,
|
|
pools *jobPools,
|
|
snmpResultCh chan<- *pb.SnmpResult,
|
|
mikrotikResultCh chan<- *pb.MikrotikResult,
|
|
credTestResultCh chan<- *pb.CredentialTestResult,
|
|
monitoringCheckCh chan<- *pb.MonitoringCheck,
|
|
checkResultCh chan<- *pb.CheckResult,
|
|
lldpTopologyResultCh chan<- *pb.LldpTopologyResult,
|
|
) bool {
|
|
switch msg.Event {
|
|
case "phx_reply":
|
|
slog.Debug("channel reply", "topic", msg.Topic)
|
|
|
|
case "jobs", "discovery_job", "backup_job":
|
|
var payload struct {
|
|
Binary string `json:"binary"`
|
|
}
|
|
if err := json.Unmarshal(msg.Payload, &payload); err != nil {
|
|
slog.Error("decode job payload", "error", err)
|
|
return false
|
|
}
|
|
if len(payload.Binary) > maxJobPayloadBytes {
|
|
slog.Error("job payload too large", "size", len(payload.Binary), "max", maxJobPayloadBytes)
|
|
return false
|
|
}
|
|
bin, err := base64.StdEncoding.DecodeString(payload.Binary)
|
|
if err != nil {
|
|
slog.Error("decode base64", "error", err)
|
|
return false
|
|
}
|
|
var jobList pb.AgentJobList
|
|
if err := proto.Unmarshal(bin, &jobList); err != nil {
|
|
slog.Error("unmarshal job list", "error", err)
|
|
return false
|
|
}
|
|
slog.Info("received jobs", "count", len(jobList.Jobs))
|
|
for _, job := range jobList.Jobs {
|
|
dispatchJob(ctx, job, pools, snmpResultCh, mikrotikResultCh, credTestResultCh, monitoringCheckCh, checkResultCh, lldpTopologyResultCh)
|
|
}
|
|
|
|
case "check_jobs":
|
|
var payload struct {
|
|
Binary string `json:"binary"`
|
|
}
|
|
if err := json.Unmarshal(msg.Payload, &payload); err != nil {
|
|
slog.Error("decode check_jobs payload", "error", err)
|
|
return false
|
|
}
|
|
if len(payload.Binary) > maxJobPayloadBytes {
|
|
slog.Error("check_jobs payload too large", "size", len(payload.Binary), "max", maxJobPayloadBytes)
|
|
return false
|
|
}
|
|
bin, err := base64.StdEncoding.DecodeString(payload.Binary)
|
|
if err != nil {
|
|
slog.Error("decode check_jobs base64", "error", err)
|
|
return false
|
|
}
|
|
var checkList pb.CheckList
|
|
if err := proto.Unmarshal(bin, &checkList); err != nil {
|
|
slog.Error("unmarshal check list", "error", err)
|
|
return false
|
|
}
|
|
slog.Info("received checks", "count", len(checkList.Checks))
|
|
for _, check := range checkList.Checks {
|
|
executeCheck(ctx, check, pools, checkResultCh)
|
|
}
|
|
|
|
case "restart":
|
|
slog.Info("restart requested by server")
|
|
return true
|
|
|
|
case "update":
|
|
var payload struct {
|
|
URL string `json:"url"`
|
|
Checksum string `json:"checksum"`
|
|
}
|
|
if err := json.Unmarshal(msg.Payload, &payload); err != nil || payload.URL == "" || payload.Checksum == "" {
|
|
slog.Error("invalid update payload")
|
|
return false
|
|
}
|
|
slog.Info("update requested", "url", sanitizeURL(payload.URL))
|
|
if err := doSelfUpdate(payload.URL, payload.Checksum); err != nil {
|
|
slog.Error("self-update failed", "error", err)
|
|
}
|
|
|
|
default:
|
|
slog.Debug("ignoring event", "event", msg.Event)
|
|
}
|
|
return false
|
|
}
|
|
|
|
// jobPools holds the worker pools for each job type.
|
|
type jobPools struct {
|
|
snmp *workerPool
|
|
mikrotik *workerPool
|
|
ping *workerPool
|
|
checks *workerPool
|
|
}
|
|
|
|
// dispatchJob routes a job to the appropriate worker pool.
|
|
func dispatchJob(
|
|
ctx context.Context,
|
|
job *pb.AgentJob,
|
|
pools *jobPools,
|
|
snmpResultCh chan<- *pb.SnmpResult,
|
|
mikrotikResultCh chan<- *pb.MikrotikResult,
|
|
credTestResultCh chan<- *pb.CredentialTestResult,
|
|
monitoringCheckCh chan<- *pb.MonitoringCheck,
|
|
checkResultCh chan<- *pb.CheckResult,
|
|
lldpTopologyResultCh chan<- *pb.LldpTopologyResult,
|
|
) {
|
|
slog.Info("starting job", "job_id", job.JobId, "type", job.JobType)
|
|
|
|
var ok bool
|
|
switch job.JobType {
|
|
case pb.JobType_MIKROTIK:
|
|
ok = pools.mikrotik.submit(ctx, func() { executeMikrotikJob(ctx, job, mikrotikResultCh) })
|
|
case pb.JobType_TEST_CREDENTIALS:
|
|
ok = pools.snmp.submit(ctx, func() { executeCredentialTest(ctx, job, credTestResultCh) })
|
|
case pb.JobType_PING:
|
|
ok = pools.ping.submit(ctx, func() { executePingJob(ctx, job, monitoringCheckCh) })
|
|
case pb.JobType_LLDP_TOPOLOGY:
|
|
ok = pools.snmp.submit(ctx, func() { executeLldpTopologyJob(ctx, job, lldpTopologyResultCh) })
|
|
default:
|
|
ok = pools.snmp.submit(ctx, func() { executeSnmpJob(ctx, job, snmpResultCh) })
|
|
}
|
|
if !ok {
|
|
slog.Warn("job dropped, pool full", "job_id", job.JobId)
|
|
}
|
|
}
|
|
|
|
// nextBackoff doubles the current delay (capped at max) and adds up to 25% jitter.
|
|
func nextBackoff(current, maxDelay time.Duration) time.Duration {
|
|
next := current * 2
|
|
if next > maxDelay {
|
|
next = maxDelay
|
|
}
|
|
jitter := time.Duration(rand.Int64N(int64(next / 4)))
|
|
return next + jitter
|
|
}
|
|
|
|
// zeroBytes overwrites a byte slice with zeros.
|
|
// SECURITY: Go strings are immutable and cannot be zeroed in place. This utility
|
|
// is for zeroing byte slices (e.g., password buffers) to limit credential lifetime
|
|
// in memory. Credentials stored as Go strings (ssh.go, snmp.go, mikrotik.go)
|
|
// cannot benefit from this until the protocol layer supports []byte credentials.
|
|
func zeroBytes(b []byte) {
|
|
for i := range b {
|
|
b[i] = 0
|
|
}
|
|
}
|
|
|
|
func strPtr(s string) *string { return &s }
|
|
|
|
// executeCheck dispatches a check to the worker pool.
|
|
func executeCheck(ctx context.Context, check *pb.Check, pools *jobPools, checkResultCh chan<- *pb.CheckResult) {
|
|
ok := pools.checks.submit(ctx, func() {
|
|
result := ExecuteCheck(ctx, check)
|
|
select {
|
|
case checkResultCh <- result:
|
|
case <-ctx.Done():
|
|
}
|
|
})
|
|
if !ok {
|
|
slog.Warn("check rejected (pool full)", "check_id", check.Id, "type", check.CheckType)
|
|
}
|
|
}
|