diff --git a/internal/command/presentation.go b/internal/command/presentation.go index f5b84a1..b69ffb2 100644 --- a/internal/command/presentation.go +++ b/internal/command/presentation.go @@ -8,6 +8,7 @@ import ( "github.com/lettermint/lettermint-cli/internal/api" "github.com/lettermint/lettermint-cli/internal/presentation" + "github.com/lettermint/lettermint-cli/internal/update" "github.com/spf13/cobra" "github.com/spf13/pflag" ) @@ -128,12 +129,22 @@ func (a *app) configurePresentation(root *cobra.Command) { if original := c.RunE; original != nil { c.RunE = func(cmd *cobra.Command, args []string) error { ui := a.ui(cmd) + check := a.checkUpdate(cmd, key) + defer check.Stop() + if check != nil && check.Cached != "" { + _ = ui.UpdateAvailable(a.version, check.Cached, update.Instructions(a.version, check.Cached)) + } // These commands produce raw bytes or help and must stay undecorated. if key != "messages content" && !strings.HasPrefix(key, "completion") && key != "version" && cmd != root { ui.StartProgress(cmd.Context(), "Running "+key) } defer ui.StopProgress() - return original(cmd, args) + err := original(cmd, args) + ui.StopProgress() + if latest := check.Finish(err == nil && key != "webhooks listen"); latest != "" { + _ = ui.UpdateAvailable(a.version, latest, update.Instructions(a.version, latest)) + } + return err } } for _, child := range c.Commands() { @@ -143,6 +154,17 @@ func (a *app) configurePresentation(root *cobra.Command) { wrap(root) } +func (a *app) checkUpdate(c *cobra.Command, key string) *update.Check { + if a.noInput || !a.ui(c).UpdateNotifications() || c == c.Root() || + key == "version" || key == "messages content" || strings.HasPrefix(key, "completion") { + return nil + } + if a.updates == nil { + a.updates = update.New() + } + return a.updates.Start(c.Context(), a.version) +} + func commandExample(key string) string { switch key { case "auth login": diff --git a/internal/command/root.go b/internal/command/root.go index c7ab1ac..f34377a 100644 --- a/internal/command/root.go +++ b/internal/command/root.go @@ -11,6 +11,7 @@ import ( "github.com/lettermint/lettermint-cli/internal/config" "github.com/lettermint/lettermint-cli/internal/listener" "github.com/lettermint/lettermint-cli/internal/presentation" + "github.com/lettermint/lettermint-cli/internal/update" "github.com/lettermint/lettermint-cli/skills" "github.com/spf13/cobra" "io" @@ -27,6 +28,7 @@ type app struct { display presentation.Options presenter *presentation.Presenter scope presentation.Context + updates *update.Checker } func New(version, clientID string) *cobra.Command { diff --git a/internal/command/update_test.go b/internal/command/update_test.go new file mode 100644 index 0000000..6dc1592 --- /dev/null +++ b/internal/command/update_test.go @@ -0,0 +1,249 @@ +package command + +import ( + "bytes" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/lettermint/lettermint-cli/internal/presentation" + "github.com/lettermint/lettermint-cli/internal/update" + "github.com/spf13/cobra" +) + +func updateFixture(t *testing.T) (*update.Checker, *atomic.Int32) { + t.Helper() + requests := &atomic.Int32{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + requests.Add(1) + _, _ = io.WriteString(w, `{"tag_name":"v1.2.0","assets":[ + {"name":"lettermint_1.2.0_linux_amd64.tar.gz","state":"uploaded","size":100}, + {"name":"install.sh","state":"uploaded","size":100}, + {"name":"checksums.txt","state":"uploaded","size":100}, + {"name":"provenance.jsonl","state":"uploaded","size":100}]}`) + })) + t.Cleanup(server.Close) + return &update.Checker{ + CachePath: filepath.Join(t.TempDir(), "cache", "update.json"), + HTTP: server.Client(), Endpoint: server.URL, + Now: func() time.Time { return time.Date(2026, 9, 27, 12, 0, 0, 0, time.UTC) }, + OS: "linux", Arch: "amd64", + }, requests +} + +func updateRoot(a *app, run func(*cobra.Command, []string) error) *cobra.Command { + root := &cobra.Command{Use: "lettermint", SilenceErrors: true, SilenceUsage: true} + root.PersistentFlags().BoolVar(&a.noInput, "no-input", false, "Do not prompt") + for _, key := range []string{"skills list", "version", "completion", "messages content", "webhooks listen"} { + parts := strings.Fields(key) + parent := root + if len(parts) == 2 { + parent = &cobra.Command{Use: parts[0]} + root.AddCommand(parent) + } + parent.AddCommand(&cobra.Command{Use: parts[len(parts)-1], Args: cobra.NoArgs, RunE: run}) + } + a.configurePresentation(root) + return root +} + +func updateApp(c *update.Checker, out, diagnostic io.Writer, options presentation.Options, stdoutTTY, stderrTTY bool) *app { + return &app{ + version: "1.0.0", updates: c, display: options, + presenter: presentation.WithTerminals(out, diagnostic, options, + presentation.Terminal{TTY: stdoutTTY}, presentation.Terminal{TTY: stderrTTY}), + } +} + +func waitForUpdate(t *testing.T, c *update.Checker) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + data, _ := os.ReadFile(c.CachePath) + var cached struct{ Latest string } + if json.Unmarshal(data, &cached) == nil && cached.Latest == "1.2.0" { + return + } + time.Sleep(time.Millisecond) + } + t.Fatal("update was not cached") +} + +func TestFreshUpdateFollowsSuccessfulOutput(t *testing.T) { + c, requests := updateFixture(t) + var out, diagnostic bytes.Buffer + a := updateApp(c, &out, &diagnostic, presentation.Options{Plain: true}, true, true) + root := updateRoot(a, func(cmd *cobra.Command, _ []string) error { + waitForUpdate(t, c) + if diagnostic.Len() != 0 { + t.Fatal("background worker wrote output") + } + _, err := io.WriteString(cmd.OutOrStdout(), "result\n") + return err + }) + root.SetOut(&out) + root.SetErr(&diagnostic) + root.SetArgs([]string{"skills", "list"}) + if err := root.Execute(); err != nil { + t.Fatal(err) + } + want := "Update available: 1.0.0 -> 1.2.0\n" + + "https://github.com/lettermint/lettermint-cli/releases/tag/v1.2.0\n" + + strings.Join(update.Instructions(a.version, "1.2.0"), "\n") + "\n" + + "Installation guide: https://github.com/lettermint/lettermint-cli/blob/main/docs/installation.md\n" + if out.String() != "result\n" || diagnostic.String() != want || requests.Load() != 1 { + t.Fatalf("stdout=%q stderr=%q requests=%d", out.String(), diagnostic.String(), requests.Load()) + } +} + +func TestUpdateSuppression(t *testing.T) { + for _, tc := range []struct { + name string + args []string + options presentation.Options + stdoutTTY, stderrTTY bool + }{ + {"json", []string{"skills", "list"}, presentation.Options{JSON: true}, true, true}, + {"pipe", []string{"skills", "list"}, presentation.Options{}, false, true}, + {"plain-pipe", []string{"skills", "list"}, presentation.Options{Plain: true}, false, true}, + {"stderr-file", []string{"skills", "list"}, presentation.Options{}, true, false}, + {"no-input", []string{"skills", "list", "--no-input"}, presentation.Options{}, true, true}, + {"version", []string{"version"}, presentation.Options{}, true, true}, + {"content", []string{"messages", "content"}, presentation.Options{}, true, true}, + {"completion", []string{"completion"}, presentation.Options{}, true, true}, + {"dynamic-completion", []string{"__complete", "skills", ""}, presentation.Options{}, true, true}, + {"help", []string{"skills", "--help"}, presentation.Options{}, true, true}, + {"help-command", []string{"help", "skills"}, presentation.Options{}, true, true}, + {"root", nil, presentation.Options{}, true, true}, + {"invalid-arguments", []string{"skills", "list", "extra"}, presentation.Options{}, true, true}, + } { + t.Run(tc.name, func(t *testing.T) { + c, requests := updateFixture(t) + var out, diagnostic bytes.Buffer + a := updateApp(c, &out, &diagnostic, tc.options, tc.stdoutTTY, tc.stderrTTY) + ran := false + root := updateRoot(a, func(cmd *cobra.Command, _ []string) error { + ran = true + _, err := io.WriteString(cmd.OutOrStdout(), "{\"ok\":true}\n") + return err + }) + root.SetOut(&out) + root.SetErr(&diagnostic) + root.SetArgs(tc.args) + err := root.Execute() + if (err != nil) != (tc.name == "invalid-arguments") { + t.Fatal(err) + } + if ran && out.String() != "{\"ok\":true}\n" { + t.Fatalf("result changed: %q", out.String()) + } + if requests.Load() != 0 || strings.Contains(diagnostic.String(), "Update available") { + t.Fatal("suppressed command checked for an update") + } + if _, err := os.Stat(filepath.Dir(c.CachePath)); !os.IsNotExist(err) { + t.Fatal("suppressed command opened the cache") + } + }) + } +} + +func TestListenerAndFailedCommandDeferFreshAlert(t *testing.T) { + for _, listener := range []bool{false, true} { + t.Run(map[bool]string{false: "failure", true: "listener"}[listener], func(t *testing.T) { + c, requests := updateFixture(t) + var out, diagnostic bytes.Buffer + a := updateApp(c, &out, &diagnostic, presentation.Options{Plain: true}, true, true) + failure := errors.New("command failed") + root := updateRoot(a, func(cmd *cobra.Command, _ []string) error { + waitForUpdate(t, c) + _, _ = io.WriteString(cmd.OutOrStdout(), "event\n") + if diagnostic.Len() != 0 { + t.Fatal("background worker mixed an alert into command output") + } + if !listener { + return failure + } + return nil + }) + root.SetOut(&out) + root.SetErr(&diagnostic) + args := []string{"skills", "list"} + if listener { + args = []string{"webhooks", "listen"} + } + root.SetArgs(args) + err := root.Execute() + if listener && err != nil || !listener && err != failure { + t.Fatalf("command error changed: %v", err) + } + if diagnostic.Len() != 0 || out.String() != "event\n" { + t.Fatalf("stdout=%q stderr=%q", out.String(), diagnostic.String()) + } + // The next command must display the cached result before its own work. + next := updateRoot(a, func(*cobra.Command, []string) error { + if !strings.Contains(diagnostic.String(), "Update available: 1.0.0 -> 1.2.0") { + t.Fatal("cached alert was not displayed before command work") + } + return nil + }) + next.SetArgs([]string{"skills", "list"}) + if err := next.Execute(); err != nil { + t.Fatal(err) + } + if requests.Load() != 1 { + t.Fatal("cached alert made a new request") + } + }) + } +} + +func TestRequiredFlagsValidatedBeforeUpdateCheck(t *testing.T) { + c, requests := updateFixture(t) + var out, diagnostic bytes.Buffer + a := updateApp(c, &out, &diagnostic, presentation.Options{}, true, true) + root := updateRoot(a, func(*cobra.Command, []string) error { + t.Fatal("command ran without required flag") + return nil + }) + cmd, _, err := root.Find([]string{"skills", "list"}) + if err != nil { + t.Fatal(err) + } + cmd.Flags().String("required", "", "Required input") + if err := cmd.MarkFlagRequired("required"); err != nil { + t.Fatal(err) + } + root.SetArgs([]string{"skills", "list"}) + if root.Execute() == nil || requests.Load() != 0 { + t.Fatal("missing required flag did not prevent the check") + } + if _, err := os.Stat(filepath.Dir(c.CachePath)); !os.IsNotExist(err) { + t.Fatal("opened cache before flag validation") + } +} + +type failedNoticeWriter struct{} + +func (failedNoticeWriter) Write([]byte) (int, error) { return 0, io.ErrClosedPipe } + +func TestNoticeWriteFailureDoesNotFailCommand(t *testing.T) { + c, _ := updateFixture(t) + var out bytes.Buffer + a := updateApp(c, &out, failedNoticeWriter{}, presentation.Options{Plain: true}, true, true) + root := updateRoot(a, func(*cobra.Command, []string) error { + waitForUpdate(t, c) + return nil + }) + root.SetArgs([]string{"skills", "list"}) + if err := root.Execute(); err != nil { + t.Fatal("notice changed command result:", err) + } +} diff --git a/internal/presentation/presentation.go b/internal/presentation/presentation.go index e150773..a74b84f 100644 --- a/internal/presentation/presentation.go +++ b/internal/presentation/presentation.go @@ -165,6 +165,24 @@ func (p *Presenter) Notice(text string) error { return p.write(p.err, p.diagnostic, Text(text)+"\n") } +func (p *Presenter) UpdateNotifications() bool { + return !p.JSON() && p.output.TTY && p.diagnostic.TTY +} + +func (p *Presenter) UpdateAvailable(current, latest string, instructions []string) error { + if !p.UpdateNotifications() { + return nil + } + p.StopProgress() + text := fmt.Sprintf("%s %s -> %s\n", p.heading("Update available:"), Text(strings.TrimPrefix(current, "v")), Text(latest)) + + "https://github.com/lettermint/lettermint-cli/releases/tag/v" + Text(latest) + "\n" + for _, line := range instructions { + text += Text(line) + "\n" + } + text += "Installation guide: https://github.com/lettermint/lettermint-cli/blob/main/docs/installation.md\n" + return p.write(p.err, p.diagnostic, text) +} + func (p *Presenter) Prompt(text string) error { p.StopProgress() return p.write(p.err, p.diagnostic, Text(text)+" [y/N]: ") diff --git a/internal/presentation/update_test.go b/internal/presentation/update_test.go new file mode 100644 index 0000000..0113971 --- /dev/null +++ b/internal/presentation/update_test.go @@ -0,0 +1,43 @@ +package presentation + +import ( + "bytes" + "strings" + "testing" + + "github.com/charmbracelet/colorprofile" +) + +func TestUpdateNoticeModes(t *testing.T) { + for _, tc := range []struct { + name string + options Options + profile colorprofile.Profile + color bool + }{ + {"terminal", Options{}, colorprofile.TrueColor, true}, + {"plain", Options{Plain: true, Color: "always"}, colorprofile.TrueColor, false}, + {"never", Options{Color: "never"}, colorprofile.TrueColor, false}, + {"no-color-terminal", Options{}, colorprofile.NoTTY, false}, + {"forced-color", Options{Color: "always"}, colorprofile.NoTTY, true}, + {"json", Options{JSON: true}, colorprofile.TrueColor, false}, + } { + t.Run(tc.name, func(t *testing.T) { + var out, diagnostic bytes.Buffer + terminal := Terminal{TTY: true, Width: 80, Profile: tc.profile} + p := WithTerminals(&out, &diagnostic, tc.options, terminal, terminal) + if err := p.UpdateAvailable("v1.0.0", "1.2.0", []string{"Update with Homebrew:", " brew update && brew upgrade --cask lettermint"}); err != nil { + t.Fatal(err) + } + if out.Len() != 0 || strings.Contains(diagnostic.String(), "\x1b") != tc.color { + t.Fatalf("stdout=%q stderr=%q", out.String(), diagnostic.String()) + } + if tc.options.JSON != (diagnostic.Len() == 0) { + t.Fatal("incorrect notice visibility") + } + if !tc.options.JSON && !strings.Contains(diagnostic.String(), " brew update && brew upgrade --cask lettermint\n") { + t.Fatal("missing update command") + } + }) + } +} diff --git a/internal/update/instructions.go b/internal/update/instructions.go new file mode 100644 index 0000000..01a60b8 --- /dev/null +++ b/internal/update/instructions.go @@ -0,0 +1,135 @@ +package update + +import ( + "bytes" + "crypto/sha256" + "encoding/json" + "fmt" + "io" + "os" + "path/filepath" + "runtime" + "strings" + "unicode" +) + +type installation struct { + method string + binDir string +} + +// Instructions describes an update. It never runs an installer or package manager. +func Instructions(current, latest string) []string { + executable, _ := os.Executable() + installed := detectInstallation(executable, runtime.GOOS, os.Getenv("LOCALAPPDATA"), current) + return installed.instructions(latest, runtime.GOOS, runtime.GOARCH) +} + +func detectInstallation(executable, system, localAppData, current string) installation { + resolved, err := filepath.EvalSymlinks(executable) + if err != nil || stable(current) == "" { + return installation{} + } + binDir := filepath.Dir(resolved) + if system == "darwin" && filepath.Base(resolved) == "lettermint" { + versionDir := filepath.Dir(resolved) + packageDir := filepath.Dir(versionDir) + if filepath.Base(filepath.Dir(packageDir)) == "Caskroom" && + filepath.Base(packageDir) == "lettermint" && filepath.Base(versionDir) == stable(current) { + return installation{method: "homebrew"} + } + } + if system == "windows" && localAppData != "" { + root := filepath.Join(localAppData, "Lettermint CLI") + expected, err := filepath.EvalSymlinks(filepath.Join(root, "bin", "lettermint.exe")) + if err != nil || !strings.EqualFold(resolved, expected) { + return installation{} + } + var record struct{ Manager, Version string } + // Windows PowerShell 5.1 writes UTF-8 with a byte-order mark. + data := bytes.TrimPrefix(ownershipRecord(filepath.Join(root, "install.json")), []byte{0xef, 0xbb, 0xbf}) + if json.Unmarshal(data, &record) == nil && record.Manager == "lettermint-powershell" && record.Version == "v"+stable(current) { + return installation{method: "powershell"} + } + } + if (system == "darwin" || system == "linux") && filepath.Base(resolved) == "lettermint" { + record := strings.Split(strings.TrimSuffix(string(ownershipRecord(filepath.Join(binDir, ".lettermint-install"))), "\n"), "\n") + if len(record) == 3 && record[0] == "lettermint-shell-v1" && record[1] == "v"+stable(current) && + strings.IndexFunc(binDir, func(r rune) bool { return unicode.IsControl(r) || unicode.In(r, unicode.Cf) }) < 0 && + matchesHash(resolved, record[2]) { + return installation{method: "shell", binDir: binDir} + } + } + return installation{} +} + +func ownershipRecord(path string) []byte { + info, err := os.Lstat(path) + if err != nil || !info.Mode().IsRegular() || info.Size() > 4096 { + return nil + } + f, err := os.Open(path) + if err != nil { + return nil + } + defer f.Close() + data, err := io.ReadAll(io.LimitReader(f, 4097)) + if err != nil || len(data) > 4096 { + return nil + } + return data +} + +func matchesHash(path, expected string) bool { + if len(expected) != sha256.Size*2 { + return false + } + f, err := os.Open(path) + if err != nil { + return false + } + defer f.Close() + hash := sha256.New() + if _, err := io.Copy(hash, f); err != nil { + return false + } + return fmt.Sprintf("%x", hash.Sum(nil)) == expected +} + +func (i installation) instructions(latest, system, arch string) []string { + version := stable(latest) + if version == "" { + return nil + } + base := "https://github.com/lettermint/lettermint-cli/releases/download/v" + version + switch i.method { + case "homebrew": + return []string{"Update with Homebrew:", " brew update && brew upgrade --cask lettermint"} + case "shell": + return []string{ + "Update with the shell installer:", + " curl -fsSL " + base + "/install.sh | sh -s -- --version v" + version + " --bin-dir " + shellQuote(i.binDir), + } + case "powershell": + return []string{ + "Update in PowerShell:", + ` Invoke-WebRequest -UseBasicParsing -Uri '` + base + `/install.ps1' -OutFile "$env:TEMP\lettermint.ps1"`, + ` powershell -NoProfile -ExecutionPolicy AllSigned -File "$env:TEMP\lettermint.ps1" -Version v` + version, + } + default: + lines := []string{"Installation method not detected. Use your package manager or the manual installation guide."} + label := map[string]string{"darwin": "macOS", "linux": "Linux", "windows": "Windows"}[system] + if label != "" && (arch == "amd64" || arch == "arm64") { + extension := ".tar.gz" + if system == "windows" { + extension = ".zip" + } + lines = append(lines, fmt.Sprintf("Download for %s (%s): %s/lettermint_%s_%s_%s%s", label, arch, base, version, system, arch, extension)) + } + return lines + } +} + +func shellQuote(value string) string { + return "'" + strings.ReplaceAll(value, "'", "'\"'\"'") + "'" +} diff --git a/internal/update/instructions_test.go b/internal/update/instructions_test.go new file mode 100644 index 0000000..a5403da --- /dev/null +++ b/internal/update/instructions_test.go @@ -0,0 +1,163 @@ +package update + +import ( + "crypto/sha256" + "fmt" + "os" + "os/exec" + "path/filepath" + "runtime" + "strings" + "testing" +) + +func installFile(t *testing.T, path, contents string) { + t.Helper() + if err := os.MkdirAll(filepath.Dir(path), 0700); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(path, []byte(contents), 0600); err != nil { + t.Fatal(err) + } +} + +func TestDetectHomebrewFromResolvedExecutable(t *testing.T) { + for _, prefix := range []string{"opt/homebrew", "usr/local", "custom/brew"} { + t.Run(prefix, func(t *testing.T) { + executable := filepath.Join(t.TempDir(), prefix, "Caskroom", "lettermint", "1.0.0", "lettermint") + installFile(t, executable, "binary") + got := detectInstallation(executable, "darwin", "", "1.0.0") + if got.method != "homebrew" { + t.Fatalf("installation = %+v", got) + } + if got := strings.Join(got.instructions("1.2.0", "darwin", "arm64"), "\n"); got != "Update with Homebrew:\n brew update && brew upgrade --cask lettermint" { + t.Fatal(got) + } + link := filepath.Join(t.TempDir(), "lettermint") + if err := os.Symlink(executable, link); err != nil { + t.Skip("symlinks unavailable:", err) + } + if detectInstallation(link, "darwin", "", "1.0.0").method != "homebrew" { + t.Fatal("did not resolve the executable symlink") + } + }) + } +} + +func TestDetectShellOwnershipAndCustomDirectory(t *testing.T) { + for _, system := range []string{"darwin", "linux"} { + t.Run(system, func(t *testing.T) { + bin := filepath.Join(t.TempDir(), "Bob's custom tools") + executable := filepath.Join(bin, "lettermint") + installFile(t, executable, "installed binary") + hash := sha256.Sum256([]byte("installed binary")) + marker := filepath.Join(bin, ".lettermint-install") + installFile(t, marker, fmt.Sprintf("lettermint-shell-v1\nv1.0.0\n%x\n", hash)) + got := detectInstallation(executable, system, "", "1.0.0") + resolved, err := filepath.EvalSymlinks(bin) + if err != nil { + t.Fatal(err) + } + if got.method != "shell" || got.binDir != resolved { + t.Fatalf("installation = %+v", got) + } + lines := got.instructions("1.2.0", system, "amd64") + if len(lines) != 2 || !strings.Contains(lines[1], "/releases/download/v1.2.0/install.sh | sh -s -- --version v1.2.0 --bin-dir ") || + !strings.HasSuffix(lines[1], shellQuote(resolved)) { + t.Fatal(lines) + } + // A copied or changed executable is no longer owned by this record. + installFile(t, executable, "changed binary") + if detectInstallation(executable, system, "", "1.0.0").method != "" { + t.Fatal("accepted a stale executable hash") + } + }) + } +} + +func TestDetectPowerShellOwnership(t *testing.T) { + for _, bom := range []string{"", "\xef\xbb\xbf"} { + t.Run(fmt.Sprintf("bom=%t", bom != ""), func(t *testing.T) { + localAppData := t.TempDir() + root := filepath.Join(localAppData, "Lettermint CLI") + executable := filepath.Join(root, "bin", "lettermint.exe") + installFile(t, executable, "binary") + installFile(t, filepath.Join(root, "install.json"), bom+`{"manager":"lettermint-powershell","version":"v1.0.0"}`) + got := detectInstallation(executable, "windows", localAppData, "1.0.0") + if got.method != "powershell" { + t.Fatalf("installation = %+v", got) + } + lines := got.instructions("1.2.0", "windows", "arm64") + if len(lines) != 3 || !strings.Contains(lines[1], "/releases/download/v1.2.0/install.ps1' -OutFile ") || + lines[2] != ` powershell -NoProfile -ExecutionPolicy AllSigned -File "$env:TEMP\lettermint.ps1" -Version v1.2.0` { + t.Fatal(lines) + } + // The record must belong to the running executable, not another install. + other := filepath.Join(t.TempDir(), "lettermint.exe") + installFile(t, other, "binary") + if detectInstallation(other, "windows", localAppData, "1.0.0").method != "" { + t.Fatal("used the ownership record from another install") + } + }) + } +} + +func TestUnknownOrInvalidOwnership(t *testing.T) { + for _, tc := range []struct { + name, system, executable, marker, data string + }{ + {"manual-macos", "darwin", "bin/lettermint", "", ""}, + {"manual-linux", "linux", "bin/lettermint", "", ""}, + {"homebrew-lookalike", "darwin", "Caskroom/other/1.0.0/lettermint", "", ""}, + {"homebrew-wrong-version", "darwin", "Caskroom/lettermint/2.0.0/lettermint", "", ""}, + {"bad-shell-owner", "linux", "bin/lettermint", "bin/.lettermint-install", "another-installer\nv1.0.0\nhash\n"}, + {"bad-shell-hash", "linux", "bin/lettermint", "bin/.lettermint-install", "lettermint-shell-v1\nv1.0.0\nwrong\n"}, + {"manual-windows", "windows", "Lettermint CLI/bin/lettermint.exe", "", ""}, + {"bad-powershell-owner", "windows", "Lettermint CLI/bin/lettermint.exe", "Lettermint CLI/install.json", `{"manager":"other","version":"v1.0.0"}`}, + {"bad-powershell-version", "windows", "Lettermint CLI/bin/lettermint.exe", "Lettermint CLI/install.json", `{"manager":"lettermint-powershell","version":"v2.0.0"}`}, + {"corrupt-powershell-record", "windows", "Lettermint CLI/bin/lettermint.exe", "Lettermint CLI/install.json", "{"}, + } { + t.Run(tc.name, func(t *testing.T) { + root := t.TempDir() + executable := filepath.Join(root, filepath.FromSlash(tc.executable)) + installFile(t, executable, "binary") + if tc.marker != "" { + installFile(t, filepath.Join(root, filepath.FromSlash(tc.marker)), tc.data) + } + if got := detectInstallation(executable, tc.system, root, "1.0.0"); got.method != "" { + t.Fatalf("guessed installation method: %+v", got) + } + }) + } +} + +func TestManualInstructionsMatchPlatform(t *testing.T) { + for _, system := range []string{"darwin", "linux", "windows"} { + for _, arch := range []string{"amd64", "arm64"} { + lines := (installation{}).instructions("1.2.0", system, arch) + extension := ".tar.gz" + if system == "windows" { + extension = ".zip" + } + want := fmt.Sprintf("/releases/download/v1.2.0/lettermint_1.2.0_%s_%s%s", system, arch, extension) + if len(lines) != 2 || !strings.HasSuffix(lines[1], want) || !strings.Contains(lines[0], "manual installation guide") { + t.Fatal(lines) + } + } + } + if lines := (installation{}).instructions("bad version", "linux", "amd64"); len(lines) != 0 { + t.Fatal("rendered an invalid release version") + } +} + +func TestShellQuoteKeepsLiteralPaths(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("POSIX shell is not required on Windows") + } + for _, path := range []string{"/home/user/.local/bin", "/home/Bob's tools", "/tmp/$HOME/$(echo changed)/`echo changed`"} { + out, err := exec.Command("sh", "-c", "printf '%s' "+shellQuote(path)).Output() + if err != nil || string(out) != path { + t.Fatalf("path=%q output=%q error=%v", path, out, err) + } + } +} diff --git a/internal/update/update.go b/internal/update/update.go new file mode 100644 index 0000000..4e0c1a1 --- /dev/null +++ b/internal/update/update.go @@ -0,0 +1,269 @@ +// Package update checks for complete stable releases without changing the CLI. +package update + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "os" + "path/filepath" + "regexp" + "runtime" + "strings" + "time" + + "github.com/gofrs/flock" +) + +const interval = 24 * time.Hour + +// Checker keeps release requests separate from the authenticated API client. +// Its fields allow tests to use a local server, temporary cache, and fixed clock. +type Checker struct { + CachePath string + HTTP *http.Client + Endpoint string + Now func() time.Time + OS, Arch string +} + +func New() *Checker { + dir, err := os.UserCacheDir() + if err != nil { + return nil + } + return &Checker{ + CachePath: filepath.Join(dir, "lettermint", "update.json"), + HTTP: &http.Client{Timeout: 3 * time.Second}, + Endpoint: "https://api.github.com/repos/lettermint/lettermint-cli/releases/latest", + Now: time.Now, + OS: runtime.GOOS, + Arch: runtime.GOARCH, + } +} + +type state struct { + CheckedAt time.Time `json:"checked_at"` + Latest string `json:"latest,omitempty"` + AlertedAt time.Time `json:"alerted_at"` +} + +// Check owns one worker. Only the caller may display Cached or Finish's result. +type Check struct { + Cached string + checker *Checker + current string + cancel context.CancelFunc + done chan struct{} + ctx context.Context +} + +// Start claims any cached alert and starts a request only when a check is due. +// A held lock or an unavailable cache disables this invocation's check. +func (c *Checker) Start(ctx context.Context, current string) *Check { + if c == nil || stable(current) == "" || ctx.Err() != nil { + return nil + } + lock := c.lock() + if lock == nil { + return nil + } + s := c.read() + checkCtx, cancel := context.WithCancel(ctx) + check := &Check{checker: c, current: current, cancel: cancel, done: make(chan struct{}), ctx: ctx} + check.Cached = c.claim(&s, current) + if recent(s.CheckedAt, c.Now()) { + _ = lock.Unlock() + close(check.done) + return check + } + go func() { + defer close(check.done) + defer lock.Unlock() + requestCtx, stop := context.WithTimeout(checkCtx, 3*time.Second) + defer stop() + latest, err := c.fetch(requestCtx) + // Exit cancellation must not consume the next invocation's check. + // A request timeout is a completed failure and does consume it. + if checkCtx.Err() != nil { + return + } + s.CheckedAt = c.Now() + if err == nil { + s.Latest = latest + } + _ = c.write(s) + }() + return check +} + +// Stop cancels network work and waits for cleanup, not for the request timeout. +func (c *Check) Stop() { + if c != nil { + c.cancel() + <-c.done + } +} + +// Finish returns an unclaimed alert only after successful command completion. +func (c *Check) Finish(success bool) string { + if c == nil { + return "" + } + c.Stop() + if !success || c.Cached != "" || c.ctx.Err() != nil { + return "" + } + lock := c.checker.lock() + if lock == nil { + return "" + } + defer lock.Unlock() + s := c.checker.read() + return c.checker.claim(&s, c.current) +} + +func (c *Checker) claim(s *state, current string) string { + if !newer(s.Latest, current) || recent(s.AlertedAt, c.Now()) { + return "" + } + s.AlertedAt = c.Now() + if c.write(*s) != nil { + return "" + } + return stable(s.Latest) +} + +func recent(t, now time.Time) bool { + return !t.IsZero() && !t.After(now) && now.Sub(t) < interval +} + +func (c *Checker) lock() *flock.Flock { + if c.CachePath == "" || os.MkdirAll(filepath.Dir(c.CachePath), 0700) != nil { + return nil + } + lock := flock.New(c.CachePath + ".lock") + ok, err := lock.TryLock() + if err != nil || !ok { + _ = lock.Close() + return nil + } + return lock +} + +func (c *Checker) read() state { + f, err := os.Open(c.CachePath) + if err != nil { + return state{} + } + defer f.Close() + var s state + if json.NewDecoder(io.LimitReader(f, 64<<10)).Decode(&s) != nil { + return state{} + } + return s +} + +func (c *Checker) write(s state) error { + f, err := os.CreateTemp(filepath.Dir(c.CachePath), "update-*") + if err != nil { + return err + } + defer os.Remove(f.Name()) + err = json.NewEncoder(f).Encode(s) + if err == nil { + err = f.Sync() + } + closeErr := f.Close() + if err != nil { + return err + } + if closeErr != nil { + return closeErr + } + return os.Rename(f.Name(), c.CachePath) +} + +func (c *Checker) fetch(ctx context.Context) (string, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.Endpoint, nil) + if err != nil { + return "", err + } + req.Header.Set("Accept", "application/vnd.github+json") + req.Header.Set("User-Agent", "lettermint-cli-update-check") + resp, err := c.HTTP.Do(req) + if err != nil { + return "", err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("release request returned HTTP %d", resp.StatusCode) + } + var release struct { + Tag string `json:"tag_name"` + Draft bool `json:"draft"` + Prerelease bool `json:"prerelease"` + Assets []struct { + Name string `json:"name"` + State string `json:"state"` + Size int64 `json:"size"` + } `json:"assets"` + } + if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&release); err != nil { + return "", err + } + version := stable(release.Tag) + if release.Draft || release.Prerelease || version == "" || release.Tag != "v"+version { + return "", nil + } + installer, extension := "install.sh", ".tar.gz" + if c.OS == "windows" { + installer, extension = "install.ps1", ".zip" + } + required := map[string]bool{ + fmt.Sprintf("lettermint_%s_%s_%s%s", version, c.OS, c.Arch, extension): false, + installer: false, "checksums.txt": false, "provenance.jsonl": false, + } + for _, asset := range release.Assets { + if _, ok := required[asset.Name]; ok && asset.State == "uploaded" && asset.Size > 0 { + required[asset.Name] = true + } + } + for _, found := range required { + if !found { + return "", nil + } + } + return version, nil +} + +// Only stable semantic versions are eligible. Build metadata has no precedence. +var stableVersion = regexp.MustCompile(`^v?((?:0|[1-9][0-9]*)\.(?:0|[1-9][0-9]*)\.(?:0|[1-9][0-9]*))(?:\+[0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*)?$`) + +func stable(version string) string { + match := stableVersion.FindStringSubmatch(version) + if match == nil { + return "" + } + return match[1] +} + +func newer(candidate, current string) bool { + a, b := stable(candidate), stable(current) + if a == "" || b == "" { + return false + } + x, y := strings.Split(a, "."), strings.Split(b, ".") + for i := range x { + // Compare decimal components without integer overflow. + if len(x[i]) != len(y[i]) { + return len(x[i]) > len(y[i]) + } + if x[i] != y[i] { + return x[i] > y[i] + } + } + return false +} diff --git a/internal/update/update_test.go b/internal/update/update_test.go new file mode 100644 index 0000000..9640414 --- /dev/null +++ b/internal/update/update_test.go @@ -0,0 +1,328 @@ +package update + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync/atomic" + "testing" + "time" +) + +func fixture(t *testing.T, handler http.HandlerFunc) *Checker { + t.Helper() + server := httptest.NewServer(handler) + t.Cleanup(server.Close) + return &Checker{ + CachePath: filepath.Join(t.TempDir(), "update.json"), + HTTP: server.Client(), Endpoint: server.URL, + Now: func() time.Time { return time.Date(2026, 9, 27, 12, 0, 0, 0, time.UTC) }, + OS: "linux", Arch: "amd64", + } +} + +func release(version, system, arch string) map[string]any { + installer, suffix := "install.sh", ".tar.gz" + if system == "windows" { + installer, suffix = "install.ps1", ".zip" + } + assets := []map[string]any{} + for _, name := range []string{ + fmt.Sprintf("lettermint_%s_%s_%s%s", version, system, arch, suffix), + installer, "checksums.txt", "provenance.jsonl", + } { + assets = append(assets, map[string]any{"name": name, "state": "uploaded", "size": 100}) + } + return map[string]any{"tag_name": "v" + version, "draft": false, "prerelease": false, "assets": assets} +} + +func complete(t *testing.T, check *Check, success bool) string { + t.Helper() + if check == nil { + t.Fatal("check did not start") + } + t.Cleanup(check.Stop) + select { + case <-check.done: + case <-time.After(5 * time.Second): + t.Fatal("check did not finish") + } + return check.Finish(success) +} + +func TestStableVersionComparison(t *testing.T) { + for _, tc := range []struct { + candidate, current string + want bool + }{ + {"1.10.0", "v1.9.9", true}, {"v2.0.0", "1.99.99", true}, + {"1.2.4", "1.2.3+build.1", true}, {"1.2.3+other", "1.2.3+build", false}, + {"1.2.3", "1.2.3", false}, {"1.2.3", "1.3.0", false}, + {"1.2.3-rc.1", "1.2.2", false}, {"1.2.3", "1.2.3-rc.1", false}, + {"1.2.3", "dev", false}, {"1.2.3", "1.2.2-SNAPSHOT", false}, + {"01.2.3", "1.2.2", false}, {"1.2", "1.0.0", false}, + {"1.2.3\n", "1.0.0", false}, {"1.2.3+", "1.0.0", false}, + {"999999999999999999999999.0.0", "9.0.0", true}, + } { + t.Run(tc.candidate+"/"+tc.current, func(t *testing.T) { + if got := newer(tc.candidate, tc.current); got != tc.want { + t.Fatalf("newer = %v, want %v", got, tc.want) + } + }) + } +} + +func TestCompleteReleaseForEachPlatform(t *testing.T) { + for _, system := range []string{"linux", "darwin", "windows"} { + for _, arch := range []string{"amd64", "arm64"} { + t.Run(system+"/"+arch, func(t *testing.T) { + c := fixture(t, func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "" || r.Header.Get("Cookie") != "" { + t.Error("release request contains credentials") + } + if r.Method != http.MethodGet || r.Header.Get("User-Agent") != "lettermint-cli-update-check" { + t.Error("unexpected release request") + } + _ = json.NewEncoder(w).Encode(release("1.10.0", system, arch)) + }) + c.OS, c.Arch = system, arch + if got := complete(t, c.Start(context.Background(), "1.9.0"), true); got != "1.10.0" { + t.Fatalf("alert = %q", got) + } + }) + } + } +} + +func TestIncompleteAndNonStableReleases(t *testing.T) { + for _, name := range []string{"archive", "installer", "checksums", "provenance", "uploading", "empty", "draft", "prerelease", "rc-tag", "invalid-tag", "wrong-platform"} { + t.Run(name, func(t *testing.T) { + data := release("1.2.0", "linux", "amd64") + assets := data["assets"].([]map[string]any) + switch name { + case "archive", "installer", "checksums", "provenance": + index := map[string]int{"archive": 0, "installer": 1, "checksums": 2, "provenance": 3}[name] + data["assets"] = append(assets[:index:index], assets[index+1:]...) + case "uploading": + assets[0]["state"] = "new" + case "empty": + assets[0]["size"] = 0 + case "draft", "prerelease": + data[name] = true + case "rc-tag": + data["tag_name"] = "v1.2.0-rc.1" + case "invalid-tag": + data["tag_name"] = "v1.2.0\n" + case "wrong-platform": + data = release("1.2.0", "windows", "arm64") + } + c := fixture(t, func(w http.ResponseWriter, _ *http.Request) { _ = json.NewEncoder(w).Encode(data) }) + if got := complete(t, c.Start(context.Background(), "1.0.0"), true); got != "" { + t.Fatalf("unexpected alert: %s", got) + } + }) + } +} + +func TestDailyChecksAndAlerts(t *testing.T) { + var requests atomic.Int32 + c := fixture(t, func(w http.ResponseWriter, _ *http.Request) { + requests.Add(1) + _ = json.NewEncoder(w).Encode(release("1.2.0", "linux", "amd64")) + }) + now := c.Now() + c.Now = func() time.Time { return now } + if got := complete(t, c.Start(context.Background(), "1.0.0"), true); got != "1.2.0" { + t.Fatal(got) + } + for _, elapsed := range []time.Duration{0, 23 * time.Hour, time.Hour} { + now = now.Add(elapsed) + check := c.Start(context.Background(), "1.0.0") + want := "" + if elapsed == time.Hour { + want = "1.2.0" + } + if check.Cached != want || complete(t, check, true) != "" { + t.Fatalf("cached alert = %q, want %q", check.Cached, want) + } + } + if requests.Load() != 2 { + t.Fatalf("requests = %d", requests.Load()) + } + // An installed update suppresses the cached notice immediately. + now = now.Add(24 * time.Hour) + check := c.Start(context.Background(), "1.2.0") + if check.Cached != "" || complete(t, check, true) != "" { + t.Fatal("alert after upgrade") + } +} + +func TestFailureAndListenerSaveNoticeForNextRun(t *testing.T) { + c := fixture(t, func(w http.ResponseWriter, _ *http.Request) { + _ = json.NewEncoder(w).Encode(release("1.2.0", "linux", "amd64")) + }) + if got := complete(t, c.Start(context.Background(), "1.0.0"), false); got != "" { + t.Fatal("failed command displayed a fresh alert") + } + check := c.Start(context.Background(), "1.0.0") + if check.Cached != "1.2.0" || complete(t, check, true) != "" { + t.Fatal("next run did not claim the saved alert") + } +} + +func TestCompletedFailuresConsumeDailyCheck(t *testing.T) { + for _, kind := range []string{"http", "json", "timeout"} { + t.Run(kind, func(t *testing.T) { + var requests atomic.Int32 + c := fixture(t, func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + switch kind { + case "http": + w.WriteHeader(http.StatusTooManyRequests) + case "json": + _, _ = io.WriteString(w, "not json") + case "timeout": + <-r.Context().Done() + } + }) + if kind == "timeout" { + c.HTTP.Timeout = 30 * time.Millisecond + } + for range 2 { + if got := complete(t, c.Start(context.Background(), "1.0.0"), true); got != "" { + t.Fatal(got) + } + } + if requests.Load() != 1 || c.read().CheckedAt.IsZero() { + t.Fatal("failure was not cached") + } + }) + } +} + +func TestExitCancellationIsPromptAndDoesNotConsumeCheck(t *testing.T) { + started := make(chan struct{}, 2) + c := fixture(t, func(w http.ResponseWriter, r *http.Request) { + started <- struct{}{} + <-r.Context().Done() + }) + for range 2 { + check := c.Start(context.Background(), "1.0.0") + <-started + start := time.Now() + if check.Finish(true) != "" || time.Since(start) > time.Second { + t.Fatal("exit waited for the request timeout") + } + if !c.read().CheckedAt.IsZero() { + t.Fatal("cancellation consumed the next check") + } + } +} + +func TestConcurrentRunsSkipHeldLock(t *testing.T) { + started, releaseResponse := make(chan struct{}), make(chan struct{}) + c := fixture(t, func(w http.ResponseWriter, r *http.Request) { + close(started) + select { + case <-releaseResponse: + _ = json.NewEncoder(w).Encode(release("1.2.0", "linux", "amd64")) + case <-r.Context().Done(): + } + }) + first := c.Start(context.Background(), "1.0.0") + t.Cleanup(first.Stop) + <-started + if second := c.Start(context.Background(), "1.0.0"); second != nil { + second.Stop() + t.Fatal("concurrent check did not skip the held lock") + } + close(releaseResponse) + if complete(t, first, true) != "1.2.0" { + t.Fatal("missing alert") + } + third := c.Start(context.Background(), "1.0.0") + if third.Cached != "" || complete(t, third, true) != "" { + t.Fatal("duplicate alert") + } +} + +func TestCorruptCacheAndUnavailableCache(t *testing.T) { + c := fixture(t, func(w http.ResponseWriter, _ *http.Request) { + _ = json.NewEncoder(w).Encode(release("1.2.0", "linux", "amd64")) + }) + if err := os.WriteFile(c.CachePath, []byte("{"), 0600); err != nil { + t.Fatal(err) + } + if complete(t, c.Start(context.Background(), "1.0.0"), true) != "1.2.0" { + t.Fatal("did not recover from corrupt cache") + } + entries, _ := os.ReadDir(filepath.Dir(c.CachePath)) + for _, entry := range entries { + if strings.HasPrefix(entry.Name(), "update-") { + t.Fatal("temporary cache file was not removed") + } + } + c.CachePath = filepath.Join(c.CachePath, "not-a-directory", "update.json") + if check := c.Start(context.Background(), "1.0.0"); check != nil { + check.Stop() + t.Fatal("started with unavailable cache") + } +} + +func TestInvalidBuildAndCanceledContextHaveNoSideEffects(t *testing.T) { + c := fixture(t, func(http.ResponseWriter, *http.Request) { t.Error("unexpected request") }) + c.CachePath = filepath.Join(t.TempDir(), "not-created", "update.json") + for _, version := range []string{"dev", "test", "1.0.0-rc.1", "1.0.0-SNAPSHOT"} { + if c.Start(context.Background(), version) != nil { + t.Fatal("check started for", version) + } + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if c.Start(ctx, "1.0.0") != nil { + t.Fatal("check started after cancellation") + } + if _, err := os.Stat(filepath.Dir(c.CachePath)); !os.IsNotExist(err) { + t.Fatal("created cache directory") + } +} + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +func TestRequestDeadline(t *testing.T) { + c := fixture(t, func(http.ResponseWriter, *http.Request) { t.Error("unexpected server request") }) + c.HTTP = &http.Client{Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) { + deadline, ok := r.Context().Deadline() + if remaining := time.Until(deadline); !ok || remaining <= 0 || remaining > 3*time.Second { + t.Error("request does not have a three-second deadline") + } + return nil, errors.New("offline") + })} + if complete(t, c.Start(context.Background(), "1.0.0"), true) != "" || c.read().CheckedAt.IsZero() { + t.Fatal("offline check was not saved as a completed failure") + } +} + +func TestCacheWriteFailureReleasesLock(t *testing.T) { + c := fixture(t, func(w http.ResponseWriter, _ *http.Request) { + _ = json.NewEncoder(w).Encode(release("1.2.0", "linux", "amd64")) + }) + // A directory at the cache file path prevents atomic replacement on all OSes. + if err := os.Mkdir(c.CachePath, 0700); err != nil { + t.Fatal(err) + } + for range 2 { + if got := complete(t, c.Start(context.Background(), "1.0.0"), true); got != "" { + t.Fatal("claimed an alert without a writable cache") + } + } +}