diff --git a/pkg/commands/git.go b/pkg/commands/git.go index 53701a04e..5904dfe4d 100644 --- a/pkg/commands/git.go +++ b/pkg/commands/git.go @@ -145,7 +145,7 @@ func NewGitCommandAux( gitHubCommands := git_commands.NewGitHubCommands(gitCommon) hostingServiceCommands := git_commands.NewHostingServiceCommand(gitCommon) - branchLoader := git_commands.NewBranchLoader(cmn, gitCommon, cmd, branchCommands.CurrentBranchInfo, configCommands) + branchLoader := git_commands.NewBranchLoader(cmn, gitCommon, cmd, branchCommands.CurrentBranchInfo, branchCommands.HasLocalOnlyCommits, configCommands) commitFileLoader := git_commands.NewCommitFileLoader(cmn, cmd) commitLoader := git_commands.NewCommitLoader(cmn, cmd, statusCommands.WorkingTreeState, gitCommon) reflogCommitLoader := git_commands.NewReflogCommitLoader(cmn, cmd) diff --git a/pkg/commands/git_commands/branch_loader.go b/pkg/commands/git_commands/branch_loader.go index 10d9ad30c..6091a4679 100644 --- a/pkg/commands/git_commands/branch_loader.go +++ b/pkg/commands/git_commands/branch_loader.go @@ -44,6 +44,7 @@ type BranchLoader struct { *GitCommon cmd oscommands.ICmdObjBuilder getCurrentBranchInfo func() (BranchInfo, error) + hasLocalOnlyCommits func(*models.Branch) (bool, error) config BranchLoaderConfigCommands } @@ -52,6 +53,7 @@ func NewBranchLoader( gitCommon *GitCommon, cmd oscommands.ICmdObjBuilder, getCurrentBranchInfo func() (BranchInfo, error), + hasLocalOnlyCommits func(*models.Branch) (bool, error), config BranchLoaderConfigCommands, ) *BranchLoader { return &BranchLoader{ @@ -59,6 +61,7 @@ func NewBranchLoader( GitCommon: gitCommon, cmd: cmd, getCurrentBranchInfo: getCurrentBranchInfo, + hasLocalOnlyCommits: hasLocalOnlyCommits, config: config, } } @@ -135,24 +138,63 @@ func (self *BranchLoader) Load(reflogCommits []*models.Commit, branch.UpstreamBranch = match.Merge } - // If the branch already existed, take over its BehindBaseBranch value - // to reduce flicker + // If the branch already existed, take over the values that are + // determined in the background, to reduce flicker if oldBranch, found := lo.Find(oldBranches, func(b *models.Branch) bool { return b.Name == branch.Name }); found { branch.BehindBaseBranch.Store(oldBranch.BehindBaseBranch.Load()) + branch.UpstreamRewritten.Store(oldBranch.UpstreamRewritten.Load()) } } - if loadExtraInfo && self.UserConfig().Gui.ShowDivergenceFromBaseBranch != "none" { + if loadExtraInfo { + if self.UserConfig().Gui.ShowDivergenceFromBaseBranch != "none" { + onWorker(func() error { + return self.GetBehindBaseBranchValuesForAllBranches(branches, mainBranches, renderFunc) + }) + } + onWorker(func() error { - return self.GetBehindBaseBranchValuesForAllBranches(branches, mainBranches, renderFunc) + return self.checkForRewrittenUpstreams(branches, renderFunc) }) } return branches, nil } +// For each branch that has diverged from its upstream, determines whether the +// divergence comes from the upstream branch having been rewritten, and stores +// the answer in the branch. A branch that we can't determine it for keeps the +// answer "no", so that we don't offer anything we aren't sure about. +func (self *BranchLoader) checkForRewrittenUpstreams(branches []*models.Branch, renderFunc func()) error { + t := time.Now() + errg := errgroup.Group{} + + for _, branch := range branches { + if !branch.IsAheadForPull() || !branch.IsBehindForPull() { + branch.UpstreamRewritten.Store(false) + continue + } + + errg.Go(func() error { + hasLocalOnlyCommits, err := self.hasLocalOnlyCommits(branch) + if err != nil { + // Not worth bothering the user about; it only means that we + // don't show this branch differently. + self.Log.Errorf("Failed to check whether branch %s has commits of its own: %v", branch.Name, err) + } + branch.UpstreamRewritten.Store(err == nil && !hasLocalOnlyCommits) + return nil + }) + } + + err := errg.Wait() + self.Log.Debugf("time to check for rewritten upstreams for all branches: %s", time.Since(t)) + renderFunc() + return err +} + func (self *BranchLoader) GetBehindBaseBranchValuesForAllBranches( branches []*models.Branch, mainBranches *MainBranches, diff --git a/pkg/commands/git_commands/branch_loader_test.go b/pkg/commands/git_commands/branch_loader_test.go index 080c5e2f7..d549f4c39 100644 --- a/pkg/commands/git_commands/branch_loader_test.go +++ b/pkg/commands/git_commands/branch_loader_test.go @@ -6,8 +6,10 @@ import ( "testing" "time" + "github.com/go-errors/errors" "github.com/jesseduffield/lazygit/pkg/commands/models" "github.com/jesseduffield/lazygit/pkg/commands/oscommands" + "github.com/sasha-s/go-deadlock" "github.com/stretchr/testify/assert" ) @@ -291,3 +293,62 @@ func TestGetBehindBaseBranchValuesForAllBranches_LegacyPath(t *testing.T) { runner.CheckForMissingCalls() } + +func TestCheckForRewrittenUpstreams(t *testing.T) { + branch := func(name string, ahead string, behind string) *models.Branch { + return &models.Branch{ + Name: name, + UpstreamRemote: "origin", + UpstreamBranch: name, + AheadForPull: ahead, + BehindForPull: behind, + } + } + + notDiverged := branch("not-diverged", "0", "2") + rewritten := branch("rewritten", "3", "5") + ownCommits := branch("own-commits", "3", "5") + failing := branch("failing", "1", "1") + + // A branch that is no longer diverged must lose the value it had before + notDiverged.UpstreamRewritten.Store(true) + + branches := []*models.Branch{notDiverged, rewritten, ownCommits, failing} + + var mutex deadlock.Mutex + queried := []string{} + hasLocalOnlyCommits := func(branch *models.Branch) (bool, error) { + mutex.Lock() + queried = append(queried, branch.Name) + mutex.Unlock() + + switch branch.Name { + case "own-commits": + return true, nil + case "failing": + return false, errors.New("error") + default: + return false, nil + } + } + + gitCommon := buildGitCommon(commonDeps{}) + loader := &BranchLoader{ + Common: gitCommon.Common, + GitCommon: gitCommon, + cmd: gitCommon.cmd, + hasLocalOnlyCommits: hasLocalOnlyCommits, + } + + rendered := false + err := loader.checkForRewrittenUpstreams(branches, func() { rendered = true }) + assert.NoError(t, err) + assert.True(t, rendered, "renderFunc should have been called") + + assert.ElementsMatch(t, []string{"rewritten", "own-commits", "failing"}, queried, + "only diverged branches should be looked at") + assert.False(t, notDiverged.UpstreamRewritten.Load()) + assert.True(t, rewritten.UpstreamRewritten.Load()) + assert.False(t, ownCommits.UpstreamRewritten.Load()) + assert.False(t, failing.UpstreamRewritten.Load(), "a failed check should not claim anything") +} diff --git a/pkg/commands/models/branch.go b/pkg/commands/models/branch.go index 29b8fccf0..b01e5a7ef 100644 --- a/pkg/commands/models/branch.go +++ b/pkg/commands/models/branch.go @@ -47,6 +47,13 @@ type Branch struct { // determined yet, or up to date with base branch. (We don't need to // distinguish the two, as we don't draw anything in both cases.) BehindBaseBranch atomic.Int32 + + // Whether the branch has diverged from its upstream because the upstream + // branch was rewritten, and not because the branch has commits of its own. + // Such a branch can be reset to its upstream without losing anything. + // False for branches that haven't diverged, and for those we haven't + // determined it for yet. + UpstreamRewritten atomic.Bool } func (b *Branch) FullRefName() string { diff --git a/pkg/gui/presentation/branches.go b/pkg/gui/presentation/branches.go index de3a5bb11..75abd640f 100644 --- a/pkg/gui/presentation/branches.go +++ b/pkg/gui/presentation/branches.go @@ -26,6 +26,8 @@ type branchColorPattern struct { var branchColorPatterns []branchColorPattern +var dimYellow = style.FgYellow.SetDim() + func GetBranchListDisplayStrings( branches []*models.Branch, getItemOperation func(item types.HasUrn) types.ItemOperation, @@ -223,7 +225,12 @@ func BranchStatus( } else if branch.RemoteBranchNotStoredLocally() { result = style.FgMagenta.Sprint("?") } else if branch.IsBehindForPull() && branch.IsAheadForPull() { - result = style.FgYellow.Sprintf("↓%s↑%s", branch.BehindForPull, branch.AheadForPull) + // A branch that diverged only because its upstream was rewritten + // has no commits of its own, and fast-forwarding it resolves the + // divergence. Dim it to set it apart from a branch whose + // divergence needs a decision. + divergenceStyle := lo.Ternary(branch.UpstreamRewritten.Load(), dimYellow, style.FgYellow) + result = divergenceStyle.Sprintf("↓%s↑%s", branch.BehindForPull, branch.AheadForPull) } else if branch.IsBehindForPull() { result = style.FgYellow.Sprintf("↓%s", branch.BehindForPull) } else if branch.IsAheadForPull() { diff --git a/pkg/gui/presentation/branches_test.go b/pkg/gui/presentation/branches_test.go index c8fecf140..01c61a0c4 100644 --- a/pkg/gui/presentation/branches_test.go +++ b/pkg/gui/presentation/branches_test.go @@ -457,3 +457,30 @@ func TestGetBranchTextStyle(t *testing.T) { SetCustomBranches(patterns) assert.Equal(t, style.FgRed, GetBranchTextStyle("feature/ISSUE-1")) } + +func TestBranchStatus(t *testing.T) { + oldColorLevel := color.ForceSetColorLevel(terminfo.ColorLevelMillions) + defer color.ForceSetColorLevel(oldColorLevel) + + c := common.NewDummyCommon() + + divergedBranch := func(upstreamRewritten bool) *models.Branch { + branch := &models.Branch{ + Name: "branch", + UpstreamRemote: "origin", + UpstreamBranch: "branch", + AheadForPull: "3", + BehindForPull: "5", + } + branch.UpstreamRewritten.Store(upstreamRewritten) + return branch + } + + status := func(branch *models.Branch) string { + return BranchStatus(branch, types.ItemOperationNone, c.Tr, time.Time{}, c.UserConfig()) + } + + assert.Equal(t, "\x1b[33m↓5↑3\x1b[0m", status(divergedBranch(false))) + assert.Equal(t, "\x1b[33;2m↓5↑3\x1b[0m", status(divergedBranch(true)), + "a branch whose upstream was rewritten should be dimmed") +}