diff --git a/internal/app/components/customhandlers/get_tool.go b/internal/app/components/customhandlers/get_tool.go index dbb3585..d8184ce 100644 --- a/internal/app/components/customhandlers/get_tool.go +++ b/internal/app/components/customhandlers/get_tool.go @@ -32,28 +32,53 @@ func (g *GetTool) Handle(ctx context.Context, args []string, out io.Writer, _ co fileName := filepath.Base(source) destination := filepath.Join(g.cfg.ToolsPath, fileName) + // The tool is downloaded next to its destination and promoted only once it + // is complete and executable: a tool that is already there survives a + // failed fetch and is never seen half-written. go-getter does not replace + // a file it finds at its destination — it resumes the download from the + // file's size, so a tool that changed upstream came back as a splice of + // both versions, and one no larger than the old copy was not fetched at + // all — which is why the leftover of an interrupted run goes first. + partial := destination + ".part" + + err := removeIfExists(partial) + if err != nil { + return int(domain.ErrorResult), errors.WithMessage(err, "[components.GetTool] failed to remove a partial download") + } + progressTracker := components.NewDownloadProgressTracker(out, 5*time.Second) c := getter.Client{ Ctx: ctx, Src: args[0], - Dst: destination, + Dst: partial, Mode: getter.ClientModeFile, ProgressListener: progressTracker, } _, _ = out.Write([]byte("Getting tool from " + source + " to " + destination + " ...\n")) - err := c.Get() + err = c.Get() if err != nil { + _ = os.Remove(partial) + return int(domain.ErrorResult), errors.WithMessage(err, "[components.GetTool] failed to get tool") } - err = os.Chmod(destination, 0700) + err = os.Chmod(partial, 0700) if err != nil { + _ = os.Remove(partial) _, _ = out.Write([]byte("Failed to chmod tool")) + return int(domain.ErrorResult), errors.WithMessage(err, "[components.GetTool] failed to chmod tool") } + err = os.Rename(partial, destination) + if err != nil { + _ = os.Remove(partial) + + return int(domain.ErrorResult), errors.WithMessage(err, "[components.GetTool] failed to replace the previous tool") + } + err = config.UpdateEnvPath(g.cfg) if err != nil { _, _ = out.Write([]byte("Failed to update PATH with tools directories")) @@ -63,3 +88,12 @@ func (g *GetTool) Handle(ctx context.Context, args []string, out io.Writer, _ co return int(domain.SuccessResult), nil } + +func removeIfExists(path string) error { + err := os.Remove(path) + if err != nil && !errors.Is(err, os.ErrNotExist) { + return err + } + + return nil +} diff --git a/internal/app/components/customhandlers/get_tool_test.go b/internal/app/components/customhandlers/get_tool_test.go new file mode 100644 index 0000000..fa6066f --- /dev/null +++ b/internal/app/components/customhandlers/get_tool_test.go @@ -0,0 +1,109 @@ +package customhandlers_test + +import ( + "bytes" + "context" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "testing" + + "github.com/gameap/daemon/internal/app/components/customhandlers" + "github.com/gameap/daemon/internal/app/config" + "github.com/gameap/daemon/internal/app/contracts" + "github.com/gameap/daemon/internal/app/domain" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +const toolName = "install-tool.sh" + +// http.FileServer answers HEAD with Accept-Ranges and serves Range requests, +// which is what lets go-getter resume onto a file that is already there. +func Test_GetTool_ReplacesThePreviousTool(t *testing.T) { + published := bytes.Repeat([]byte("published line\n"), 100) + + srcDir := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(srcDir, toolName), published, 0o644)) + server := httptest.NewServer(http.FileServer(http.Dir(srcDir))) + defer server.Close() + + tests := []struct { + name string + existing []byte + partial []byte + }{ + { + name: "no previous tool", + }, + { + name: "smaller previous tool", + existing: bytes.Repeat([]byte("previous line\n"), 10), + }, + { + name: "larger previous tool", + existing: bytes.Repeat([]byte("previous line\n"), 200), + }, + { + name: "leftover of an interrupted download", + existing: bytes.Repeat([]byte("previous line\n"), 10), + partial: bytes.Repeat([]byte("interrupted line\n"), 10), + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + toolsDir := t.TempDir() + destination := filepath.Join(toolsDir, toolName) + if tt.existing != nil { + require.NoError(t, os.WriteFile(destination, tt.existing, 0o600)) + } + if tt.partial != nil { + require.NoError(t, os.WriteFile(destination+".part", tt.partial, 0o600)) + } + handler := customhandlers.NewGetTool(&config.Config{ToolsPath: toolsDir}) + out := &bytes.Buffer{} + + code, err := handler.Handle( + context.Background(), + []string{server.URL + "/" + toolName}, + out, + contracts.ExecutorOptions{}, + ) + + require.NoError(t, err, out.String()) + assert.Equal(t, int(domain.SuccessResult), code) + got, err := os.ReadFile(destination) + require.NoError(t, err) + assert.Equal(t, published, got) + assert.NoFileExists(t, destination+".part") + }) + } +} + +func Test_GetTool_KeepsThePreviousToolWhenTheDownloadFails(t *testing.T) { + previous := []byte("previous tool\n") + + server := httptest.NewServer(http.NotFoundHandler()) + defer server.Close() + + toolsDir := t.TempDir() + destination := filepath.Join(toolsDir, toolName) + require.NoError(t, os.WriteFile(destination, previous, 0o600)) + handler := customhandlers.NewGetTool(&config.Config{ToolsPath: toolsDir}) + out := &bytes.Buffer{} + + code, err := handler.Handle( + context.Background(), + []string{server.URL + "/" + toolName}, + out, + contracts.ExecutorOptions{}, + ) + + require.Error(t, err) + assert.Equal(t, int(domain.ErrorResult), code) + got, err := os.ReadFile(destination) + require.NoError(t, err) + assert.Equal(t, previous, got, "the previous tool must survive a failed fetch") + assert.NoFileExists(t, destination+".part") +}