From ae1c39960f964a5ec352be4e1d43c2a8cc9ba837 Mon Sep 17 00:00:00 2001 From: rcourtman Date: Mon, 29 Dec 2025 09:18:05 +0000 Subject: [PATCH] fix: Remove duplicate AI chat response streaming (issue #947) Content was being streamed twice: 1. During each iteration of the tool loop (intended for intermediate feedback) 2. Again after the loop ended with finalContent (redundant) This caused duplicate responses when using Ollama and other providers. --- internal/ai/service.go | 4 +- internal/dockeragent/agent.go | 6 +- internal/dockeragent/agent_collect_test.go | 12 +- internal/dockeragent/agent_cpu_test.go | 22 +- internal/dockeragent/agent_flow_test.go | 72 ++ internal/dockeragent/agent_http_test.go | 6 + internal/dockeragent/container_update_test.go | 5 +- internal/dockeragent/deps.go | 22 +- internal/dockeragent/self_update_test.go | 912 ++++++++++++++++++ internal/dockeragent/swarm_coverage_test.go | 13 +- 10 files changed, 1033 insertions(+), 41 deletions(-) create mode 100644 internal/dockeragent/self_update_test.go diff --git a/internal/ai/service.go b/internal/ai/service.go index ce88ab0ff..52f4b8fa0 100644 --- a/internal/ai/service.go +++ b/internal/ai/service.go @@ -1730,8 +1730,8 @@ Always execute the commands rather than telling the user how to do it.` } } - // Stream the final content - callback(StreamEvent{Type: "content", Data: finalContent}) + // Don't stream finalContent here - it was already streamed in the iteration above + // Sending it again causes duplicate responses (issue #947) callback(StreamEvent{Type: "done"}) return &ExecuteResponse{ diff --git a/internal/dockeragent/agent.go b/internal/dockeragent/agent.go index 5a2acab39..ca16b0f4e 100644 --- a/internal/dockeragent/agent.go +++ b/internal/dockeragent/agent.go @@ -1248,12 +1248,12 @@ func (a *Agent) disableSelf(ctx context.Context) error { } // Remove Unraid startup script if present to prevent restart on reboot. - if err := removeFileIfExists("/boot/config/go.d/pulse-docker-agent.sh"); err != nil { + if err := removeFileIfExists(unraidStartupScriptPath); err != nil { a.logger.Warn().Err(err).Msg("Failed to remove Unraid startup script") } // Best-effort log cleanup (ignore errors). - _ = removeFileIfExists("/var/log/pulse-docker-agent.log") + _ = removeFileIfExists(agentLogPath) return nil } @@ -1632,7 +1632,7 @@ func randomDuration(max time.Duration) time.Duration { return 0 } - n, err := rand.Int(rand.Reader, big.NewInt(int64(max))) + n, err := randIntFn(rand.Reader, big.NewInt(int64(max))) if err != nil { return 0 } diff --git a/internal/dockeragent/agent_collect_test.go b/internal/dockeragent/agent_collect_test.go index 149354155..c77892056 100644 --- a/internal/dockeragent/agent_collect_test.go +++ b/internal/dockeragent/agent_collect_test.go @@ -76,7 +76,7 @@ func TestCollectContainer(t *testing.T) { } summary := containertypes.Summary{ - ID: "container1", + ID: "container-123456", Names: []string{"/app"}, Image: "nginx@sha256:abc123", ImageID: "sha256:local", @@ -119,7 +119,7 @@ func TestCollectContainer(t *testing.T) { agent := &Agent{ logger: logger, prevContainerCPU: map[string]cpuSample{ - "container1": {totalUsage: 1}, + "container-123456": {totalUsage: 1}, }, docker: &fakeDockerClient{ containerInspectWithRawFn: func(context.Context, string, bool) (containertypes.InspectResponse, []byte, error) { @@ -130,7 +130,7 @@ func TestCollectContainer(t *testing.T) { }, } - summary := containertypes.Summary{ID: "container1", Names: []string{"/app"}, State: "exited"} + summary := containertypes.Summary{ID: "container-123456", Names: []string{"/app"}, State: "exited"} if _, err := agent.collectContainer(context.Background(), summary); err != nil { t.Fatalf("unexpected error: %v", err) } @@ -149,7 +149,7 @@ func TestCollectContainer(t *testing.T) { logger: logger, } - if _, err := agent.collectContainer(context.Background(), containertypes.Summary{ID: "container1"}); err == nil { + if _, err := agent.collectContainer(context.Background(), containertypes.Summary{ID: "container-123456"}); err == nil { t.Fatal("expected error") } }) @@ -169,7 +169,7 @@ func TestCollectContainer(t *testing.T) { logger: logger, } - if _, err := agent.collectContainer(context.Background(), containertypes.Summary{ID: "container1"}); err == nil { + if _, err := agent.collectContainer(context.Background(), containertypes.Summary{ID: "container-123456"}); err == nil { t.Fatal("expected error") } }) @@ -189,7 +189,7 @@ func TestCollectContainer(t *testing.T) { logger: logger, } - if _, err := agent.collectContainer(context.Background(), containertypes.Summary{ID: "container1"}); err == nil { + if _, err := agent.collectContainer(context.Background(), containertypes.Summary{ID: "container-123456"}); err == nil { t.Fatal("expected error") } }) diff --git a/internal/dockeragent/agent_cpu_test.go b/internal/dockeragent/agent_cpu_test.go index 0dcb227d5..c1cbb1d04 100644 --- a/internal/dockeragent/agent_cpu_test.go +++ b/internal/dockeragent/agent_cpu_test.go @@ -31,11 +31,11 @@ func TestCalculateContainerCPUPercent(t *testing.T) { }, } - got := agent.calculateContainerCPUPercent("container1", stats) + got := agent.calculateContainerCPUPercent("container-123456", stats) if got <= 0 { t.Fatalf("expected percent > 0, got %f", got) } - if _, ok := agent.prevContainerCPU["container1"]; !ok { + if _, ok := agent.prevContainerCPU["container-123456"]; !ok { t.Fatal("expected current sample to be stored") } }) @@ -57,11 +57,11 @@ func TestCalculateContainerCPUPercent(t *testing.T) { }, } - got := agent.calculateContainerCPUPercent("container1", stats) + got := agent.calculateContainerCPUPercent("container-123456", stats) if got != 0 { t.Fatalf("expected 0, got %f", got) } - if _, ok := agent.prevContainerCPU["container1"]; !ok { + if _, ok := agent.prevContainerCPU["container-123456"]; !ok { t.Fatal("expected sample to be stored") } }) @@ -70,7 +70,7 @@ func TestCalculateContainerCPUPercent(t *testing.T) { agent := &Agent{ logger: logger, prevContainerCPU: map[string]cpuSample{ - "container1": { + "container-123456": { totalUsage: 100, systemUsage: 1000, onlineCPUs: 2, @@ -89,7 +89,7 @@ func TestCalculateContainerCPUPercent(t *testing.T) { PreCPUStats: containertypes.CPUStats{}, } - got := agent.calculateContainerCPUPercent("container1", stats) + got := agent.calculateContainerCPUPercent("container-123456", stats) if got <= 0 { t.Fatalf("expected percent > 0, got %f", got) } @@ -100,7 +100,7 @@ func TestCalculateContainerCPUPercent(t *testing.T) { logger: logger, cpuCount: 4, prevContainerCPU: map[string]cpuSample{ - "container1": { + "container-123456": { totalUsage: 100, systemUsage: 1000, onlineCPUs: 0, @@ -119,7 +119,7 @@ func TestCalculateContainerCPUPercent(t *testing.T) { PreCPUStats: containertypes.CPUStats{}, } - got := agent.calculateContainerCPUPercent("container1", stats) + got := agent.calculateContainerCPUPercent("container-123456", stats) if got <= 0 { t.Fatalf("expected percent > 0, got %f", got) } @@ -129,7 +129,7 @@ func TestCalculateContainerCPUPercent(t *testing.T) { agent := &Agent{ logger: logger, prevContainerCPU: map[string]cpuSample{ - "container1": { + "container-123456": { totalUsage: 100, systemUsage: 1000, onlineCPUs: 0, @@ -148,7 +148,7 @@ func TestCalculateContainerCPUPercent(t *testing.T) { PreCPUStats: containertypes.CPUStats{}, } - got := agent.calculateContainerCPUPercent("container1", stats) + got := agent.calculateContainerCPUPercent("container-123456", stats) if got != 0 { t.Fatalf("expected 0, got %f", got) } @@ -170,7 +170,7 @@ func TestCalculateContainerCPUPercent(t *testing.T) { } for i := 0; i < 10; i++ { - _ = agent.calculateContainerCPUPercent("container1", stats) + _ = agent.calculateContainerCPUPercent("container-123456", stats) } }) } diff --git a/internal/dockeragent/agent_flow_test.go b/internal/dockeragent/agent_flow_test.go index e3d0e5201..b1380f1cd 100644 --- a/internal/dockeragent/agent_flow_test.go +++ b/internal/dockeragent/agent_flow_test.go @@ -97,6 +97,7 @@ func TestStopTimer(t *testing.T) { t.Run("timer fired and drained", func(t *testing.T) { timer := time.NewTimer(0) + time.Sleep(time.Millisecond) stopTimer(timer) select { case <-timer.C: @@ -220,3 +221,74 @@ func TestCollectOnce(t *testing.T) { } }) } + +func TestRun(t *testing.T) { + t.Run("stop requested on startup", func(t *testing.T) { + swap(t, &connectRuntimeFn, func(_ RuntimeKind, _ *zerolog.Logger) (dockerClient, systemtypes.Info, RuntimeKind, error) { + return &fakeDockerClient{ + infoFunc: func(context.Context) (systemtypes.Info, error) { + return systemtypes.Info{}, ErrStopRequested + }, + }, systemtypes.Info{}, RuntimeDocker, nil + }) + + agent := &Agent{ + cfg: Config{ + Interval: 10 * time.Millisecond, + }, + docker: &fakeDockerClient{ + infoFunc: func(context.Context) (systemtypes.Info, error) { + return systemtypes.Info{}, ErrStopRequested + }, + }, + logger: zerolog.Nop(), + } + + if err := agent.Run(context.Background()); err != nil { + t.Fatalf("expected nil, got %v", err) + } + }) + + t.Run("ticker and update timer", func(t *testing.T) { + swap(t, &randomDurationFn, func(time.Duration) time.Duration { + return -5 * time.Second + }) + swap(t, &hostmetricsCollect, func(context.Context, []string) (hostmetrics.Snapshot, error) { + return hostmetrics.Snapshot{}, nil + }) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + })) + defer server.Close() + + agent := &Agent{ + cfg: Config{ + Interval: 5 * time.Millisecond, + DisableAutoUpdate: true, + }, + docker: &fakeDockerClient{ + infoFunc: func(context.Context) (systemtypes.Info, error) { + return systemtypes.Info{ID: "daemon", ServerVersion: "24.0.0"}, nil + }, + }, + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: server.URL, Token: "token"}}, + httpClients: map[bool]*http.Client{false: server.Client()}, + reportBuffer: buffer.New[agentsdocker.Report](10), + } + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + done <- agent.Run(ctx) + }() + + time.Sleep(20 * time.Millisecond) + cancel() + + if err := <-done; !errors.Is(err, context.Canceled) { + t.Fatalf("expected context canceled, got %v", err) + } + }) +} diff --git a/internal/dockeragent/agent_http_test.go b/internal/dockeragent/agent_http_test.go index 073ad8c9d..edc786740 100644 --- a/internal/dockeragent/agent_http_test.go +++ b/internal/dockeragent/agent_http_test.go @@ -195,6 +195,12 @@ func TestSendReportToTarget(t *testing.T) { }) 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 { diff --git a/internal/dockeragent/container_update_test.go b/internal/dockeragent/container_update_test.go index 850d3f088..df392352a 100644 --- a/internal/dockeragent/container_update_test.go +++ b/internal/dockeragent/container_update_test.go @@ -34,7 +34,7 @@ func baseInspect() containertypes.InspectResponse { Config: &containertypes.Config{ Image: "nginx:latest", }, - NetworkSettings: &network.NetworkSettings{ + NetworkSettings: &containertypes.NetworkSettings{ Networks: map[string]*network.EndpointSettings{ "net1": {Aliases: []string{"app"}}, "net2": {Aliases: []string{"app2"}}, @@ -189,9 +189,6 @@ func TestUpdateContainer_Errors(t *testing.T) { containerStopFn: func(context.Context, string, containertypes.StopOptions) error { return nil }, - containerRenameFn: func(context.Context, string, string) error { - return nil - }, containerCreateFn: func(context.Context, *containertypes.Config, *containertypes.HostConfig, *network.NetworkingConfig, *v1.Platform, string) (containertypes.CreateResponse, error) { return containertypes.CreateResponse{ID: "new123"}, nil }, diff --git a/internal/dockeragent/deps.go b/internal/dockeragent/deps.go index 4522b60c0..3e804d0d5 100644 --- a/internal/dockeragent/deps.go +++ b/internal/dockeragent/deps.go @@ -2,6 +2,7 @@ package dockeragent import ( "context" + "crypto/rand" "io" "os" "os/exec" @@ -14,15 +15,16 @@ import ( ) var ( - connectRuntimeFn = connectRuntime - hostmetricsCollect = hostmetrics.Collect - newTickerFn = time.NewTicker - newTimerFn = time.NewTimer - randomDurationFn = randomDuration - nowFn = time.Now - sleepFn = time.Sleep + connectRuntimeFn = connectRuntime + hostmetricsCollect = hostmetrics.Collect + newTickerFn = time.NewTicker + newTimerFn = time.NewTimer + randomDurationFn = randomDuration + nowFn = time.Now + sleepFn = time.Sleep buildRuntimeCandidatesFn = buildRuntimeCandidates tryRuntimeCandidateFn = tryRuntimeCandidate + randIntFn = rand.Int osExecutableFn = os.Executable osCreateTempFn = os.CreateTemp closeFileFn = func(f *os.File) error { return f.Close() } @@ -45,8 +47,10 @@ var ( "/etc/machine-id", "/var/lib/dbus/machine-id", } - unraidVersionPath = "/etc/unraid-version" - unraidPersistPath = "/boot/config/plugins/pulse-docker-agent/pulse-docker-agent" + unraidVersionPath = "/etc/unraid-version" + unraidPersistPath = "/boot/config/plugins/pulse-docker-agent/pulse-docker-agent" + unraidStartupScriptPath = "/boot/config/go.d/pulse-docker-agent.sh" + agentLogPath = "/var/log/pulse-docker-agent.log" openProcUptime = func() (io.ReadCloser, error) { return os.Open("/proc/uptime") } diff --git a/internal/dockeragent/self_update_test.go b/internal/dockeragent/self_update_test.go new file mode 100644 index 000000000..8523af339 --- /dev/null +++ b/internal/dockeragent/self_update_test.go @@ -0,0 +1,912 @@ +package dockeragent + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + + "github.com/rs/zerolog" +) + +func TestDetermineSelfUpdateArch_Coverage(t *testing.T) { + t.Run("known arches", func(t *testing.T) { + swap(t, &goArch, "amd64") + if got := determineSelfUpdateArch(); got != "linux-amd64" { + t.Fatalf("expected linux-amd64, got %q", got) + } + + swap(t, &goArch, "arm64") + if got := determineSelfUpdateArch(); got != "linux-arm64" { + t.Fatalf("expected linux-arm64, got %q", got) + } + + swap(t, &goArch, "arm") + if got := determineSelfUpdateArch(); got != "linux-armv7" { + t.Fatalf("expected linux-armv7, got %q", got) + } + }) + + t.Run("uname fallback", func(t *testing.T) { + swap(t, &goArch, "other") + swap(t, &unameMachine, func() (string, error) { + return "x86_64", nil + }) + if got := determineSelfUpdateArch(); got != "linux-amd64" { + t.Fatalf("expected linux-amd64, got %q", got) + } + }) + + t.Run("uname error", func(t *testing.T) { + swap(t, &goArch, "other") + swap(t, &unameMachine, func() (string, error) { + return "", errors.New("boom") + }) + if got := determineSelfUpdateArch(); got != "" { + t.Fatalf("expected empty result, got %q", got) + } + }) +} + +func TestResolveSymlink(t *testing.T) { + dir := t.TempDir() + target := filepath.Join(dir, "target") + if err := os.WriteFile(target, []byte("data"), 0600); err != nil { + t.Fatalf("write target: %v", err) + } + link := filepath.Join(dir, "link") + if err := os.Symlink(target, link); err != nil { + t.Fatalf("symlink: %v", err) + } + + got, err := resolveSymlink(link) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != target { + t.Fatalf("expected %q, got %q", target, got) + } + + if _, err := resolveSymlink(filepath.Join(dir, "missing")); err == nil { + t.Fatal("expected error for missing symlink") + } +} + +func TestVerifyELFMagic(t *testing.T) { + dir := t.TempDir() + valid := filepath.Join(dir, "valid") + if err := os.WriteFile(valid, []byte{0x7f, 'E', 'L', 'F', 0x01}, 0600); err != nil { + t.Fatalf("write valid: %v", err) + } + if err := verifyELFMagic(valid); err != nil { + t.Fatalf("expected valid ELF, got %v", err) + } + + invalid := filepath.Join(dir, "invalid") + if err := os.WriteFile(invalid, []byte("nope"), 0600); err != nil { + t.Fatalf("write invalid: %v", err) + } + if err := verifyELFMagic(invalid); err == nil { + t.Fatal("expected error for invalid magic") + } +} + +func TestCheckForUpdates(t *testing.T) { + t.Run("dev version skips", func(t *testing.T) { + swap(t, &Version, "dev") + agent := &Agent{logger: zerolog.Nop()} + agent.checkForUpdates(context.Background()) + }) + + t.Run("no target skips", func(t *testing.T) { + swap(t, &Version, "1.0.0") + agent := &Agent{logger: zerolog.Nop()} + agent.checkForUpdates(context.Background()) + }) + + t.Run("request creation error", func(t *testing.T) { + swap(t, &Version, "1.0.0") + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com/\x7f"}}, + } + agent.checkForUpdates(context.Background()) + }) + + t.Run("http error", func(t *testing.T) { + swap(t, &Version, "1.0.0") + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return nil, errors.New("boom") + })} + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + agent.checkForUpdates(context.Background()) + }) + + t.Run("non-200 status", func(t *testing.T) { + swap(t, &Version, "1.0.0") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer server.Close() + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: server.URL, Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: server.Client(), + }, + } + agent.checkForUpdates(context.Background()) + }) + + t.Run("decode error", func(t *testing.T) { + swap(t, &Version, "1.0.0") + 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(), + targets: []TargetConfig{{URL: server.URL, Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: server.Client(), + }, + } + agent.checkForUpdates(context.Background()) + }) + + t.Run("server dev version", func(t *testing.T) { + swap(t, &Version, "1.0.0") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"version":"dev"}`)) + })) + defer server.Close() + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: server.URL, Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: server.Client(), + }, + } + agent.checkForUpdates(context.Background()) + }) + + t.Run("up to date", func(t *testing.T) { + swap(t, &Version, "v1.2.3") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"version":"1.2.3"}`)) + })) + defer server.Close() + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: server.URL, Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: server.Client(), + }, + } + agent.checkForUpdates(context.Background()) + }) + + t.Run("update success", func(t *testing.T) { + swap(t, &Version, "1.2.3") + called := false + swap(t, &selfUpdateFunc, func(*Agent, context.Context) error { + called = true + return nil + }) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"version":"1.2.4"}`)) + })) + defer server.Close() + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: server.URL, Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: server.Client(), + }, + } + agent.checkForUpdates(context.Background()) + if !called { + t.Fatal("expected selfUpdate to be called") + } + }) + + t.Run("update error", func(t *testing.T) { + swap(t, &Version, "1.2.3") + swap(t, &selfUpdateFunc, func(*Agent, context.Context) error { + return errors.New("update failed") + }) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"version":"1.2.4"}`)) + })) + defer server.Close() + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: server.URL, Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: server.Client(), + }, + } + agent.checkForUpdates(context.Background()) + }) +} + +type sizeReadCloser struct { + remaining int64 +} + +func (s *sizeReadCloser) Read(p []byte) (int, error) { + if s.remaining <= 0 { + return 0, io.EOF + } + if int64(len(p)) > s.remaining { + p = p[:s.remaining] + } + for i := range p { + p[i] = 0 + } + s.remaining -= int64(len(p)) + return len(p), nil +} + +func (s *sizeReadCloser) Close() error { + return nil +} + +func elfBytes() []byte { + return []byte{0x7f, 'E', 'L', 'F', 0x01, 0x02, 0x03} +} + +func sha256Hex(data []byte) string { + sum := sha256.Sum256(data) + return hex.EncodeToString(sum[:]) +} + +func TestSelfUpdate(t *testing.T) { + t.Run("no target", func(t *testing.T) { + agent := &Agent{logger: zerolog.Nop()} + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("executable error", func(t *testing.T) { + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + } + swap(t, &osExecutableFn, func() (string, error) { + return "", errors.New("no exec") + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("request creation error", func(t *testing.T) { + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com/\x7f", Token: "token"}}, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("request error", func(t *testing.T) { + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: {Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return nil, errors.New("send failed") + })}, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("status error", func(t *testing.T) { + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusInternalServerError, + Status: http.StatusText(http.StatusInternalServerError), + Body: io.NopCloser(strings.NewReader("fail")), + Header: make(http.Header), + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("create temp error", func(t *testing.T) { + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + body := elfBytes() + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + + swap(t, &osExecutableFn, func() (string, error) { + return filepath.Join(t.TempDir(), "missing", "exec"), nil + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("copy 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: http.Header{"X-Checksum-Sha256": []string{"ignored"}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("too large", func(t *testing.T) { + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: &sizeReadCloser{remaining: (100 * 1024 * 1024) + 1}, + Header: http.Header{"X-Checksum-Sha256": []string{"ignored"}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("close error", func(t *testing.T) { + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + body := elfBytes() + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + swap(t, &closeFileFn, func(*os.File) error { + return errors.New("close failed") + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("invalid elf", func(t *testing.T) { + body := []byte("bad") + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("missing checksum", func(t *testing.T) { + body := elfBytes() + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: make(http.Header), + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("checksum mismatch", func(t *testing.T) { + body := elfBytes() + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{"bad"}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("chmod error", func(t *testing.T) { + body := elfBytes() + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + swap(t, &osChmodFn, func(string, os.FileMode) error { + return errors.New("chmod failed") + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("rename backup error", func(t *testing.T) { + body := elfBytes() + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + swap(t, &osRenameFn, func(string, string) error { + return errors.New("rename failed") + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("rename replace error", func(t *testing.T) { + body := elfBytes() + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + + calls := 0 + swap(t, &osRenameFn, func(old, new string) error { + calls++ + if calls == 2 { + return errors.New("rename failed") + } + return os.Rename(old, new) + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("unraid read error", func(t *testing.T) { + body := elfBytes() + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + unraidPath := filepath.Join(dir, "unraid-version") + if err := os.WriteFile(unraidPath, []byte("1"), 0600); err != nil { + t.Fatalf("write unraid: %v", err) + } + swap(t, &unraidVersionPath, unraidPath) + persist := filepath.Join(dir, "persist") + if err := os.WriteFile(persist, []byte("old"), 0600); err != nil { + t.Fatalf("write persist: %v", err) + } + swap(t, &unraidPersistPath, persist) + swap(t, &osReadFileFn, func(string) ([]byte, error) { + return nil, errors.New("read failed") + }) + swap(t, &syscallExecFn, func(string, []string, []string) error { + return errors.New("exec failed") + }) + + _ = agent.selfUpdate(context.Background()) + }) + + t.Run("unraid write error", func(t *testing.T) { + body := elfBytes() + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + unraidPath := filepath.Join(dir, "unraid-version") + if err := os.WriteFile(unraidPath, []byte("1"), 0600); err != nil { + t.Fatalf("write unraid: %v", err) + } + swap(t, &unraidVersionPath, unraidPath) + persist := filepath.Join(dir, "persist") + if err := os.WriteFile(persist, []byte("old"), 0600); err != nil { + t.Fatalf("write persist: %v", err) + } + swap(t, &unraidPersistPath, persist) + swap(t, &osWriteFileFn, func(string, []byte, os.FileMode) error { + return errors.New("write failed") + }) + swap(t, &syscallExecFn, func(string, []string, []string) error { + return errors.New("exec failed") + }) + + _ = agent.selfUpdate(context.Background()) + }) + + t.Run("unraid rename error", func(t *testing.T) { + body := elfBytes() + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + unraidPath := filepath.Join(dir, "unraid-version") + if err := os.WriteFile(unraidPath, []byte("1"), 0600); err != nil { + t.Fatalf("write unraid: %v", err) + } + swap(t, &unraidVersionPath, unraidPath) + persist := filepath.Join(dir, "persist") + if err := os.WriteFile(persist, []byte("old"), 0600); err != nil { + t.Fatalf("write persist: %v", err) + } + swap(t, &unraidPersistPath, persist) + swap(t, &osRenameFn, func(old, new string) error { + if strings.HasSuffix(new, ".tmp") { + return os.Rename(old, new) + } + return errors.New("rename failed") + }) + swap(t, &syscallExecFn, func(string, []string, []string) error { + return errors.New("exec failed") + }) + + _ = agent.selfUpdate(context.Background()) + }) + + t.Run("exec error", func(t *testing.T) { + body := elfBytes() + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + swap(t, &syscallExecFn, func(string, []string, []string) error { + return errors.New("exec failed") + }) + + if err := agent.selfUpdate(context.Background()); err == nil { + t.Fatal("expected error") + } + }) + + t.Run("exec success", func(t *testing.T) { + body := elfBytes() + client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"X-Checksum-Sha256": []string{sha256Hex(body)}}, + }, nil + })} + + agent := &Agent{ + logger: zerolog.Nop(), + targets: []TargetConfig{{URL: "http://example.com", Token: "token"}}, + httpClients: map[bool]*http.Client{ + false: client, + }, + } + dir := t.TempDir() + execPath := filepath.Join(dir, "exec") + if err := os.WriteFile(execPath, elfBytes(), 0700); err != nil { + t.Fatalf("write exec: %v", err) + } + swap(t, &osExecutableFn, func() (string, error) { + return execPath, nil + }) + swap(t, &syscallExecFn, func(string, []string, []string) error { + return nil + }) + + if err := agent.selfUpdate(context.Background()); err != nil { + t.Fatalf("unexpected error: %v", err) + } + }) +} diff --git a/internal/dockeragent/swarm_coverage_test.go b/internal/dockeragent/swarm_coverage_test.go index 53e239a9f..c3749727d 100644 --- a/internal/dockeragent/swarm_coverage_test.go +++ b/internal/dockeragent/swarm_coverage_test.go @@ -24,6 +24,7 @@ func TestMapSwarmService(t *testing.T) { "com.docker.stack.namespace": "stack", }, }, + Mode: swarmtypes.ServiceMode{Replicated: &swarmtypes.ReplicatedService{}}, TaskTemplate: swarmtypes.TaskSpec{ ContainerSpec: &swarmtypes.ContainerSpec{ Image: "nginx:latest", @@ -128,7 +129,7 @@ func TestMapSwarmTask(t *testing.T) { ID: "task2", ServiceID: "svc2", Status: swarmtypes.TaskStatus{ - State: swarmtypes.TaskStateCompleted, + State: swarmtypes.TaskStateComplete, ContainerStatus: &swarmtypes.ContainerStatus{ ContainerID: "container-full", }, @@ -169,7 +170,7 @@ func TestCollectSwarmDataFromManager(t *testing.T) { } info := systemtypes.Info{ - Swarm: systemtypes.SwarmInfo{ + Swarm: swarmtypes.Info{ NodeID: "node1", }, } @@ -198,8 +199,8 @@ func TestCollectSwarmData(t *testing.T) { t.Run("inactive swarm returns info only", func(t *testing.T) { agent := &Agent{supportsSwarm: true, cfg: Config{SwarmScope: swarmScopeNode}} info := systemtypes.Info{ - Swarm: systemtypes.SwarmInfo{ - NodeID: "node1", + Swarm: swarmtypes.Info{ + NodeID: "node1", LocalNodeState: swarmtypes.LocalNodeStatePending, }, } @@ -250,7 +251,7 @@ func TestCollectSwarmData(t *testing.T) { } info := systemtypes.Info{ - Swarm: systemtypes.SwarmInfo{ + Swarm: swarmtypes.Info{ NodeID: "node1", ControlAvailable: true, LocalNodeState: swarmtypes.LocalNodeStateActive, @@ -291,7 +292,7 @@ func TestCollectSwarmData(t *testing.T) { } info := systemtypes.Info{ - Swarm: systemtypes.SwarmInfo{ + Swarm: swarmtypes.Info{ NodeID: "node1", ControlAvailable: true, LocalNodeState: swarmtypes.LocalNodeStateActive,