diff --git a/agent.go b/agent.go index 1899259..b9c65ec 100644 --- a/agent.go +++ b/agent.go @@ -4,6 +4,7 @@ import ( "context" "encoding/base64" "encoding/json" + "errors" "fmt" "log/slog" "math/rand/v2" @@ -22,6 +23,8 @@ import ( var osExit = os.Exit var doSelfUpdate = selfUpdate +var errRestartRequested = fmt.Errorf("restart requested") + const maxJobPayloadBytes = 4 << 20 // 4 MB — well above any legitimate job list // channelMsg is the WebSocket channel message format (JSON wrapper around binary protobuf). @@ -49,6 +52,10 @@ func runAgent(ctx context.Context, wsURL, token string) { if ctx.Err() != nil { return } + if errors.Is(err, errRestartRequested) { + osExit(0) + return + } if err != nil { slog.Error("agent disconnected", "error", err) } @@ -239,7 +246,10 @@ func runSession(ctx context.Context, baseURL, token string) error { slog.Warn("invalid message", "error", err) continue } - handleMessage(ctx, msg, pools, snmpResultCh, mikrotikResultCh, credTestResultCh, monitoringCheckCh) + if handleMessage(ctx, msg, pools, snmpResultCh, mikrotikResultCh, credTestResultCh, monitoringCheckCh) { + flushSnmpBatch() + return errRestartRequested + } case result := <-snmpResultCh: snmpBatch = append(snmpBatch, result) @@ -290,6 +300,7 @@ func runSession(ctx context.Context, baseURL, token string) error { } // handleMessage dispatches incoming channel messages. +// Returns true if the session should end (e.g. restart requested). func handleMessage( ctx context.Context, msg channelMsg, @@ -298,7 +309,7 @@ func handleMessage( mikrotikResultCh chan<- *pb.MikrotikResult, credTestResultCh chan<- *pb.CredentialTestResult, monitoringCheckCh chan<- *pb.MonitoringCheck, -) { +) bool { switch msg.Event { case "phx_reply": slog.Debug("channel reply", "topic", msg.Topic) @@ -309,21 +320,21 @@ func handleMessage( } if err := json.Unmarshal(msg.Payload, &payload); err != nil { slog.Error("decode job payload", "error", err) - return + return false } if len(payload.Binary) > maxJobPayloadBytes { slog.Error("job payload too large", "size", len(payload.Binary), "max", maxJobPayloadBytes) - return + return false } bin, err := base64.StdEncoding.DecodeString(payload.Binary) if err != nil { slog.Error("decode base64", "error", err) - return + return false } var jobList pb.AgentJobList if err := proto.Unmarshal(bin, &jobList); err != nil { slog.Error("unmarshal job list", "error", err) - return + return false } slog.Info("received jobs", "count", len(jobList.Jobs)) for _, job := range jobList.Jobs { @@ -331,8 +342,8 @@ func handleMessage( } case "restart": - slog.Info("restart requested by server, exiting") - osExit(0) + slog.Info("restart requested by server") + return true case "update": var payload struct { @@ -341,7 +352,7 @@ func handleMessage( } if err := json.Unmarshal(msg.Payload, &payload); err != nil || payload.URL == "" || payload.Checksum == "" { slog.Error("invalid update payload") - return + return false } slog.Info("update requested", "url", payload.URL) if err := doSelfUpdate(payload.URL, payload.Checksum); err != nil { @@ -351,6 +362,7 @@ func handleMessage( default: slog.Debug("ignoring event", "event", msg.Event) } + return false } // jobPools holds the worker pools for each job type. diff --git a/agent_test.go b/agent_test.go index bbdf277..7ae7c36 100644 --- a/agent_test.go +++ b/agent_test.go @@ -165,20 +165,14 @@ func TestHandleMessage(t *testing.T) { }) t.Run("restart", func(t *testing.T) { - origExit := osExit - defer func() { osExit = origExit }() - - var exitCode int - osExit = func(code int) { exitCode = code } - snmpCh := make(chan *pb.SnmpResult, 1) mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) - handleMessage(context.Background(), channelMsg{Event: "restart", Payload: json.RawMessage(`{}`)}, testPools(t), snmpCh, mtCh, credCh, monCh) + shouldEnd := handleMessage(context.Background(), channelMsg{Event: "restart", Payload: json.RawMessage(`{}`)}, testPools(t), snmpCh, mtCh, credCh, monCh) - if exitCode != 0 { - t.Errorf("expected exit code 0, got %d", exitCode) + if !shouldEnd { + t.Error("expected handleMessage to return true for restart") } })