package dockeragent import ( "bytes" "context" "encoding/json" "errors" "io" "math" "net/http" "net/http/httptest" "os" "path/filepath" "runtime" "strings" "testing" "time" containertypes "github.com/moby/moby/api/types/container" systemtypes "github.com/moby/moby/api/types/system" agentsdocker "github.com/rcourtman/pulse-go-rewrite/pkg/agents/docker" "github.com/rs/zerolog" ) func TestSendReport(t *testing.T) { t.Run("marshal error", func(t *testing.T) { agent := &Agent{logger: zerolog.Nop()} report := agentsdocker.Report{ Host: agentsdocker.HostInfo{ CPUUsagePercent: math.NaN(), }, } if err := agent.sendReport(context.Background(), report); err == nil { t.Fatal("expected marshal error") } }) t.Run("stop requested", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusBadRequest) _, _ = w.Write([]byte(`{"error":"host was removed","code":"invalid_report"}`)) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", targets: []TargetConfig{{URL: server.URL, Token: "token"}}, httpClients: map[bool]*http.Client{ false: server.Client(), }, } if err := agent.sendReport(context.Background(), agentsdocker.Report{}); !errors.Is(err, ErrStopRequested) { t.Fatalf("expected ErrStopRequested, got %v", err) } }) t.Run("errors join", func(t *testing.T) { agent := &Agent{ logger: zerolog.Nop(), targets: []TargetConfig{{URL: "http://one", Token: "t1"}, {URL: "http://two", Token: "t2"}}, httpClients: map[bool]*http.Client{ false: {Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { return nil, errors.New("send failed") })}, }, } if err := agent.sendReport(context.Background(), agentsdocker.Report{}); err == nil { t.Fatal("expected error") } }) t.Run("large payload succeeds", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), targets: []TargetConfig{{URL: server.URL, Token: "token"}}, httpClients: map[bool]*http.Client{ false: server.Client(), }, } report := agentsdocker.Report{ Containers: []agentsdocker.Container{ {ID: strings.Repeat("a", 500000)}, }, } if err := agent.sendReport(context.Background(), report); err != nil { t.Fatalf("unexpected error: %v", err) } }) } func TestSendReportToTarget(t *testing.T) { t.Run("request error", func(t *testing.T) { agent := &Agent{logger: zerolog.Nop()} if err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: "http://example.com/\x7f"}, []byte(`{}`), 0); err == nil { t.Fatal("expected error") } }) t.Run("host removed", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusBadRequest) _, _ = w.Write([]byte(`{"error":"host was removed","code":"invalid_report"}`)) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, } err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, []byte(`{}`), 0) if !errors.Is(err, ErrStopRequested) { t.Fatalf("expected ErrStopRequested, got %v", err) } }) t.Run("command continue on nil error", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"commands":[{"id":"cmd1","type":"unknown"}]}`)) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), httpClients: map[bool]*http.Client{ false: server.Client(), }, } if err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, []byte(`{}`), 0); err != nil { t.Fatalf("unexpected error: %v", err) } }) t.Run("status error", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusBadRequest) _, _ = w.Write([]byte("bad request")) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), httpClients: map[bool]*http.Client{ false: server.Client(), }, } if err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, []byte(`{}`), 0); err == nil { t.Fatal("expected error") } }) t.Run("status error body too large", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusBadRequest) _, _ = w.Write([]byte(strings.Repeat("x", maxPulseResponseBodyBytes+1))) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), httpClients: map[bool]*http.Client{ false: server.Client(), }, } err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, []byte(`{}`), 0) if err == nil || !strings.Contains(err.Error(), "read error response") { t.Fatalf("expected oversized error response failure, got %v", err) } }) t.Run("status error with empty body", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusInternalServerError) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), httpClients: map[bool]*http.Client{ false: server.Client(), }, } if err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, []byte(`{}`), 0); err == nil { t.Fatal("expected error") } }) t.Run("read error", func(t *testing.T) { client := &http.Client{ Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { return &http.Response{ StatusCode: http.StatusOK, Body: errReadCloser{err: errors.New("read failed")}, Header: make(http.Header), }, nil }), } agent := &Agent{ logger: zerolog.Nop(), httpClients: map[bool]*http.Client{ false: client, }, } if err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: "http://example", Token: "token"}, []byte(`{}`), 0); err == nil { t.Fatal("expected error") } }) t.Run("success body too large", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte(strings.Repeat("x", maxPulseResponseBodyBytes+1))) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), httpClients: map[bool]*http.Client{ false: server.Client(), }, } err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, []byte(`{}`), 0) if err == nil || !strings.Contains(err.Error(), "read response") { t.Fatalf("expected oversized response failure, got %v", err) } }) t.Run("invalid json response", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte("{")) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), httpClients: map[bool]*http.Client{ false: server.Client(), }, } if err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, []byte(`{}`), 0); err != nil { t.Fatalf("unexpected error: %v", err) } }) t.Run("empty response body", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), httpClients: map[bool]*http.Client{ false: server.Client(), }, } if err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, []byte(`{}`), 0); err != nil { t.Fatalf("unexpected error: %v", err) } }) t.Run("stop command", func(t *testing.T) { prevPath := os.Getenv("PATH") _ = os.Setenv("PATH", "") t.Cleanup(func() { _ = os.Setenv("PATH", prevPath) }) var ackBody bytes.Buffer server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { switch { case strings.HasSuffix(r.URL.Path, "/report"): w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"commands":[{"id":"cmd1","type":"stop"}]}`)) case strings.Contains(r.URL.Path, "/commands/"): body, _ := io.ReadAll(r.Body) ackBody.Write(body) w.WriteHeader(http.StatusOK) default: w.WriteHeader(http.StatusNotFound) } })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, } err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, []byte(`{}`), 0) if !errors.Is(err, ErrStopRequested) { t.Fatalf("expected ErrStopRequested, got %v", err) } }) t.Run("command error bubbles", func(t *testing.T) { prevPath := os.Getenv("PATH") _ = os.Setenv("PATH", "") t.Cleanup(func() { _ = os.Setenv("PATH", prevPath) }) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { switch { case strings.HasSuffix(r.URL.Path, "/report"): w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"commands":[{"id":"cmd1","type":"stop"}]}`)) case strings.Contains(r.URL.Path, "/commands/"): w.WriteHeader(http.StatusInternalServerError) _, _ = w.Write([]byte("boom")) default: w.WriteHeader(http.StatusNotFound) } })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, } if err := agent.sendReportToTarget(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, []byte(`{}`), 0); err == nil { t.Fatal("expected error") } }) } func TestSendCommandAck(t *testing.T) { t.Run("missing host id", func(t *testing.T) { agent := &Agent{} if err := agent.sendCommandAck(context.Background(), TargetConfig{URL: "http://example"}, "cmd", "status", "msg"); err == nil { t.Fatal("expected error") } }) t.Run("marshal error", func(t *testing.T) { agent := &Agent{ hostID: "host1", jsonMarshalFn: func(any) ([]byte, error) { return nil, errors.New("marshal failed") }, } if err := agent.sendCommandAck(context.Background(), TargetConfig{URL: "http://example"}, "cmd", "status", "msg"); err == nil { t.Fatal("expected error") } }) t.Run("request error", func(t *testing.T) { agent := &Agent{hostID: "host1"} badURL := "http://example.com/\x7f" if err := agent.sendCommandAck(context.Background(), TargetConfig{URL: badURL}, "cmd", "status", "msg"); err == nil { t.Fatal("expected error") } }) t.Run("client error", func(t *testing.T) { agent := &Agent{ hostID: "host1", httpClients: map[bool]*http.Client{ false: {Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { return nil, errors.New("send failed") })}, }, } if err := agent.sendCommandAck(context.Background(), TargetConfig{URL: "http://example", Token: "token"}, "cmd", "status", "msg"); err == nil { t.Fatal("expected error") } }) t.Run("status error", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusInternalServerError) _, _ = w.Write([]byte("boom")) })) defer server.Close() agent := &Agent{ hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, } if err := agent.sendCommandAck(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, "cmd", "status", "msg"); err == nil { t.Fatal("expected error") } }) t.Run("status error body too large", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusInternalServerError) _, _ = w.Write([]byte(strings.Repeat("x", maxPulseResponseBodyBytes+1))) })) defer server.Close() agent := &Agent{ hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, } err := agent.sendCommandAck(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, "cmd", "status", "msg") if err == nil || !strings.Contains(err.Error(), "read acknowledgement error response") { t.Fatalf("expected oversized acknowledgement response failure, got %v", err) } }) t.Run("success", func(t *testing.T) { var got agentsdocker.CommandAck server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { body, _ := io.ReadAll(r.Body) _ = json.Unmarshal(body, &got) w.WriteHeader(http.StatusOK) })) defer server.Close() agent := &Agent{ hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, } if err := agent.sendCommandAck(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, "cmd", "completed", "ok"); err != nil { t.Fatalf("unexpected error: %v", err) } if got.Status != "completed" { t.Fatalf("expected status to be sent, got %q", got.Status) } }) } func TestHandleCommand(t *testing.T) { agent := &Agent{logger: zerolog.Nop()} if err := agent.handleCommand(context.Background(), TargetConfig{}, agentsdocker.Command{Type: "unknown"}); err != nil { t.Fatalf("unexpected error: %v", err) } t.Run("stop command", func(t *testing.T) { prev := os.Getenv("PATH") _ = os.Setenv("PATH", "") t.Cleanup(func() { _ = os.Setenv("PATH", prev) }) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, } t.Cleanup(func() { _ = agent.Close() }) err := agent.handleCommand(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, agentsdocker.Command{ID: "cmd", Type: agentsdocker.CommandTypeStop}) if !errors.Is(err, ErrStopRequested) { t.Fatalf("expected ErrStopRequested, got %v", err) } }) t.Run("update command", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, docker: &fakeDockerClient{ containerInspectFn: func(context.Context, string) (containertypes.InspectResponse, error) { return containertypes.InspectResponse{}, errors.New("inspect failed") }, }, } cmd := agentsdocker.Command{ ID: "cmd2", Type: agentsdocker.CommandTypeUpdateContainer, Payload: map[string]any{ "containerId": "container1", }, } if err := agent.handleCommand(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, cmd); err != nil { t.Fatalf("unexpected error: %v", err) } }) t.Run("check updates command", func(t *testing.T) { collectAttempted := make(chan struct{}, 1) ackPath := make(chan string, 1) registryChecker := NewRegistryChecker(zerolog.Nop()) registryChecker.MarkChecked() registryChecker.cacheDigest("cached-key", "sha256:cached") server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { ackPath <- r.URL.Path w.WriteHeader(http.StatusOK) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", registryChecker: registryChecker, httpClients: map[bool]*http.Client{ false: server.Client(), }, docker: &fakeDockerClient{ infoFunc: func(context.Context) (systemtypes.Info, error) { select { case collectAttempted <- struct{}{}: default: } return systemtypes.Info{}, errors.New("info failed") }, }, } t.Cleanup(func() { _ = agent.Close() }) cmd := agentsdocker.Command{ ID: "cmd3", Type: agentsdocker.CommandTypeCheckUpdates, } if err := agent.handleCommand(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, cmd); err != nil { t.Fatalf("unexpected error: %v", err) } select { case gotPath := <-ackPath: if !strings.HasSuffix(gotPath, "/commands/cmd3/ack") { t.Fatalf("unexpected ack path: %s", gotPath) } case <-time.After(500 * time.Millisecond): t.Fatal("expected check-updates acknowledgement request") } registryChecker.mu.RLock() lastFullCheck := registryChecker.lastFullCheck registryChecker.mu.RUnlock() if !lastFullCheck.IsZero() { t.Fatalf("expected ForceCheck to reset lastFullCheck, got %s", lastFullCheck) } registryChecker.cache.mu.RLock() cacheLen := len(registryChecker.cache.entries) registryChecker.cache.mu.RUnlock() if cacheLen != 0 { t.Fatalf("expected ForceCheck to clear cache, found %d entries", cacheLen) } select { case <-collectAttempted: case <-time.After(500 * time.Millisecond): t.Fatal("expected check-updates command to trigger immediate collection") } }) t.Run("check updates command ack error does not propagate", func(t *testing.T) { registryChecker := NewRegistryChecker(zerolog.Nop()) registryChecker.MarkChecked() registryChecker.cacheDigest("cached-key", "sha256:cached") agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", registryChecker: registryChecker, manualCheckCollect: func(context.Context) (agentsdocker.Report, error) { return agentsdocker.Report{}, nil }, } t.Cleanup(func() { _ = agent.Close() }) err := agent.handleCheckUpdatesCommand(context.Background(), TargetConfig{URL: "http://example.com/\x7f", Token: "token"}, agentsdocker.Command{ ID: "cmd4", Type: agentsdocker.CommandTypeCheckUpdates, }) if err != nil { t.Fatalf("expected nil error on ack failure, got: %v", err) } registryChecker.mu.RLock() lastFullCheck := registryChecker.lastFullCheck registryChecker.mu.RUnlock() if !lastFullCheck.IsZero() { t.Fatalf("expected ForceCheck to reset lastFullCheck, got %s", lastFullCheck) } }) } func TestMonitoringOnlyCollectorRejectsForgedReportCommand(t *testing.T) { var ack agentsdocker.CommandAck server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if strings.HasSuffix(r.URL.Path, "/report") { w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"commands":[{"id":"forged-update","type":"update_container","payload":{"containerId":"container-1"}}]}`)) return } if strings.HasSuffix(r.URL.Path, "/commands/forged-update/ack") { if err := json.NewDecoder(r.Body).Decode(&ack); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } w.WriteHeader(http.StatusNoContent) return } http.NotFound(w, r) })) defer server.Close() mutations := 0 helper := &helperInventoryStub{} agent := &Agent{ cfg: Config{HelperInventory: helper}, hostID: "safe-host", docker: &fakeDockerClient{ containerInspectFn: func(context.Context, string) (containertypes.InspectResponse, error) { mutations++ return containertypes.InspectResponse{}, errors.New("must not execute") }, }, httpClients: map[bool]*http.Client{false: server.Client()}, logger: zerolog.Nop(), } target := TargetConfig{Name: "primary", URL: server.URL, Token: "token", Authoritative: true} if err := agent.sendReportToTarget(context.Background(), target, []byte("ignored"), 0); err != nil { t.Fatalf("safe report response: %v", err) } if mutations != 0 { t.Fatalf("forged report command reached direct runtime %d times", mutations) } if ack.Status != agentsdocker.CommandStatusFailed || !strings.Contains(ack.Message, "Monitoring-only collector") { t.Fatalf("rejected command acknowledgement = %+v", ack) } } func TestHandleStopCommand(t *testing.T) { if runtime.GOOS == "windows" { t.Skip("systemd stop-command behavior is not available on Windows") } t.Run("disable error sends failure ack", func(t *testing.T) { writeSystemctl(t, "echo 'access denied' >&2\nexit 1") var ack agentsdocker.CommandAck server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { body, _ := io.ReadAll(r.Body) _ = json.Unmarshal(body, &ack) w.WriteHeader(http.StatusOK) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, } t.Cleanup(func() { _ = agent.Close() }) if err := agent.handleStopCommand(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, agentsdocker.Command{ID: "cmd"}); err != nil { t.Fatalf("unexpected error: %v", err) } if ack.Status != agentsdocker.CommandStatusFailed { t.Fatalf("expected failed status, got %q", ack.Status) } }) t.Run("disable error ack failure", func(t *testing.T) { writeSystemctl(t, "echo 'access denied' >&2\nexit 1") agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", } t.Cleanup(func() { _ = agent.Close() }) if err := agent.handleStopCommand(context.Background(), TargetConfig{URL: "http://example.com/\x7f", Token: "token"}, agentsdocker.Command{ID: "cmd"}); err != nil { t.Fatalf("unexpected error: %v", err) } }) t.Run("success returns stop requested", func(t *testing.T) { prev := os.Getenv("PATH") _ = os.Setenv("PATH", "") t.Cleanup(func() { _ = os.Setenv("PATH", prev) }) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, } t.Cleanup(func() { _ = agent.Close() }) if err := agent.handleStopCommand(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, agentsdocker.Command{ID: "cmd"}); !errors.Is(err, ErrStopRequested) { t.Fatalf("expected ErrStopRequested, got %v", err) } }) t.Run("completion ack error", func(t *testing.T) { prev := os.Getenv("PATH") _ = os.Setenv("PATH", "") t.Cleanup(func() { _ = os.Setenv("PATH", prev) }) agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", httpClients: map[bool]*http.Client{ false: {Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { return nil, errors.New("send failed") })}, }, } t.Cleanup(func() { _ = agent.Close() }) if err := agent.handleStopCommand(context.Background(), TargetConfig{URL: "http://example", Token: "token"}, agentsdocker.Command{ID: "cmd"}); err == nil { t.Fatal("expected error") } }) t.Run("stop service goroutine executes", func(t *testing.T) { marker := filepath.Join(t.TempDir(), "called") writeSystemctl(t, "if [ \"$1\" = \"disable\" ]; then exit 0; fi\nif [ \"$1\" = \"stop\" ]; then : > "+marker+"; exit 2; fi\nexit 0") server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) })) defer server.Close() agent := &Agent{ logger: zerolog.Nop(), hostID: "host1", httpClients: map[bool]*http.Client{ false: server.Client(), }, newTimerFn: immediateTimer, } // Close joins the stop-service goroutine before the fake systemctl is cleaned up t.Cleanup(func() { _ = agent.Close() }) if err := agent.handleStopCommand(context.Background(), TargetConfig{URL: server.URL, Token: "token"}, agentsdocker.Command{ID: "cmd"}); !errors.Is(err, ErrStopRequested) { t.Fatalf("expected ErrStopRequested, got %v", err) } deadline := time.Now().Add(200 * time.Millisecond) for { if _, err := os.Stat(marker); err == nil { break } if time.Now().After(deadline) { t.Fatal("expected stopSystemdService to be invoked") } time.Sleep(5 * time.Millisecond) } }) } func TestDisableSelf(t *testing.T) { prev := os.Getenv("PATH") _ = os.Setenv("PATH", "") t.Cleanup(func() { _ = os.Setenv("PATH", prev) }) baseDir := t.TempDir() scriptDir := filepath.Join(baseDir, "script") if err := os.MkdirAll(scriptDir, 0700); err != nil { t.Fatalf("mkdir: %v", err) } if err := os.WriteFile(filepath.Join(scriptDir, "file"), []byte("x"), 0600); err != nil { t.Fatalf("write: %v", err) } logDir := filepath.Join(baseDir, "logs") if err := os.MkdirAll(logDir, 0700); err != nil { t.Fatalf("mkdir: %v", err) } if err := os.WriteFile(filepath.Join(logDir, "file"), []byte("x"), 0600); err != nil { t.Fatalf("write: %v", err) } swap(t, &unraidStartupScriptPath, scriptDir) swap(t, &agentLogPath, logDir) agent := &Agent{logger: zerolog.Nop()} if err := agent.disableSelf(context.Background()); err != nil { t.Fatalf("unexpected error: %v", err) } }