From 0b05b273c01ead9176343581aa1bbd8f27b6aaa5 Mon Sep 17 00:00:00 2001 From: Sameen Karim Date: Fri, 2 Oct 2026 10:58:14 -0400 Subject: [PATCH 1/3] new release update checker --- README.md | 27 ++ cmd/root.go | 72 ++- cmd/root_test.go | 248 +++++++++++ docs/src/content/docs/reference/cli.md | 37 ++ go.mod | 3 +- go.sum | 2 + internal/update/update.go | 290 ++++++++++++ internal/update/update_test.go | 592 +++++++++++++++++++++++++ 8 files changed, 1261 insertions(+), 10 deletions(-) create mode 100644 internal/update/update.go create mode 100644 internal/update/update_test.go diff --git a/README.md b/README.md index 893c0688..4faab13d 100644 --- a/README.md +++ b/README.md @@ -12,6 +12,33 @@ gh extension install github/gh-stack Requires the [GitHub CLI](https://cli.github.com/) (`gh`) v2.0+ and Git 2.36+. +### Updating + +```sh +gh extension upgrade stack +``` + +Official, unpinned stable release installations check for a newer +[latest release](https://github.com/github/gh-stack/releases/latest) in the +background, at most once every 24 hours. When an update is available, successful +commands can append an upgrade notice to stderr, also at most once every 24 +hours. This includes non-interactive use; stdout and JSON output are unchanged. +The check and reminder timestamps are shared across repositories for your user. +They are stored as YAML in `gh-stack/state.yml` beneath GitHub CLI's state +directory (`~/.local/state/gh` by default on macOS/Linux). + +Commands never wait for a release check. Short commands may finish before a +notice is ready, and failed or interrupted checks still count toward the daily +check limit. Completed checks are cached so a later command can show the notice. +No upgrades happen automatically, and development, locally linked, prerelease, +and pinned installations are excluded. Help, version, and completion commands +do not run the notifier. + +Set `GH_STACK_NO_UPDATE_NOTIFIER=1` to disable these checks and notices, including +in CI or scripts. Any non-empty value disables the notifier. Optional check +failures do not affect command exit codes; set `GH_DEBUG=1` for diagnostics. +Install a release containing this feature to receive notices for future releases. + ## AI agent integration Install the gh-stack skill so your AI coding agent knows how to work with stacked PRs and the `gh stack` CLI: diff --git a/cmd/root.go b/cmd/root.go index f4100c74..98bb3498 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -1,18 +1,27 @@ package cmd import ( + "context" "errors" "fmt" "os" "github.com/github/gh-stack/internal/config" "github.com/github/gh-stack/internal/theme" + "github.com/github/gh-stack/internal/update" "github.com/spf13/cobra" ) +type updateResult struct { + notify update.Notification + err error +} + func RootCmd() *cobra.Command { - cfg := config.New() + return newRootCmd(config.New(), startUpdateCheck) +} +func newRootCmd(cfg *config.Config, startCheck func(context.Context, string) <-chan updateResult) *cobra.Command { root := &cobra.Command{ Use: "stack ", Short: "Manage stacked branches and pull requests", @@ -36,10 +45,31 @@ locally, then push to GitHub to create your stack of PRs.`, Version: Version, SilenceUsage: true, SilenceErrors: true, - // Honor GH_STACK_THEME (auto|light|dark) before any command renders - PersistentPreRun: func(_ *cobra.Command, _ []string) { - theme.ApplyOverride() - }, + } + + var updates <-chan updateResult + root.PersistentPreRun = func(cmd *cobra.Command, _ []string) { + theme.ApplyOverride() + updates = nil + if update.Enabled(root.Version) && !isHelpOrCompletionCommand(cmd) { + updates = startCheck(cmd.Context(), root.Version) + } + } + root.PersistentPostRun = func(cmd *cobra.Command, _ []string) { + if cmd.Context().Err() != nil { + return + } + select { + case result := <-updates: + if result.notify != nil { + result.err = errors.Join(result.err, result.notify(cmd.ErrOrStderr())) + } + if result.err != nil && os.Getenv("GH_DEBUG") != "" { + fmt.Fprintf(cmd.ErrOrStderr(), "debug: gh-stack update notification: %v\n", result.err) + } + default: + // Never wait for a release check if the command finishes quickly + } } root.SetVersionTemplate("gh stack version {{.Version}}\n") @@ -153,15 +183,39 @@ locally, then push to GitHub to create your stack of PRs.`, return root } -func Execute() { - cmd := RootCmd() +func startUpdateCheck(ctx context.Context, version string) <-chan updateResult { + results := make(chan updateResult, 1) + go func() { + notify, err := update.Check(ctx, version) + results <- updateResult{notify: notify, err: err} + }() + return results +} + +func isHelpOrCompletionCommand(cmd *cobra.Command) bool { + for c := cmd; c != nil; c = c.Parent() { + switch c.Name() { + case "help", "completion", cobra.ShellCompRequestCmd, cobra.ShellCompNoDescRequestCmd: + return true + } + } + return false +} + +func execute(cmd *cobra.Command, args []string) error { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() // Wrap in a "gh" parent so help output shows "gh stack" instead of just "stack". wrapCmd := &cobra.Command{Use: "gh", SilenceUsage: true, SilenceErrors: true} wrapCmd.AddCommand(cmd) - wrapCmd.SetArgs(append([]string{"stack"}, os.Args[1:]...)) + wrapCmd.SetArgs(append([]string{"stack"}, args...)) + return wrapCmd.ExecuteContext(ctx) +} - if err := wrapCmd.Execute(); err != nil { +func Execute() { + cmd := RootCmd() + if err := execute(cmd, os.Args[1:]); err != nil { var exitErr *ExitError if errors.As(err, &exitErr) { os.Exit(exitErr.Code) diff --git a/cmd/root_test.go b/cmd/root_test.go index 47328fd9..18416872 100644 --- a/cmd/root_test.go +++ b/cmd/root_test.go @@ -2,8 +2,15 @@ package cmd import ( "bytes" + "context" + "errors" + "fmt" + "io" "testing" + "time" + "github.com/github/gh-stack/internal/config" + "github.com/spf13/cobra" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -39,3 +46,244 @@ func TestRootCmd_HelpOutput(t *testing.T) { assert.Contains(t, output, "Learn more:") assert.Contains(t, output, "https://gh.io/stacks") } + +func newUpdateTestRoot(t *testing.T, startCheck func(context.Context, string) <-chan updateResult) (*cobra.Command, *bytes.Buffer, *bytes.Buffer) { + t.Helper() + t.Setenv("GH_STACK_NO_UPDATE_NOTIFIER", "") + t.Setenv("GH_DEBUG", "") + cfg, outR, errR := config.NewTestConfig() + t.Cleanup(func() { + cfg.Out.Close() + cfg.Err.Close() + outR.Close() + errR.Close() + }) + root := newRootCmd(cfg, startCheck) + root.Version = "0.1.1" + stdout, stderr := &bytes.Buffer{}, &bytes.Buffer{} + root.SetOut(stdout) + root.SetErr(stderr) + return root, stdout, stderr +} + +func readyUpdate(result updateResult) <-chan updateResult { + results := make(chan updateResult, 1) + results <- result + return results +} + +func testUpdateNotification(out io.Writer) error { + _, err := fmt.Fprint(out, "\nA new release of gh-stack is available: 0.1.1 -> 0.2.0\nTo upgrade, run: gh extension upgrade stack\n") + return err +} + +func TestRootCmd_UpdateNotice(t *testing.T) { + var checkContext context.Context + checks := 0 + root, stdout, stderr := newUpdateTestRoot(t, func(ctx context.Context, version string) <-chan updateResult { + checks++ + checkContext = ctx + assert.Equal(t, "0.1.1", version) + return readyUpdate(updateResult{notify: testUpdateNotification}) + }) + t.Setenv("CI", "true") + t.Setenv("GH_NO_EXTENSION_UPDATE_NOTIFIER", "1") + root.AddCommand(&cobra.Command{ + Use: "probe", + RunE: func(cmd *cobra.Command, _ []string) error { + fmt.Fprintln(cmd.OutOrStdout(), `{"ok":true}`) + fmt.Fprintln(cmd.ErrOrStderr(), "Command complete.") + return nil + }, + }) + + require.NoError(t, execute(root, []string{"probe"})) + assert.Equal(t, 1, checks) + assert.Equal(t, "{\"ok\":true}\n", stdout.String()) + assert.Equal(t, "Command complete.\n\nA new release of gh-stack is available: 0.1.1 -> 0.2.0\nTo upgrade, run: gh extension upgrade stack\n", stderr.String()) + require.ErrorIs(t, checkContext.Err(), context.Canceled) +} + +func TestRootCmd_UpdateNoticeExcludedCommands(t *testing.T) { + tests := []struct { + name string + args []string + version string + disabled bool + wantErr bool + }{ + {name: "no command"}, + {name: "help flag", args: []string{"--help"}}, + {name: "help command", args: []string{"help", "probe"}}, + {name: "subcommand help", args: []string{"probe", "--help"}}, + {name: "version", args: []string{"--version"}}, + {name: "completion generation", args: []string{"completion", "bash"}}, + {name: "completion request", args: []string{cobra.ShellCompRequestCmd, "probe", ""}}, + {name: "completion without descriptions", args: []string{cobra.ShellCompNoDescRequestCmd, "probe", ""}}, + {name: "invalid command", args: []string{"unknown"}, wantErr: true}, + {name: "invalid flag", args: []string{"probe", "--unknown"}, wantErr: true}, + {name: "invalid argument", args: []string{"probe", "extra"}, wantErr: true}, + {name: "opt-out", args: []string{"probe"}, disabled: true}, + {name: "development build", args: []string{"probe"}, version: "dev"}, + {name: "prerelease build", args: []string{"probe"}, version: "0.2.0-rc.1"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + checks := 0 + root, stdout, stderr := newUpdateTestRoot(t, func(context.Context, string) <-chan updateResult { + checks++ + return readyUpdate(updateResult{notify: testUpdateNotification}) + }) + if tt.version != "" { + root.Version = tt.version + } + if tt.disabled { + t.Setenv("GH_STACK_NO_UPDATE_NOTIFIER", "1") + } + root.AddCommand(&cobra.Command{ + Use: "probe", + Args: cobra.NoArgs, + RunE: func(*cobra.Command, []string) error { return nil }, + }) + root.SetArgs(append([]string{}, tt.args...)) + err := root.Execute() + if tt.wantErr { + require.Error(t, err) + } else { + require.NoError(t, err) + } + assert.Zero(t, checks) + assert.NotContains(t, stdout.String()+stderr.String(), "A new release") + if tt.name == "version" { + assert.Equal(t, "gh stack version 0.1.1\n", stdout.String()) + } + }) + } +} + +func TestRootCmd_UpdateNoticePreservesCommandErrors(t *testing.T) { + for _, commandErr := range []error{ + errors.New("command failed"), + ErrConflict, + fmt.Errorf("wrapped: %w", ErrInvalidArgs), + ErrSilent, + context.Canceled, + } { + t.Run(commandErr.Error(), func(t *testing.T) { + var checkContext context.Context + root, stdout, stderr := newUpdateTestRoot(t, func(ctx context.Context, _ string) <-chan updateResult { + checkContext = ctx + return readyUpdate(updateResult{notify: testUpdateNotification}) + }) + root.AddCommand(&cobra.Command{ + Use: "probe", + RunE: func(cmd *cobra.Command, _ []string) error { + fmt.Fprintln(cmd.ErrOrStderr(), "Original diagnostic.") + return commandErr + }, + }) + err := execute(root, []string{"probe"}) + require.ErrorIs(t, err, commandErr) + assert.Empty(t, stdout.String()) + assert.Equal(t, "Original diagnostic.\n", stderr.String()) + require.ErrorIs(t, checkContext.Err(), context.Canceled) + }) + } +} + +func TestRootCmd_UpdateDiagnostics(t *testing.T) { + tests := []struct { + name string + debug bool + result updateResult + }{ + {name: "quiet failure", result: updateResult{err: errors.New("offline")}}, + {name: "debug failure", debug: true, result: updateResult{err: errors.New("offline")}}, + {name: "recovered cache", debug: true, result: updateResult{notify: testUpdateNotification, err: errors.New("invalid cache")}}, + {name: "notification failure", debug: true, result: updateResult{notify: func(io.Writer) error { return errors.New("state is unwritable") }}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + root, stdout, stderr := newUpdateTestRoot(t, func(context.Context, string) <-chan updateResult { + return readyUpdate(tt.result) + }) + if tt.debug { + t.Setenv("GH_DEBUG", "1") + } + root.AddCommand(&cobra.Command{Use: "probe", RunE: func(*cobra.Command, []string) error { return nil }}) + require.NoError(t, execute(root, []string{"probe"})) + assert.Empty(t, stdout.String()) + if tt.debug { + assert.Contains(t, stderr.String(), "debug: gh-stack update notification:") + } else { + assert.Empty(t, stderr.String()) + } + if tt.name == "recovered cache" { + assert.Contains(t, stderr.String(), "gh extension upgrade stack") + assert.Contains(t, stderr.String(), "invalid cache") + } + }) + } +} + +func TestRootCmd_DoesNotWaitForUpdate(t *testing.T) { + for _, fail := range []bool{false, true} { + t.Run(fmt.Sprintf("command failure=%t", fail), func(t *testing.T) { + results := make(chan updateResult, 1) + workerDone := make(chan struct{}) + root, stdout, stderr := newUpdateTestRoot(t, func(ctx context.Context, _ string) <-chan updateResult { + go func() { + <-ctx.Done() + results <- updateResult{err: ctx.Err()} + close(workerDone) + }() + return results + }) + root.AddCommand(&cobra.Command{ + Use: "probe", + RunE: func(*cobra.Command, []string) error { + if fail { + return ErrConflict + } + return nil + }, + }) + done := make(chan error, 1) + go func() { done <- execute(root, []string{"probe"}) }() + select { + case err := <-done: + if fail { + require.ErrorIs(t, err, ErrConflict) + } else { + require.NoError(t, err) + } + case <-time.After(5 * time.Second): + // Unblock a regressed post-run hook before failing the test. + results <- updateResult{} + <-done + <-workerDone + t.Fatal("command waited for an unfinished update check") + } + select { + case <-workerDone: + case <-time.After(5 * time.Second): + t.Fatal("check was not canceled or its result sender was blocked") + } + assert.Empty(t, stdout.String()) + assert.Empty(t, stderr.String()) + }) + } +} + +func TestStartUpdateCheck_BufferedResult(t *testing.T) { + t.Setenv("GH_STACK_NO_UPDATE_NOTIFIER", "") + results := startUpdateCheck(context.Background(), "dev") + assert.Equal(t, 1, cap(results), "a late result must not block its sender") + select { + case result := <-results: + assert.Nil(t, result.notify) + require.NoError(t, result.err) + case <-time.After(5 * time.Second): + t.Fatal("development build check did not return") + } +} diff --git a/docs/src/content/docs/reference/cli.md b/docs/src/content/docs/reference/cli.md index e31b5300..a07929e2 100644 --- a/docs/src/content/docs/reference/cli.md +++ b/docs/src/content/docs/reference/cli.md @@ -23,6 +23,42 @@ All linked worktrees share `/gh-stack` and gh-stack recovery journal For repositories created with `git init --separate-git-dir`, main-worktree invocation and existing absolute/relative `core.worktree` backlinks are supported, including settings in the main `config.worktree`. The discovery limitation is only linked invocation without a main-worktree backlink. If the operation requires that main owner, it fails with actionable guidance; unaffected worktrees continue. Administration directories are never used as checkout destinations. +### Updating + +```sh +gh extension upgrade stack +``` + +Official, unpinned stable release installations check the +[latest gh-stack release](https://github.com/github/gh-stack/releases/latest) +in the background, at most once every 24 hours. Successful commands can append +an upgrade notice to stderr, also at most once every 24 hours, until you upgrade. +The timestamps are stored per user and shared across repositories, independently +of GitHub CLI's own update checks. +The YAML state file is `gh-stack/state.yml` beneath GitHub CLI's state directory: +`~/.local/state/gh/gh-stack/state.yml` by default on macOS/Linux, or +`$XDG_STATE_HOME/gh/gh-stack/state.yml` when `XDG_STATE_HOME` is set. + +Notices also appear in non-interactive use, including CI. Stdout and JSON output +remain unchanged. No upgrade happens automatically, and commands never wait for +the check: a short command may finish before a notice is ready. Completed checks +are cached for later commands; failed or interrupted attempts still count toward +the daily check limit. + +Development builds, locally linked installations, prereleases, and pinned +installations are excluded. Help, version, and shell completion commands do not +run the notifier. Users must first install a release containing this feature to +receive notices about subsequent releases. + +To disable the notifier, including its network checks: + +```sh +GH_STACK_NO_UPDATE_NOTIFIER=1 gh stack view +``` + +Any non-empty value disables it. Update-check failures never change the command's +exit code; use `GH_DEBUG=1` to see diagnostic errors when a check completes. + --- ## Stack Management @@ -695,6 +731,7 @@ gh stack feedback "Support for reordering branches" |----------|--------|-------------| | `GH_STACK_THEME` | `auto` (default), `light`, `dark` | Controls the color palette of the interactive screens (`submit`, `modify`, `view`) and all colored command output. Colors adapt to your terminal background automatically; set this to force the light or dark palette when a terminal doesn't report its background (some SSH or `tmux` setups). | | `GH_STACK_HYPERLINKS` | `0`, `1` | Disables or enables OSC 8 hyperlinks when terminal detection is incorrect. Unsupported terminals show the full URL by default. | +| `GH_STACK_NO_UPDATE_NOTIFIER` | Any non-empty value | Disables gh-stack's background release checks and upgrade notices. For example, set to `1` in automation that requires quiet stderr. | ```sh # Force the light palette for one command diff --git a/go.mod b/go.mod index eed0bc59..05c0a983 100644 --- a/go.mod +++ b/go.mod @@ -15,9 +15,11 @@ require ( github.com/muesli/termenv v0.16.0 github.com/spf13/cobra v1.10.2 github.com/stretchr/testify v1.11.1 + golang.org/x/mod v0.36.0 golang.org/x/sys v0.45.0 golang.org/x/term v0.43.0 golang.org/x/text v0.37.0 + gopkg.in/yaml.v3 v3.0.1 ) require ( @@ -61,5 +63,4 @@ require ( github.com/yuin/goldmark v1.8.2 // indirect github.com/yuin/goldmark-emoji v1.0.6 // indirect golang.org/x/net v0.55.0 // indirect - gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/go.sum b/go.sum index e0970f3b..8d2ac163 100644 --- a/go.sum +++ b/go.sum @@ -148,6 +148,8 @@ golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5y golang.org/x/exp v0.0.0-20231006140011-7918f672742d h1:jtJma62tbqLibJ5sFQz8bKtEM8rJBtfilJ2qTU199MI= golang.org/x/exp v0.0.0-20231006140011-7918f672742d/go.mod h1:ldy0pHrwJyGW56pPQzzkH36rKxoZW1tw7ZJpeKx+hdo= golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= +golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4= +golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= diff --git a/internal/update/update.go b/internal/update/update.go new file mode 100644 index 00000000..ccbb34f7 --- /dev/null +++ b/internal/update/update.go @@ -0,0 +1,290 @@ +package update + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "os" + "path/filepath" + "runtime" + "strings" + "time" + + "github.com/cli/go-gh/v2/pkg/auth" + ghconfig "github.com/cli/go-gh/v2/pkg/config" + "golang.org/x/mod/semver" + "gopkg.in/yaml.v3" +) + +const ( + checkInterval = 24 * time.Hour + requestTimeout = 2 * time.Second + latestReleaseURL = "https://api.github.com/repos/github/gh-stack/releases/latest" +) + +// Notification records and writes a pending notice when a command succeeds. +type Notification func(io.Writer) error + +// Enabled excludes opt-outs and builds that cannot be compared to stable releases. +func Enabled(version string) bool { + return os.Getenv("GH_STACK_NO_UPDATE_NOTIFIER") == "" && stableVersion(version) != "" +} + +// Check returns a pending notice without displaying it. A recovered state error +// may accompany a notice so callers can report diagnostics without losing it. +func Check(ctx context.Context, version string) (Notification, error) { + if !Enabled(version) { + return nil, nil + } + + assetSuffix := runtime.GOOS + "-" + runtime.GOARCH + if runtime.GOOS == "windows" { + assetSuffix += ".exe" + } + c := checker{ + statePath: filepath.Join(ghconfig.StateDir(), "gh-stack", "state.yml"), + executable: os.Executable, + client: &http.Client{Timeout: requestTimeout}, + token: func() string { + token, _ := auth.TokenForHost("github.com") + return token + }, + now: time.Now, + assetSuffix: assetSuffix, + } + return c.check(ctx, version) +} + +type checker struct { + statePath string + executable func() (string, error) + client *http.Client + token func() string + now func() time.Time + assetSuffix string +} + +type state struct { + LastCheckedAt time.Time `yaml:"last_checked_at"` + LastNotifiedAt time.Time `yaml:"last_notified_at"` + LatestVersion string `yaml:"latest_version,omitempty"` +} + +var errInvalidState = errors.New("invalid update notification state") + +func (c checker) check(ctx context.Context, version string) (Notification, error) { + if !Enabled(version) { + return nil, nil + } + ctx, cancel := context.WithTimeout(ctx, requestTimeout) + defer cancel() + + if err := ctx.Err(); err != nil { + return nil, err + } + current := stableVersion(version) + eligible, err := c.eligible(current) + if err != nil || !eligible { + return nil, err + } + if !filepath.IsAbs(c.statePath) { + return nil, fmt.Errorf("update state path must be absolute: %s", c.statePath) + } + + now := c.now() + s, err := readState(c.statePath, now) + var diagnostic error + switch { + case err == nil, errors.Is(err, os.ErrNotExist): + case errors.Is(err, errInvalidState): + diagnostic = err + s = state{} + default: + return nil, err + } + + if s.LastCheckedAt.IsZero() || now.Sub(s.LastCheckedAt) >= checkInterval { + // Reserve the daily attempt before the request, including attempts + // interrupted when a short command exits. + s.LastCheckedAt = now + s.LatestVersion = "" + if err := writeState(c.statePath, s); err != nil { + return nil, errors.Join(diagnostic, err) + } + latest, err := c.fetchLatest(ctx, current) + if err != nil { + return nil, errors.Join(diagnostic, err) + } + s.LatestVersion = latest + if err := writeState(c.statePath, s); err != nil { + return nil, errors.Join(diagnostic, err) + } + } + + if !s.noticeDue(current, now) { + return nil, diagnostic + } + return func(out io.Writer) error { + now := c.now() + // Re-read before delivery so a previously prepared notice cannot + // overwrite a later check or reminder. + s, err := readState(c.statePath, now) + if err != nil { + return err + } + if !s.noticeDue(current, now) { + return nil + } + s.LastNotifiedAt = now + if err := writeState(c.statePath, s); err != nil { + return err + } + _, err = fmt.Fprintf(out, + "\nA new release of gh-stack is available: %s -> %s\nTo upgrade, run: gh extension upgrade stack\n", + strings.TrimPrefix(current, "v"), strings.TrimPrefix(s.LatestVersion, "v")) + return err + }, diagnostic +} + +func (s state) noticeDue(current string, now time.Time) bool { + return s.LatestVersion != "" && + semver.Compare(s.LatestVersion, current) > 0 && + now.Sub(s.LastCheckedAt) < checkInterval && + (s.LastNotifiedAt.IsZero() || now.Sub(s.LastNotifiedAt) >= checkInterval) +} + +func stableVersion(version string) string { + version = "v" + strings.TrimPrefix(version, "v") + if !semver.IsValid(version) || semver.Prerelease(version) != "" { + return "" + } + return version +} + +func (c checker) eligible(current string) (bool, error) { + executable, err := c.executable() + if err != nil { + return false, fmt.Errorf("locating installed executable: %w", err) + } + executable, err = filepath.EvalSymlinks(executable) + if err != nil { + return false, fmt.Errorf("resolving installed executable: %w", err) + } + data, err := os.ReadFile(filepath.Join(filepath.Dir(executable), "manifest.yml")) + if errors.Is(err, os.ErrNotExist) { + return false, nil + } + if err != nil { + return false, fmt.Errorf("reading extension manifest: %w", err) + } + var manifest struct { + Owner string + Name string + Host string + Tag string + IsPinned bool + } + if err := yaml.Unmarshal(data, &manifest); err != nil { + return false, fmt.Errorf("decoding extension manifest: %w", err) + } + if !strings.EqualFold(manifest.Host, "github.com") || + !strings.EqualFold(manifest.Owner, "github") || + !strings.EqualFold(manifest.Name, "gh-stack") || manifest.IsPinned { + return false, nil + } + installed := stableVersion(manifest.Tag) + if installed == "" || semver.Compare(installed, current) != 0 { + return false, fmt.Errorf("extension manifest tag %q does not match running version %q", manifest.Tag, current) + } + return true, nil +} + +func (c checker) fetchLatest(ctx context.Context, current string) (string, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, latestReleaseURL, nil) + if err != nil { + return "", err + } + req.Header.Set("Accept", "application/vnd.github+json") + req.Header.Set("X-GitHub-Api-Version", "2022-11-28") + req.Header.Set("User-Agent", "gh-stack/"+strings.TrimPrefix(current, "v")) + if token := c.token(); token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + resp, err := c.client.Do(req) + if err != nil { + return "", fmt.Errorf("fetching latest gh-stack release: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("fetching latest gh-stack release: HTTP %d", resp.StatusCode) + } + var release struct { + Tag string `json:"tag_name"` + Draft bool + Prerelease bool + Assets []struct{ Name string } + } + if err := json.NewDecoder(io.LimitReader(resp.Body, 1024*1024)).Decode(&release); err != nil { + return "", fmt.Errorf("decoding latest gh-stack release: %w", err) + } + if release.Draft || release.Prerelease { + return "", nil + } + latest := stableVersion(release.Tag) + if latest == "" { + return "", fmt.Errorf("latest gh-stack release has an invalid stable version: %q", release.Tag) + } + for _, asset := range release.Assets { + if strings.HasSuffix(asset.Name, c.assetSuffix) { + return latest, nil + } + } + return "", nil +} + +func readState(path string, now time.Time) (state, error) { + data, err := os.ReadFile(path) + if err != nil { + return state{}, fmt.Errorf("reading update state: %w", err) + } + var s state + if err := yaml.Unmarshal(data, &s); err != nil { + return state{}, fmt.Errorf("%w: %w", errInvalidState, err) + } + if s.LastCheckedAt.After(now) || s.LastNotifiedAt.After(now) { + return state{}, fmt.Errorf("%w: timestamp is in the future", errInvalidState) + } + if s.LatestVersion != "" && (stableVersion(s.LatestVersion) != s.LatestVersion || s.LastCheckedAt.IsZero()) { + return state{}, fmt.Errorf("%w: invalid cached release", errInvalidState) + } + return s, nil +} + +func writeState(path string, s state) error { + data, err := yaml.Marshal(s) + if err != nil { + return fmt.Errorf("encoding update state: %w", err) + } + if err := os.MkdirAll(filepath.Dir(path), 0700); err != nil { + return fmt.Errorf("creating update state directory: %w", err) + } + f, err := os.CreateTemp(filepath.Dir(path), ".state-*.yml") + if err != nil { + return fmt.Errorf("creating update state file: %w", err) + } + defer os.Remove(f.Name()) + if _, err := f.Write(data); err != nil { + f.Close() + return fmt.Errorf("writing update state: %w", err) + } + if err := f.Close(); err != nil { + return fmt.Errorf("closing update state: %w", err) + } + if err := os.Rename(f.Name(), path); err != nil { + return fmt.Errorf("replacing update state: %w", err) + } + return nil +} diff --git a/internal/update/update_test.go b/internal/update/update_test.go new file mode 100644 index 00000000..7aa72ab2 --- /dev/null +++ b/internal/update/update_test.go @@ -0,0 +1,592 @@ +package update + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "net/http" + "os" + "os/exec" + "path/filepath" + "runtime" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +func releaseResponse(tag string) *http.Response { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(fmt.Sprintf( + `{"tag_name":%q,"assets":[{"name":"linux-amd64"}]}`, tag))), + Header: make(http.Header), + } +} + +func newTestChecker(t *testing.T) checker { + t.Helper() + t.Setenv("GH_STACK_NO_UPDATE_NOTIFIER", "") + dir := t.TempDir() + executable := filepath.Join(dir, "gh-stack") + require.NoError(t, os.WriteFile(executable, nil, 0600)) + c := checker{ + statePath: filepath.Join(dir, "state", "state.yml"), + executable: func() (string, error) { return executable, nil }, + client: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return releaseResponse("v0.2.0"), nil + })}, + token: func() string { return "test-token" }, + now: func() time.Time { return time.Date(2026, 1, 15, 12, 0, 0, 0, time.UTC) }, + assetSuffix: "linux-amd64", + } + writeManifest(t, c, "v0.1.1", "") + return c +} + +func writeManifest(t *testing.T, c checker, tag, extra string) { + t.Helper() + executable, err := c.executable() + require.NoError(t, err) + data := fmt.Sprintf("owner: github\nname: gh-stack\nhost: github.com\ntag: %s\n%s", tag, extra) + require.NoError(t, os.WriteFile(filepath.Join(filepath.Dir(executable), "manifest.yml"), []byte(data), 0600)) +} + +func TestCheck_Versions(t *testing.T) { + tests := []struct { + name string + current string + latest string + wantNotice bool + wantCheck bool + }{ + {"newer release", "0.1.1", "v0.2.0", true, true}, + {"current", "0.2.0", "v0.2.0", false, true}, + {"installed ahead", "0.3.0", "v0.2.0", false, true}, + {"numeric ordering", "0.9.0", "v0.10.0", true, true}, + {"not lexical ordering", "0.10.0", "v0.9.0", false, true}, + {"installed v prefix", "v0.1.1", "v0.2.0", true, true}, + {"release without v prefix", "0.1.1", "0.2.0", true, true}, + {"build metadata", "0.2.0+local", "v0.2.0", false, true}, + {"development build", "dev", "v0.2.0", false, false}, + {"prerelease build", "0.2.0-rc.1", "v0.2.0", false, false}, + {"invalid build", "garbage", "v0.2.0", false, false}, + {"git describe build", "0.1.1-4-g12345678", "v0.2.0", false, false}, + {"invalid leading zeros", "00.1.1", "v0.2.0", false, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c := newTestChecker(t) + writeManifest(t, c, tt.current, "") + requests := 0 + c.client.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + requests++ + return releaseResponse(tt.latest), nil + }) + notice, err := c.check(context.Background(), tt.current) + require.NoError(t, err) + assert.Equal(t, tt.wantNotice, notice != nil) + assert.Equal(t, tt.wantCheck, requests == 1) + }) + } +} + +func TestCheck_Request(t *testing.T) { + for _, token := range []string{"", "test-token"} { + t.Run("token="+token, func(t *testing.T) { + c := newTestChecker(t) + t.Setenv("GH_HOST", "enterprise.example.com") + c.token = func() string { return token } + c.client.Transport = roundTripFunc(func(req *http.Request) (*http.Response, error) { + assert.Equal(t, http.MethodGet, req.Method) + assert.Equal(t, latestReleaseURL, req.URL.String()) + assert.Equal(t, "application/vnd.github+json", req.Header.Get("Accept")) + assert.Equal(t, "2022-11-28", req.Header.Get("X-GitHub-Api-Version")) + assert.Equal(t, "gh-stack/0.1.1", req.Header.Get("User-Agent")) + if token == "" { + assert.Empty(t, req.Header.Get("Authorization")) + } else { + assert.Equal(t, "Bearer "+token, req.Header.Get("Authorization")) + } + deadline, ok := req.Context().Deadline() + assert.True(t, ok) + assert.LessOrEqual(t, time.Until(deadline), requestTimeout) + assert.Greater(t, time.Until(deadline), time.Duration(0)) + return releaseResponse("v0.2.0"), nil + }) + notice, err := c.check(context.Background(), "0.1.1") + require.NoError(t, err) + require.NotNil(t, notice) + }) + } +} + +func TestCheck_Installation(t *testing.T) { + tests := []struct { + name string + manifest string + wantErr string + }{ + {"local installation", "", ""}, + {"pinned", "owner: github\nname: gh-stack\nhost: github.com\ntag: v0.1.1\nispinned: true\n", ""}, + {"fork", "owner: someone\nname: gh-stack\nhost: github.com\ntag: v0.1.1\n", ""}, + {"another extension", "owner: github\nname: gh-other\nhost: github.com\ntag: v0.1.1\n", ""}, + {"enterprise host", "owner: github\nname: gh-stack\nhost: example.com\ntag: v0.1.1\n", ""}, + {"mismatched version", "owner: github\nname: gh-stack\nhost: github.com\ntag: v0.0.1\n", "does not match"}, + {"malformed manifest", "ispinned: [", "decoding extension manifest"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c := newTestChecker(t) + executable, err := c.executable() + require.NoError(t, err) + path := filepath.Join(filepath.Dir(executable), "manifest.yml") + if tt.manifest == "" { + require.NoError(t, os.Remove(path)) + } else { + require.NoError(t, os.WriteFile(path, []byte(tt.manifest), 0600)) + } + c.client.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + t.Error("ineligible installation made a network request") + return nil, errors.New("unexpected request") + }) + notice, err := c.check(context.Background(), "0.1.1") + if tt.wantErr != "" { + require.ErrorContains(t, err, tt.wantErr) + } else { + require.NoError(t, err) + } + assert.Nil(t, notice) + assert.NoFileExists(t, c.statePath) + }) + } +} + +func TestCheck_InstallationErrors(t *testing.T) { + for _, kind := range []string{"executable lookup", "missing executable", "unreadable manifest"} { + t.Run(kind, func(t *testing.T) { + c := newTestChecker(t) + switch kind { + case "executable lookup": + c.executable = func() (string, error) { return "", errors.New("lookup failed") } + case "missing executable": + c.executable = func() (string, error) { return filepath.Join(t.TempDir(), "missing"), nil } + case "unreadable manifest": + executable, err := c.executable() + require.NoError(t, err) + path := filepath.Join(filepath.Dir(executable), "manifest.yml") + require.NoError(t, os.Remove(path)) + require.NoError(t, os.Mkdir(path, 0700)) + } + notice, err := c.check(context.Background(), "0.1.1") + require.Error(t, err) + assert.Nil(t, notice) + assert.NoFileExists(t, c.statePath) + }) + } +} + +func TestCheck_OptOut(t *testing.T) { + for _, value := range []string{"1", "true", "0"} { + t.Run(value, func(t *testing.T) { + c := newTestChecker(t) + t.Setenv("GH_STACK_NO_UPDATE_NOTIFIER", value) + c.executable = func() (string, error) { + t.Error("opt-out must not inspect the installation") + return "", errors.New("unexpected lookup") + } + notice, err := c.check(context.Background(), "0.1.1") + require.NoError(t, err) + assert.Nil(t, notice) + assert.NoFileExists(t, c.statePath) + notice, err = Check(context.Background(), "0.1.1") + require.NoError(t, err) + assert.Nil(t, notice) + }) + } +} + +func TestCheck_DailyCadence(t *testing.T) { + c := newTestChecker(t) + now := c.now() + start := now + c.now = func() time.Time { return now } + requests := 0 + c.client.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + requests++ + return releaseResponse("v0.2.0"), nil + }) + + notice, err := c.check(context.Background(), "0.1.1") + require.NoError(t, err) + require.NotNil(t, notice) + var out bytes.Buffer + require.NoError(t, notice(&out)) + assert.Equal(t, "\nA new release of gh-stack is available: 0.1.1 -> 0.2.0\nTo upgrade, run: gh extension upgrade stack\n", out.String()) + + for _, elapsed := range []time.Duration{time.Minute, checkInterval - time.Nanosecond} { + now = start.Add(elapsed) + freshChecker := c + notice, err = freshChecker.check(context.Background(), "0.1.1") + require.NoError(t, err) + assert.Nil(t, notice) + assert.Equal(t, 1, requests) + } + + now = start.Add(checkInterval) + notice, err = c.check(context.Background(), "0.1.1") + require.NoError(t, err) + require.NotNil(t, notice) + assert.Equal(t, 2, requests) + require.NoError(t, notice(io.Discard)) + s, err := readState(c.statePath, now) + require.NoError(t, err) + assert.Equal(t, now, s.LastCheckedAt) + assert.Equal(t, now, s.LastNotifiedAt) +} + +func TestCheck_PendingNotice(t *testing.T) { + c := newTestChecker(t) + now := c.now() + start := now + c.now = func() time.Time { return now } + firstNotice, err := c.check(context.Background(), "0.1.1") + require.NoError(t, err) + require.NotNil(t, firstNotice) + s, err := readState(c.statePath, now) + require.NoError(t, err) + assert.True(t, s.LastNotifiedAt.IsZero()) + + c.client.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + t.Error("pending notice must use cached metadata") + return nil, errors.New("unexpected request") + }) + now = now.Add(time.Hour) + notice, err := c.check(context.Background(), "0.1.1") + require.NoError(t, err) + require.NotNil(t, notice) + require.NoError(t, notice(io.Discard)) + s, err = readState(c.statePath, now) + require.NoError(t, err) + assert.Equal(t, start, s.LastCheckedAt) + assert.Equal(t, now, s.LastNotifiedAt) + + var out bytes.Buffer + require.NoError(t, firstNotice(&out)) + assert.Empty(t, out.String(), "a prepared notice must recheck the last notification") +} + +func TestCheck_NotificationCadenceIndependentOfChecks(t *testing.T) { + c := newTestChecker(t) + now := c.now() + require.NoError(t, writeState(c.statePath, state{ + LastCheckedAt: now.Add(-checkInterval), + LastNotifiedAt: now.Add(-time.Hour), + LatestVersion: "v0.1.2", + })) + notice, err := c.check(context.Background(), "0.1.1") + require.NoError(t, err) + assert.Nil(t, notice, "a newly checked release must not bypass the reminder cooldown") + s, err := readState(c.statePath, now) + require.NoError(t, err) + assert.Equal(t, now, s.LastCheckedAt) + assert.Equal(t, "v0.2.0", s.LatestVersion) + + c.now = func() time.Time { return now.Add(23 * time.Hour) } + notice, err = c.check(context.Background(), "0.1.1") + require.NoError(t, err) + require.NotNil(t, notice) + require.NoError(t, notice(io.Discard)) +} + +func TestCheck_UpgradedInstallation(t *testing.T) { + c := newTestChecker(t) + notice, err := c.check(context.Background(), "0.1.1") + require.NoError(t, err) + require.NotNil(t, notice) + writeManifest(t, c, "v0.2.0", "") + notice, err = c.check(context.Background(), "0.2.0") + require.NoError(t, err) + assert.Nil(t, notice) +} + +func TestNotification_ExpiredCandidate(t *testing.T) { + c := newTestChecker(t) + now := c.now() + c.now = func() time.Time { return now } + notice, err := c.check(context.Background(), "0.1.1") + require.NoError(t, err) + require.NotNil(t, notice) + now = now.Add(checkInterval) + var out bytes.Buffer + require.NoError(t, notice(&out)) + assert.Empty(t, out.String()) + s, err := readState(c.statePath, c.now()) + require.NoError(t, err) + assert.True(t, s.LastNotifiedAt.IsZero()) +} + +func TestCheck_ReleaseMetadata(t *testing.T) { + tests := []struct { + name string + body string + suffix string + wantNotice bool + wantErr bool + }{ + {"draft", `{"tag_name":"v0.2.0","draft":true,"assets":[{"name":"linux-amd64"}]}`, "linux-amd64", false, false}, + {"prerelease", `{"tag_name":"v0.2.0-rc.1","prerelease":true,"assets":[{"name":"linux-amd64"}]}`, "linux-amd64", false, false}, + {"invalid tag", `{"tag_name":"latest"}`, "linux-amd64", false, true}, + {"prerelease tag", `{"tag_name":"v0.2.0-rc.1"}`, "linux-amd64", false, true}, + {"missing tag", `{}`, "linux-amd64", false, true}, + {"invalid JSON", `{`, "linux-amd64", false, true}, + {"no assets", `{"tag_name":"v0.2.0"}`, "linux-amd64", false, false}, + {"wrong platform", `{"tag_name":"v0.2.0","assets":[{"name":"darwin-arm64"}]}`, "linux-amd64", false, false}, + {"prefixed asset", `{"tag_name":"v0.2.0","assets":[{"name":"gh-stack-linux-amd64"}]}`, "linux-amd64", true, false}, + {"windows executable", `{"tag_name":"v0.2.0","assets":[{"name":"windows-amd64.exe"}]}`, "windows-amd64.exe", true, false}, + {"not a windows executable", `{"tag_name":"v0.2.0","assets":[{"name":"windows-amd64"}]}`, "windows-amd64.exe", false, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c := newTestChecker(t) + c.assetSuffix = tt.suffix + c.client.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(tt.body))}, nil + }) + notice, err := c.check(context.Background(), "0.1.1") + if tt.wantErr { + require.Error(t, err) + } else { + require.NoError(t, err) + } + assert.Equal(t, tt.wantNotice, notice != nil) + s, err := readState(c.statePath, c.now()) + require.NoError(t, err) + assert.Equal(t, c.now(), s.LastCheckedAt) + if !tt.wantNotice { + assert.Empty(t, s.LatestVersion) + } + }) + } +} + +func TestCheck_FailedAttemptsAreThrottled(t *testing.T) { + for _, status := range []int{0, http.StatusUnauthorized, http.StatusForbidden, http.StatusNotFound, http.StatusTooManyRequests, http.StatusInternalServerError} { + t.Run(fmt.Sprint(status), func(t *testing.T) { + c := newTestChecker(t) + require.NoError(t, writeState(c.statePath, state{ + LastCheckedAt: c.now().Add(-checkInterval), + LatestVersion: "v0.1.2", + })) + requests := 0 + c.client.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + requests++ + if status == 0 { + return nil, errors.New("offline") + } + return &http.Response{StatusCode: status, Body: io.NopCloser(strings.NewReader(""))}, nil + }) + notice, err := c.check(context.Background(), "0.1.1") + require.Error(t, err) + assert.Nil(t, notice) + s, err := readState(c.statePath, c.now()) + require.NoError(t, err) + assert.Equal(t, c.now(), s.LastCheckedAt) + assert.Empty(t, s.LatestVersion, "failed checks must not make expired metadata appear fresh") + + notice, err = c.check(context.Background(), "0.1.1") + require.NoError(t, err) + assert.Nil(t, notice) + assert.Equal(t, 1, requests) + }) + } +} + +func TestCheck_CanceledAttemptIsThrottled(t *testing.T) { + c := newTestChecker(t) + started := make(chan struct{}) + c.client.Transport = roundTripFunc(func(req *http.Request) (*http.Response, error) { + close(started) + <-req.Context().Done() + return nil, req.Context().Err() + }) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + done := make(chan error, 1) + go func() { + _, err := c.check(ctx, "0.1.1") + done <- err + }() + select { + case <-started: + case <-time.After(5 * time.Second): + t.Fatal("check did not start") + } + cancel() + select { + case err := <-done: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(5 * time.Second): + t.Fatal("check did not cancel") + } + c.client.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + t.Error("canceled attempt must still be throttled") + return nil, errors.New("unexpected request") + }) + notice, err := c.check(context.Background(), "0.1.1") + require.NoError(t, err) + assert.Nil(t, notice) +} + +func TestCheck_InvalidStateRecovery(t *testing.T) { + for _, data := range []string{ + "last_checked_at: [", + "last_checked_at: invalid\n", + "last_checked_at: 2030-01-01T00:00:00Z\n", + "last_notified_at: 2030-01-01T00:00:00Z\n", + "latest_version: v0.2.0\n", + "last_checked_at: 2026-01-15T12:00:00Z\nlatest_version: garbage\n", + } { + t.Run(data, func(t *testing.T) { + c := newTestChecker(t) + require.NoError(t, os.MkdirAll(filepath.Dir(c.statePath), 0700)) + require.NoError(t, os.WriteFile(c.statePath, []byte(data), 0600)) + notice, err := c.check(context.Background(), "0.1.1") + require.ErrorIs(t, err, errInvalidState) + require.NotNil(t, notice) + require.NoError(t, notice(io.Discard)) + s, err := readState(c.statePath, c.now()) + require.NoError(t, err) + assert.Equal(t, c.now(), s.LastNotifiedAt) + }) + } +} + +func TestCheck_StateErrors(t *testing.T) { + for _, kind := range []string{"unreadable", "unwritable", "relative path"} { + t.Run(kind, func(t *testing.T) { + c := newTestChecker(t) + switch kind { + case "unreadable": + require.NoError(t, os.MkdirAll(c.statePath, 0700)) + case "unwritable": + require.NoError(t, os.WriteFile(filepath.Dir(c.statePath), nil, 0600)) + case "relative path": + c.statePath = "state.yml" + } + c.client.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + t.Error("unusable state must not cause unthrottled network requests") + return nil, errors.New("unexpected request") + }) + notice, err := c.check(context.Background(), "0.1.1") + require.Error(t, err) + assert.Nil(t, notice) + }) + } +} + +func TestNotification_StateFailure(t *testing.T) { + c := newTestChecker(t) + notice, err := c.check(context.Background(), "0.1.1") + require.NoError(t, err) + require.NotNil(t, notice) + require.NoError(t, os.Remove(c.statePath)) + require.NoError(t, os.Mkdir(c.statePath, 0700)) + var out bytes.Buffer + require.Error(t, notice(&out)) + assert.Empty(t, out.String()) +} + +func TestReadState_YAML(t *testing.T) { + c := newTestChecker(t) + require.NoError(t, os.MkdirAll(filepath.Dir(c.statePath), 0700)) + require.NoError(t, os.WriteFile(c.statePath, []byte( + "last_checked_at: 2026-01-15T12:00:00Z\n"+ + "last_notified_at: 2026-01-15T11:00:00Z\n"+ + "latest_version: v0.2.0\n"), 0600)) + + got, err := readState(c.statePath, c.now()) + require.NoError(t, err) + assert.Equal(t, state{ + LastCheckedAt: c.now(), + LastNotifiedAt: c.now().Add(-time.Hour), + LatestVersion: "v0.2.0", + }, got) +} + +func TestWriteState_ReplacesAndCleansUp(t *testing.T) { + c := newTestChecker(t) + for _, version := range []string{"v0.2.0", "v0.3.0"} { + s := state{LastCheckedAt: c.now(), LatestVersion: version} + require.NoError(t, writeState(c.statePath, s)) + data, err := os.ReadFile(c.statePath) + require.NoError(t, err) + assert.Equal(t, fmt.Sprintf( + "last_checked_at: 2026-01-15T12:00:00Z\nlast_notified_at: 0001-01-01T00:00:00Z\nlatest_version: %s\n", + version), string(data)) + got, err := readState(c.statePath, c.now()) + require.NoError(t, err) + assert.Equal(t, s, got) + } + entries, err := os.ReadDir(filepath.Dir(c.statePath)) + require.NoError(t, err) + require.Len(t, entries, 1) + assert.Equal(t, "state.yml", entries[0].Name()) + if runtime.GOOS != "windows" { + info, err := os.Stat(c.statePath) + require.NoError(t, err) + assert.Equal(t, os.FileMode(0600), info.Mode().Perm()) + } +} + +func TestCheck_InstalledBinary(t *testing.T) { + if os.Getenv("GH_STACK_UPDATE_TEST_HELPER") == "1" { + notice, err := Check(context.Background(), "0.1.1") + require.NoError(t, err) + if notice != nil { + require.NoError(t, notice(os.Stdout)) + } + return + } + + t.Setenv("GH_STACK_NO_UPDATE_NOTIFIER", "") + t.Setenv("GH_STACK_UPDATE_TEST_HELPER", "1") + t.Setenv("XDG_STATE_HOME", t.TempDir()) + t.Setenv("GH_HOST", "enterprise.example.com") + t.Setenv("GH_TOKEN", "test-token") + t.Setenv("HTTPS_PROXY", "http://127.0.0.1:0") + t.Setenv("NO_PROXY", "") + + executable, err := os.Executable() + require.NoError(t, err) + binary, err := os.ReadFile(executable) + require.NoError(t, err) + installed := filepath.Join(t.TempDir(), "gh-stack") + if runtime.GOOS == "windows" { + installed += ".exe" + } + require.NoError(t, os.WriteFile(installed, binary, 0700)) + require.NoError(t, os.WriteFile(filepath.Join(filepath.Dir(installed), "manifest.yml"), + []byte("owner: github\nname: gh-stack\nhost: github.com\ntag: v0.1.1\nispinned: false\n"), 0600)) + + statePath := filepath.Join(os.Getenv("XDG_STATE_HOME"), "gh", "gh-stack", "state.yml") + require.NoError(t, writeState(statePath, state{LastCheckedAt: time.Now(), LatestVersion: "v0.2.0"})) + + for _, wantNotice := range []bool{true, false} { + command := exec.Command(installed, "-test.run=^TestCheck_InstalledBinary$") + output, err := command.CombinedOutput() + require.NoError(t, err, "%s", output) + assert.Equal(t, wantNotice, strings.Contains(string(output), "gh extension upgrade stack"), "%s", output) + } + s, err := readState(statePath, time.Now()) + require.NoError(t, err) + assert.False(t, s.LastNotifiedAt.IsZero()) +} From c8e7b43fe17ff19eb8061e40122d917b40990d14 Mon Sep 17 00:00:00 2001 From: Sameen Karim Date: Fri, 2 Oct 2026 15:38:42 -0400 Subject: [PATCH 2/3] Show upgrade notices after command failures --- README.md | 9 ++- cmd/checkout.go | 2 + cmd/merge.go | 2 + cmd/modify.go | 7 +-- cmd/root.go | 40 +++++++++---- cmd/root_test.go | 82 ++++++++++++++++++++++---- cmd/submit.go | 3 + cmd/utils.go | 2 + cmd/utils_test.go | 1 + docs/src/content/docs/reference/cli.md | 7 ++- internal/config/config.go | 3 + internal/update/update.go | 2 +- 12 files changed, 124 insertions(+), 36 deletions(-) diff --git a/README.md b/README.md index 4faab13d..9b174659 100644 --- a/README.md +++ b/README.md @@ -20,9 +20,11 @@ gh extension upgrade stack Official, unpinned stable release installations check for a newer [latest release](https://github.com/github/gh-stack/releases/latest) in the -background, at most once every 24 hours. When an update is available, successful -commands can append an upgrade notice to stderr, also at most once every 24 -hours. This includes non-interactive use; stdout and JSON output are unchanged. +background, at most once every 24 hours. When an update is available, commands +can append an upgrade notice to stderr after success or an operational failure, +also at most once every 24 hours. On failure, the original error appears first +and the exit code is unchanged. This includes non-interactive use; stdout and +JSON output are unchanged. The check and reminder timestamps are shared across repositories for your user. They are stored as YAML in `gh-stack/state.yml` beneath GitHub CLI's state directory (`~/.local/state/gh` by default on macOS/Linux). @@ -33,6 +35,7 @@ check limit. Completed checks are cached so a later command can show the notice. No upgrades happen automatically, and development, locally linked, prerelease, and pinned installations are excluded. Help, version, and completion commands do not run the notifier. +Usage errors and explicit user cancellations do not show upgrade notices. Set `GH_STACK_NO_UPDATE_NOTIFIER=1` to disable these checks and notices, including in CI or scripts. Any non-empty value disables the notifier. Optional check diff --git a/cmd/checkout.go b/cmd/checkout.go index ba926dc3..9a1df2e3 100644 --- a/cmd/checkout.go +++ b/cmd/checkout.go @@ -660,6 +660,7 @@ func handleCompositionConflict( default: // Cancel + cfg.Canceled = true cfg.Infof("Checkout cancelled") return nil, ErrSilent } @@ -805,6 +806,7 @@ func interactiveCheckout(cfg *config.Config, sf *stack.StackFile, gitDir string) } if !ok { // The user dismissed the picker without selecting. + cfg.Canceled = true return nil, "", nil } diff --git a/cmd/merge.go b/cmd/merge.go index e8e6741b..32ed9320 100644 --- a/cmd/merge.go +++ b/cmd/merge.go @@ -364,10 +364,12 @@ func runMergeInteractive(cfg *config.Config, client github.ClientOps, stackNumbe cfg.Printf("Stack merges are atomic, so nothing was merged.") return mergeFailureExit(out.Message) case out.WatchStopped: + cfg.Canceled = true cfg.Infof("Stopped watching. Merge is still in progress. Check the pull requests on GitHub.") return ErrSilent default: // Cancelled via esc/ctrl+c before submitting. + cfg.Canceled = true cfg.Infof("Cancelled operation, nothing merged") return ErrSilent } diff --git a/cmd/modify.go b/cmd/modify.go index 5f3f7cbd..c4b6d53e 100644 --- a/cmd/modify.go +++ b/cmd/modify.go @@ -134,11 +134,8 @@ func runModify(cfg *config.Config) error { } // Handle TUI result - if m.Cancelled() { - return nil - } - - if !m.ApplyRequested() { + if m.Cancelled() || !m.ApplyRequested() { + cfg.Canceled = true return nil } diff --git a/cmd/root.go b/cmd/root.go index 98bb3498..489d566e 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -17,11 +17,16 @@ type updateResult struct { err error } +type rootCommand struct { + *cobra.Command + finish func(error) +} + func RootCmd() *cobra.Command { - return newRootCmd(config.New(), startUpdateCheck) + return newRootCmd(config.New(), startUpdateCheck).Command } -func newRootCmd(cfg *config.Config, startCheck func(context.Context, string) <-chan updateResult) *cobra.Command { +func newRootCmd(cfg *config.Config, startCheck func(context.Context, string) <-chan updateResult) *rootCommand { root := &cobra.Command{ Use: "stack ", Short: "Manage stacked branches and pull requests", @@ -50,22 +55,26 @@ locally, then push to GitHub to create your stack of PRs.`, var updates <-chan updateResult root.PersistentPreRun = func(cmd *cobra.Command, _ []string) { theme.ApplyOverride() + cfg.Canceled = false updates = nil if update.Enabled(root.Version) && !isHelpOrCompletionCommand(cmd) { updates = startCheck(cmd.Context(), root.Version) } } - root.PersistentPostRun = func(cmd *cobra.Command, _ []string) { - if cmd.Context().Err() != nil { + finish := func(err error) { + results := updates + updates = nil + if cfg.Canceled || errors.Is(err, ErrInvalidArgs) || + errors.Is(err, context.Canceled) || errors.Is(err, errInterrupt) || isInterruptError(err) { return } select { - case result := <-updates: + case result := <-results: if result.notify != nil { - result.err = errors.Join(result.err, result.notify(cmd.ErrOrStderr())) + result.err = errors.Join(result.err, result.notify(root.ErrOrStderr())) } if result.err != nil && os.Getenv("GH_DEBUG") != "" { - fmt.Fprintf(cmd.ErrOrStderr(), "debug: gh-stack update notification: %v\n", result.err) + fmt.Fprintf(root.ErrOrStderr(), "debug: gh-stack update notification: %v\n", result.err) } default: // Never wait for a release check if the command finishes quickly @@ -180,7 +189,7 @@ locally, then push to GitHub to create your stack of PRs.`, feedbackCmd.GroupID = "utils" root.AddCommand(feedbackCmd) - return root + return &rootCommand{Command: root, finish: finish} } func startUpdateCheck(ctx context.Context, version string) <-chan updateResult { @@ -202,25 +211,30 @@ func isHelpOrCompletionCommand(cmd *cobra.Command) bool { return false } -func execute(cmd *cobra.Command, args []string) error { +func execute(cmd *rootCommand, args []string) error { ctx, cancel := context.WithCancel(context.Background()) defer cancel() // Wrap in a "gh" parent so help output shows "gh stack" instead of just "stack". wrapCmd := &cobra.Command{Use: "gh", SilenceUsage: true, SilenceErrors: true} - wrapCmd.AddCommand(cmd) + wrapCmd.AddCommand(cmd.Command) wrapCmd.SetArgs(append([]string{"stack"}, args...)) - return wrapCmd.ExecuteContext(ctx) + err := wrapCmd.ExecuteContext(ctx) + var exitErr *ExitError + if err != nil && !errors.As(err, &exitErr) { + fmt.Fprintln(cmd.ErrOrStderr(), err) + } + cmd.finish(err) + return err } func Execute() { - cmd := RootCmd() + cmd := newRootCmd(config.New(), startUpdateCheck) if err := execute(cmd, os.Args[1:]); err != nil { var exitErr *ExitError if errors.As(err, &exitErr) { os.Exit(exitErr.Code) } - fmt.Fprintln(cmd.ErrOrStderr(), err) os.Exit(1) } } diff --git a/cmd/root_test.go b/cmd/root_test.go index 18416872..40c2a559 100644 --- a/cmd/root_test.go +++ b/cmd/root_test.go @@ -9,6 +9,7 @@ import ( "testing" "time" + "github.com/AlecAivazis/survey/v2/terminal" "github.com/github/gh-stack/internal/config" "github.com/spf13/cobra" "github.com/stretchr/testify/assert" @@ -47,7 +48,7 @@ func TestRootCmd_HelpOutput(t *testing.T) { assert.Contains(t, output, "https://gh.io/stacks") } -func newUpdateTestRoot(t *testing.T, startCheck func(context.Context, string) <-chan updateResult) (*cobra.Command, *bytes.Buffer, *bytes.Buffer) { +func newUpdateTestRoot(t *testing.T, startCheck func(context.Context, string) <-chan updateResult) (*rootCommand, *bytes.Buffer, *bytes.Buffer) { t.Helper() t.Setenv("GH_STACK_NO_UPDATE_NOTIFIER", "") t.Setenv("GH_DEBUG", "") @@ -162,14 +163,23 @@ func TestRootCmd_UpdateNoticeExcludedCommands(t *testing.T) { } func TestRootCmd_UpdateNoticePreservesCommandErrors(t *testing.T) { - for _, commandErr := range []error{ - errors.New("command failed"), - ErrConflict, - fmt.Errorf("wrapped: %w", ErrInvalidArgs), - ErrSilent, - context.Canceled, + for _, tt := range []struct { + name string + commandErr error + errorText string + wantNotice bool + }{ + {"untyped failure", errors.New("command failed"), "command failed\n", true}, + {"conflict", ErrConflict, "", true}, + {"API failure", ErrAPIFailure, "", true}, + {"wrapped failure", fmt.Errorf("wrapped: %w", ErrAPIFailure), "", true}, + {"already reported failure", ErrSilent, "", true}, + {"usage error", fmt.Errorf("wrapped: %w", ErrInvalidArgs), "", false}, + {"canceled context", context.Canceled, "context canceled\n", false}, + {"interrupt", errInterrupt, "interrupt\n", false}, + {"terminal interrupt", terminal.InterruptErr, terminal.InterruptErr.Error() + "\n", false}, } { - t.Run(commandErr.Error(), func(t *testing.T) { + t.Run(tt.name, func(t *testing.T) { var checkContext context.Context root, stdout, stderr := newUpdateTestRoot(t, func(ctx context.Context, _ string) <-chan updateResult { checkContext = ctx @@ -179,18 +189,66 @@ func TestRootCmd_UpdateNoticePreservesCommandErrors(t *testing.T) { Use: "probe", RunE: func(cmd *cobra.Command, _ []string) error { fmt.Fprintln(cmd.ErrOrStderr(), "Original diagnostic.") - return commandErr + return tt.commandErr }, }) err := execute(root, []string{"probe"}) - require.ErrorIs(t, err, commandErr) + require.ErrorIs(t, err, tt.commandErr) assert.Empty(t, stdout.String()) - assert.Equal(t, "Original diagnostic.\n", stderr.String()) + var want bytes.Buffer + fmt.Fprint(&want, "Original diagnostic.\n"+tt.errorText) + if tt.wantNotice { + require.NoError(t, testUpdateNotification(&want)) + } + assert.Equal(t, want.String(), stderr.String()) require.ErrorIs(t, checkContext.Err(), context.Canceled) }) } } +func TestRootCmd_CanceledCommandSuppressesNotice(t *testing.T) { + for _, tt := range []struct { + name string + interrupt bool + err error + }{ + {"successful cancellation", false, nil}, + {"silent cancellation", false, ErrSilent}, + {"reported interrupt", true, ErrSilent}, + } { + t.Run(tt.name, func(t *testing.T) { + t.Setenv("GH_STACK_NO_UPDATE_NOTIFIER", "") + cfg, outR, errR := config.NewTestConfig() + root := newRootCmd(cfg, func(context.Context, string) <-chan updateResult { + return readyUpdate(updateResult{notify: testUpdateNotification}) + }) + root.Version = "0.1.1" + root.AddCommand(&cobra.Command{ + Use: "probe", + RunE: func(*cobra.Command, []string) error { + if tt.interrupt { + printInterrupt(cfg) + } else { + cfg.Canceled = true + } + return tt.err + }, + }) + err := execute(root, []string{"probe"}) + if tt.err == nil { + require.NoError(t, err) + } else { + require.ErrorIs(t, err, tt.err) + } + output := collectOutput(cfg, outR, errR) + assert.NotContains(t, output, "gh extension upgrade stack") + if tt.interrupt { + assert.Contains(t, output, "Received interrupt, aborting operation") + } + }) + } +} + func TestRootCmd_UpdateDiagnostics(t *testing.T) { tests := []struct { name string @@ -258,7 +316,7 @@ func TestRootCmd_DoesNotWaitForUpdate(t *testing.T) { require.NoError(t, err) } case <-time.After(5 * time.Second): - // Unblock a regressed post-run hook before failing the test. + // Unblock a regressed finalizer before failing the test. results <- updateResult{} <-done <-workerDone diff --git a/cmd/submit.go b/cmd/submit.go index 5d2f1ce4..f310c00f 100644 --- a/cmd/submit.go +++ b/cmd/submit.go @@ -133,6 +133,7 @@ func runSubmit(cfg *config.Config, opts *submitOptions) error { return ErrStacksUnavailable } if !proceed { + cfg.Canceled = true return ErrStacksUnavailable } } else { @@ -222,6 +223,7 @@ func runSubmit(cfg *config.Config, opts *submitOptions) error { return ErrSilent } if cancelled { + cfg.Canceled = true cfg.Printf("Submit cancelled — no branches were pushed") return nil } @@ -710,6 +712,7 @@ func handlePendingModify(cfg *config.Config, client github.ClientOps, s *stack.S return true, promptErr } if !proceed { + cfg.Canceled = true cfg.Printf("Skipping stack recreation — run `%s` when ready", cfg.ColorCyan("gh stack submit")) return true, errInterrupt diff --git a/cmd/utils.go b/cmd/utils.go index 1f963928..4dec9a45 100644 --- a/cmd/utils.go +++ b/cmd/utils.go @@ -69,6 +69,7 @@ func isInterruptError(err error) bool { // per interrupted operation. The leading newline ensures the message starts // on its own line even if the cursor was mid-prompt. func printInterrupt(cfg *config.Config) { + cfg.Canceled = true fmt.Fprintln(cfg.Err) cfg.Infof("Received interrupt, aborting operation") } @@ -2029,6 +2030,7 @@ func resolveStackDivergence(cfg *config.Config, client github.ClientOps, sf *sta return resolveDivergenceDeleteRemote(cfg, client, sf, s, gitDir) default: // Cancel: stop the sync without touching branches or PRs. + cfg.Canceled = true cfg.Infof("Sync aborted — no changes were made") return remoteReconcileResult{stop: true}, nil } diff --git a/cmd/utils_test.go b/cmd/utils_test.go index 475e69cb..304e27f7 100644 --- a/cmd/utils_test.go +++ b/cmd/utils_test.go @@ -1029,6 +1029,7 @@ func TestPrintInterrupt_Output(t *testing.T) { printInterrupt(cfg) output := collectOutput(cfg, outR, errR) + assert.True(t, cfg.Canceled) if !strings.Contains(output, "Received interrupt, aborting operation") { t.Errorf("expected interrupt message, got: %s", output) } diff --git a/docs/src/content/docs/reference/cli.md b/docs/src/content/docs/reference/cli.md index a07929e2..05c1eeb8 100644 --- a/docs/src/content/docs/reference/cli.md +++ b/docs/src/content/docs/reference/cli.md @@ -31,8 +31,10 @@ gh extension upgrade stack Official, unpinned stable release installations check the [latest gh-stack release](https://github.com/github/gh-stack/releases/latest) -in the background, at most once every 24 hours. Successful commands can append -an upgrade notice to stderr, also at most once every 24 hours, until you upgrade. +in the background, at most once every 24 hours. Commands can append an upgrade +notice to stderr after success or an operational failure, also at most once every +24 hours, until you upgrade. On failure, the original error appears first and the +exit code is unchanged. The timestamps are stored per user and shared across repositories, independently of GitHub CLI's own update checks. The YAML state file is `gh-stack/state.yml` beneath GitHub CLI's state directory: @@ -49,6 +51,7 @@ Development builds, locally linked installations, prereleases, and pinned installations are excluded. Help, version, and shell completion commands do not run the notifier. Users must first install a release containing this feature to receive notices about subsequent releases. +Usage errors and explicit user cancellations do not show upgrade notices. To disable the notifier, including its network checks: diff --git a/internal/config/config.go b/internal/config/config.go index c924f889..686e7eb5 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -40,6 +40,9 @@ type Config struct { // NonInteractive suppresses prompts even when stdout is a terminal. NonInteractive bool + // Canceled records explicit user cancellation even when a command returns ErrSilent. + Canceled bool + // WorktreePathOnly makes checkout resolution skip imports for foreign owners. WorktreePathOnly bool diff --git a/internal/update/update.go b/internal/update/update.go index ccbb34f7..d78da27e 100644 --- a/internal/update/update.go +++ b/internal/update/update.go @@ -25,7 +25,7 @@ const ( latestReleaseURL = "https://api.github.com/repos/github/gh-stack/releases/latest" ) -// Notification records and writes a pending notice when a command succeeds. +// Notification records and writes a pending notice after a command finishes. type Notification func(io.Writer) error // Enabled excludes opt-outs and builds that cannot be compared to stable releases. From 80dfd1f44c787d4779abd4dbbababa32366f9ac8 Mon Sep 17 00:00:00 2001 From: Sameen Karim Date: Fri, 2 Oct 2026 15:43:39 -0400 Subject: [PATCH 3/3] Retry interrupted update checks after a short cooldown --- README.md | 9 ++-- docs/src/content/docs/reference/cli.md | 8 +-- internal/update/update.go | 22 ++++---- internal/update/update_test.go | 70 +++++++++++++++++++++----- 4 files changed, 80 insertions(+), 29 deletions(-) diff --git a/README.md b/README.md index 9b174659..77020b69 100644 --- a/README.md +++ b/README.md @@ -20,8 +20,8 @@ gh extension upgrade stack Official, unpinned stable release installations check for a newer [latest release](https://github.com/github/gh-stack/releases/latest) in the -background, at most once every 24 hours. When an update is available, commands -can append an upgrade notice to stderr after success or an operational failure, +background, caching successful checks for 24 hours. When an update is available, +commands can append an upgrade notice to stderr after success or an operational failure, also at most once every 24 hours. On failure, the original error appears first and the exit code is unchanged. This includes non-interactive use; stdout and JSON output are unchanged. @@ -30,8 +30,9 @@ They are stored as YAML in `gh-stack/state.yml` beneath GitHub CLI's state directory (`~/.local/state/gh` by default on macOS/Linux). Commands never wait for a release check. Short commands may finish before a -notice is ready, and failed or interrupted checks still count toward the daily -check limit. Completed checks are cached so a later command can show the notice. +notice is ready. Failed or interrupted attempts can retry after a 15-minute +cooldown instead of waiting a full day. Completed checks are cached so a later +command can show the notice. No upgrades happen automatically, and development, locally linked, prerelease, and pinned installations are excluded. Help, version, and completion commands do not run the notifier. diff --git a/docs/src/content/docs/reference/cli.md b/docs/src/content/docs/reference/cli.md index 05c1eeb8..e96b3136 100644 --- a/docs/src/content/docs/reference/cli.md +++ b/docs/src/content/docs/reference/cli.md @@ -31,8 +31,8 @@ gh extension upgrade stack Official, unpinned stable release installations check the [latest gh-stack release](https://github.com/github/gh-stack/releases/latest) -in the background, at most once every 24 hours. Commands can append an upgrade -notice to stderr after success or an operational failure, also at most once every +in the background, caching successful checks for 24 hours. Commands can append an +upgrade notice to stderr after success or an operational failure, also at most once every 24 hours, until you upgrade. On failure, the original error appears first and the exit code is unchanged. The timestamps are stored per user and shared across repositories, independently @@ -44,8 +44,8 @@ The YAML state file is `gh-stack/state.yml` beneath GitHub CLI's state directory Notices also appear in non-interactive use, including CI. Stdout and JSON output remain unchanged. No upgrade happens automatically, and commands never wait for the check: a short command may finish before a notice is ready. Completed checks -are cached for later commands; failed or interrupted attempts still count toward -the daily check limit. +are cached for later commands; failed or interrupted attempts can retry after a +15-minute cooldown instead of waiting a full day. Development builds, locally linked installations, prereleases, and pinned installations are excluded. Help, version, and shell completion commands do not diff --git a/internal/update/update.go b/internal/update/update.go index d78da27e..3b41a1e5 100644 --- a/internal/update/update.go +++ b/internal/update/update.go @@ -21,6 +21,7 @@ import ( const ( checkInterval = 24 * time.Hour + retryInterval = 15 * time.Minute requestTimeout = 2 * time.Second latestReleaseURL = "https://api.github.com/repos/github/gh-stack/releases/latest" ) @@ -68,9 +69,10 @@ type checker struct { } type state struct { - LastCheckedAt time.Time `yaml:"last_checked_at"` - LastNotifiedAt time.Time `yaml:"last_notified_at"` - LatestVersion string `yaml:"latest_version,omitempty"` + LastAttemptedAt time.Time `yaml:"last_attempted_at,omitempty"` + LastCheckedAt time.Time `yaml:"last_checked_at"` + LastNotifiedAt time.Time `yaml:"last_notified_at"` + LatestVersion string `yaml:"latest_version,omitempty"` } var errInvalidState = errors.New("invalid update notification state") @@ -106,11 +108,11 @@ func (c checker) check(ctx context.Context, version string) (Notification, error return nil, err } - if s.LastCheckedAt.IsZero() || now.Sub(s.LastCheckedAt) >= checkInterval { - // Reserve the daily attempt before the request, including attempts - // interrupted when a short command exits. - s.LastCheckedAt = now - s.LatestVersion = "" + if (s.LastCheckedAt.IsZero() || now.Sub(s.LastCheckedAt) >= checkInterval) && + (s.LastAttemptedAt.IsZero() || now.Sub(s.LastAttemptedAt) >= retryInterval) { + // Throttle interrupted attempts without refreshing or discarding + // the last successful result. + s.LastAttemptedAt = now if err := writeState(c.statePath, s); err != nil { return nil, errors.Join(diagnostic, err) } @@ -118,6 +120,8 @@ func (c checker) check(ctx context.Context, version string) (Notification, error if err != nil { return nil, errors.Join(diagnostic, err) } + now = c.now() + s.LastCheckedAt = now s.LatestVersion = latest if err := writeState(c.statePath, s); err != nil { return nil, errors.Join(diagnostic, err) @@ -254,7 +258,7 @@ func readState(path string, now time.Time) (state, error) { if err := yaml.Unmarshal(data, &s); err != nil { return state{}, fmt.Errorf("%w: %w", errInvalidState, err) } - if s.LastCheckedAt.After(now) || s.LastNotifiedAt.After(now) { + if s.LastAttemptedAt.After(now) || s.LastCheckedAt.After(now) || s.LastNotifiedAt.After(now) { return state{}, fmt.Errorf("%w: timestamp is in the future", errInvalidState) } if s.LatestVersion != "" && (stableVersion(s.LatestVersion) != s.LatestVersion || s.LastCheckedAt.IsZero()) { diff --git a/internal/update/update_test.go b/internal/update/update_test.go index 7aa72ab2..e4bb2274 100644 --- a/internal/update/update_test.go +++ b/internal/update/update_test.go @@ -370,7 +370,12 @@ func TestCheck_ReleaseMetadata(t *testing.T) { assert.Equal(t, tt.wantNotice, notice != nil) s, err := readState(c.statePath, c.now()) require.NoError(t, err) - assert.Equal(t, c.now(), s.LastCheckedAt) + assert.Equal(t, c.now(), s.LastAttemptedAt) + if tt.wantErr { + assert.True(t, s.LastCheckedAt.IsZero()) + } else { + assert.Equal(t, c.now(), s.LastCheckedAt) + } if !tt.wantNotice { assert.Empty(t, s.LatestVersion) } @@ -382,8 +387,11 @@ func TestCheck_FailedAttemptsAreThrottled(t *testing.T) { for _, status := range []int{0, http.StatusUnauthorized, http.StatusForbidden, http.StatusNotFound, http.StatusTooManyRequests, http.StatusInternalServerError} { t.Run(fmt.Sprint(status), func(t *testing.T) { c := newTestChecker(t) + now := c.now() + start := now + c.now = func() time.Time { return now } require.NoError(t, writeState(c.statePath, state{ - LastCheckedAt: c.now().Add(-checkInterval), + LastCheckedAt: start.Add(-checkInterval), LatestVersion: "v0.1.2", })) requests := 0 @@ -399,19 +407,42 @@ func TestCheck_FailedAttemptsAreThrottled(t *testing.T) { assert.Nil(t, notice) s, err := readState(c.statePath, c.now()) require.NoError(t, err) - assert.Equal(t, c.now(), s.LastCheckedAt) - assert.Empty(t, s.LatestVersion, "failed checks must not make expired metadata appear fresh") + assert.Equal(t, start, s.LastAttemptedAt) + assert.Equal(t, start.Add(-checkInterval), s.LastCheckedAt) + assert.Equal(t, "v0.1.2", s.LatestVersion) + for _, elapsed := range []time.Duration{0, retryInterval - time.Nanosecond} { + now = start.Add(elapsed) + notice, err = c.check(context.Background(), "0.1.1") + require.NoError(t, err) + assert.Nil(t, notice, "expired metadata must not be delivered while waiting to retry") + assert.Equal(t, 1, requests) + } + + now = start.Add(retryInterval) + c.client.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + requests++ + now = now.Add(time.Second) + return releaseResponse("v0.2.0"), nil + }) notice, err = c.check(context.Background(), "0.1.1") require.NoError(t, err) - assert.Nil(t, notice) - assert.Equal(t, 1, requests) + require.NotNil(t, notice) + assert.Equal(t, 2, requests) + s, err = readState(c.statePath, now) + require.NoError(t, err) + assert.Equal(t, start.Add(retryInterval), s.LastAttemptedAt) + assert.Equal(t, now, s.LastCheckedAt, "record completion, not the start of the request") + assert.Equal(t, "v0.2.0", s.LatestVersion) + require.NoError(t, notice(io.Discard)) }) } } func TestCheck_CanceledAttemptIsThrottled(t *testing.T) { c := newTestChecker(t) + now := c.now() + c.now = func() time.Time { return now } started := make(chan struct{}) c.client.Transport = roundTripFunc(func(req *http.Request) (*http.Response, error) { close(started) @@ -444,12 +475,25 @@ func TestCheck_CanceledAttemptIsThrottled(t *testing.T) { notice, err := c.check(context.Background(), "0.1.1") require.NoError(t, err) assert.Nil(t, notice) + s, err := readState(c.statePath, now) + require.NoError(t, err) + assert.Equal(t, now, s.LastAttemptedAt) + assert.True(t, s.LastCheckedAt.IsZero()) + + now = now.Add(retryInterval) + c.client.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + return releaseResponse("v0.2.0"), nil + }) + notice, err = c.check(context.Background(), "0.1.1") + require.NoError(t, err) + require.NotNil(t, notice) } func TestCheck_InvalidStateRecovery(t *testing.T) { for _, data := range []string{ "last_checked_at: [", "last_checked_at: invalid\n", + "last_attempted_at: 2030-01-01T00:00:00Z\n", "last_checked_at: 2030-01-01T00:00:00Z\n", "last_notified_at: 2030-01-01T00:00:00Z\n", "latest_version: v0.2.0\n", @@ -509,28 +553,30 @@ func TestReadState_YAML(t *testing.T) { c := newTestChecker(t) require.NoError(t, os.MkdirAll(filepath.Dir(c.statePath), 0700)) require.NoError(t, os.WriteFile(c.statePath, []byte( - "last_checked_at: 2026-01-15T12:00:00Z\n"+ + "last_attempted_at: 2026-01-15T11:59:00Z\n"+ + "last_checked_at: 2026-01-15T12:00:00Z\n"+ "last_notified_at: 2026-01-15T11:00:00Z\n"+ "latest_version: v0.2.0\n"), 0600)) got, err := readState(c.statePath, c.now()) require.NoError(t, err) assert.Equal(t, state{ - LastCheckedAt: c.now(), - LastNotifiedAt: c.now().Add(-time.Hour), - LatestVersion: "v0.2.0", + LastAttemptedAt: c.now().Add(-time.Minute), + LastCheckedAt: c.now(), + LastNotifiedAt: c.now().Add(-time.Hour), + LatestVersion: "v0.2.0", }, got) } func TestWriteState_ReplacesAndCleansUp(t *testing.T) { c := newTestChecker(t) for _, version := range []string{"v0.2.0", "v0.3.0"} { - s := state{LastCheckedAt: c.now(), LatestVersion: version} + s := state{LastAttemptedAt: c.now(), LastCheckedAt: c.now(), LatestVersion: version} require.NoError(t, writeState(c.statePath, s)) data, err := os.ReadFile(c.statePath) require.NoError(t, err) assert.Equal(t, fmt.Sprintf( - "last_checked_at: 2026-01-15T12:00:00Z\nlast_notified_at: 0001-01-01T00:00:00Z\nlatest_version: %s\n", + "last_attempted_at: 2026-01-15T12:00:00Z\nlast_checked_at: 2026-01-15T12:00:00Z\nlast_notified_at: 0001-01-01T00:00:00Z\nlatest_version: %s\n", version), string(data)) got, err := readState(c.statePath, c.now()) require.NoError(t, err)