diff --git a/README.md b/README.md index 893c068..77020b6 100644 --- a/README.md +++ b/README.md @@ -12,6 +12,37 @@ 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, 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. +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. 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. +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 +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/checkout.go b/cmd/checkout.go index ba926dc..9a1df2e 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 e8e6741..32ed932 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 5f3f7cb..c4b6d53 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 f4100c7..489d566 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -1,18 +1,32 @@ 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 +} + +type rootCommand struct { + *cobra.Command + finish func(error) +} + func RootCmd() *cobra.Command { - cfg := config.New() + return newRootCmd(config.New(), startUpdateCheck).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", @@ -36,10 +50,35 @@ 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() + cfg.Canceled = false + updates = nil + if update.Enabled(root.Version) && !isHelpOrCompletionCommand(cmd) { + updates = startCheck(cmd.Context(), root.Version) + } + } + 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 := <-results: + if result.notify != nil { + result.err = errors.Join(result.err, result.notify(root.ErrOrStderr())) + } + if result.err != nil && os.Getenv("GH_DEBUG") != "" { + 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 + } } root.SetVersionTemplate("gh stack version {{.Version}}\n") @@ -150,23 +189,52 @@ 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 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 *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.SetArgs(append([]string{"stack"}, os.Args[1:]...)) + wrapCmd.AddCommand(cmd.Command) + wrapCmd.SetArgs(append([]string{"stack"}, args...)) + err := wrapCmd.ExecuteContext(ctx) + var exitErr *ExitError + if err != nil && !errors.As(err, &exitErr) { + fmt.Fprintln(cmd.ErrOrStderr(), err) + } + cmd.finish(err) + return err +} - if err := wrapCmd.Execute(); err != nil { +func Execute() { + 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 47328fd..40c2a55 100644 --- a/cmd/root_test.go +++ b/cmd/root_test.go @@ -2,8 +2,16 @@ package cmd import ( "bytes" + "context" + "errors" + "fmt" + "io" "testing" + "time" + "github.com/AlecAivazis/survey/v2/terminal" + "github.com/github/gh-stack/internal/config" + "github.com/spf13/cobra" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -39,3 +47,301 @@ 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) (*rootCommand, *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 _, 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(tt.name, 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 tt.commandErr + }, + }) + err := execute(root, []string{"probe"}) + require.ErrorIs(t, err, tt.commandErr) + assert.Empty(t, stdout.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 + 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 finalizer 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/cmd/submit.go b/cmd/submit.go index 5d2f1ce..f310c00 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 1f96392..4dec9a4 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 475e69c..304e27f 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 e31b530..e96b313 100644 --- a/docs/src/content/docs/reference/cli.md +++ b/docs/src/content/docs/reference/cli.md @@ -23,6 +23,45 @@ 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, 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 +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 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 +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: + +```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 +734,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 eed0bc5..05c0a98 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 e0970f3..8d2ac16 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/config/config.go b/internal/config/config.go index c924f88..686e7eb 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 new file mode 100644 index 0000000..3b41a1e --- /dev/null +++ b/internal/update/update.go @@ -0,0 +1,294 @@ +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 + retryInterval = 15 * time.Minute + requestTimeout = 2 * time.Second + latestReleaseURL = "https://api.github.com/repos/github/gh-stack/releases/latest" +) + +// 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. +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 { + 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") + +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) && + (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) + } + latest, err := c.fetchLatest(ctx, current) + 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) + } + } + + 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.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()) { + 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 0000000..e4bb227 --- /dev/null +++ b/internal/update/update_test.go @@ -0,0 +1,638 @@ +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.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) + } + }) + } +} + +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: start.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, 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) + 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) + <-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) + 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", + "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_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{ + 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{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_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) + 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()) +}