From 6514203b2b430ba9d6dc6e2b95c3b9ae831bc8d4 Mon Sep 17 00:00:00 2001 From: Samuel K Date: Thu, 23 Jul 2026 22:46:57 -0500 Subject: [PATCH] feat(git): track request source branch as upstream on checkout CheckoutPR fetched a pull/merge request into a synthetic MR/PR branch with no upstream, so the local branch did not reference the remote branch it came from. After checkout, resolve the origin branch whose tip matches the request head (the head ref carries the commit but not the branch name) and set it as upstream. Forks, shallow clones, and ambiguous matches leave the branch untracked, preserving prior behavior. --- pkg/git/repo.go | 71 +++++++++++++++++++++++++++++++------------- pkg/git/repo_test.go | 31 +++++++++++++++++++ 2 files changed, 81 insertions(+), 21 deletions(-) diff --git a/pkg/git/repo.go b/pkg/git/repo.go index 20d18838a..553c3128f 100644 --- a/pkg/git/repo.go +++ b/pkg/git/repo.go @@ -9,24 +9,20 @@ import ( "github.com/devsy-org/devsy/pkg/log" ) -// Repo is a handle to a git repository at a filesystem path. type Repo struct { path string env []string runner Runner } -// RepoOption configures a Repo. type RepoOption func(*Repo) -// WithEnv appends extra environment entries applied to every operation. func WithEnv(env []string) RepoOption { return func(r *Repo) { r.env = append(r.env, env...) } } -// WithStrictHostKeyChecking appends the default SSH/terminal environment honoring the given policy. func WithStrictHostKeyChecking(strict bool) RepoOption { return WithEnv(GetDefaultExtraEnv(strict)) } @@ -38,7 +34,6 @@ func WithRunner(runner Runner) RepoOption { } } -// At returns a Repo rooted at path. func At(path string, opts ...RepoOption) *Repo { r := &Repo{path: path, runner: defaultRunner} for _, opt := range opts { @@ -50,7 +45,6 @@ func At(path string, opts ...RepoOption) *Repo { return r } -// Path returns the repository's filesystem path. func (r *Repo) Path() string { return r.path } @@ -64,7 +58,6 @@ const ( ResetHard ResetMode = "--hard" ) -// Fetch fetches the given refspec from origin. func (r *Repo) Fetch(ctx context.Context, refspec string) error { if err := r.runLogged(ctx, "fetch", "origin", refspec); err != nil { return fmt.Errorf("fetch %q: %w", refspec, err) @@ -72,7 +65,6 @@ func (r *Repo) Fetch(ctx context.Context, refspec string) error { return nil } -// Switch switches the working tree to an existing branch. func (r *Repo) Switch(ctx context.Context, branch string) error { if _, err := r.run(ctx, "switch", branch); err != nil { return fmt.Errorf("switch to branch %q: %w", branch, err) @@ -81,10 +73,8 @@ func (r *Repo) Switch(ctx context.Context, branch string) error { } // CheckoutPR fetches a pull/merge request into a local branch and switches to -// it. The request number is taken from prRef; the remote refspec is resolved -// against repoURL's hosting provider, falling back to the other known -// conventions when the detected one has no such ref (e.g. self-hosted GitLab on -// a custom domain that URL detection can't recognize). +// it, trying each provider convention so a self-hosted host that URL detection +// can't recognize still resolves against the other's refspec. func (r *Repo) CheckoutPR(ctx context.Context, repoURL, prRef string) error { number := prNumber(prRef) if number == "" { @@ -99,7 +89,11 @@ func (r *Repo) CheckoutPR(ctx context.Context, repoURL, prRef string) error { err := r.Fetch(ctx, refspec+":"+prBranch) if err == nil { - return r.Switch(ctx, prBranch) + if err := r.Switch(ctx, prBranch); err != nil { + return err + } + r.trackRequestSourceBranch(ctx, prBranch) + return nil } if !isMissingRefError(err) { return err @@ -122,7 +116,6 @@ func isMissingRefError(err error) bool { strings.Contains(msg, "not found in upstream") } -// Reset moves HEAD to commit using the given mode. func (r *Repo) Reset(ctx context.Context, commit string, mode ResetMode) error { if err := r.runLogged(ctx, "reset", string(mode), commit); err != nil { return fmt.Errorf("reset head to %q: %w", commit, err) @@ -147,13 +140,10 @@ func (r *Repo) LsTree(ctx context.Context, ref string) ([]string, error) { return splitLines(res.Stdout), nil } -// Clone clones repository into the repo's path using the given options. func (r *Repo) Clone(ctx context.Context, repository string, options ...Option) error { return r.cloneWith(ctx, repository, newCloneConfig(options...)) } -// CloneFromInfo clones the repository described by gitInfo into the repo's path, -// using the given credential helper and options. func (r *Repo) CloneFromInfo( ctx context.Context, gitInfo *GitInfo, @@ -189,7 +179,7 @@ func (r *Repo) CloneFromInfo( } // CredentialFill runs `git credential fill`, feeding request on stdin and -// returning git's response. Used to resolve credentials via configured helpers. +// returning git's response. func (r *Repo) CredentialFill(ctx context.Context, request string) (string, error) { res, err := r.runner.Run(ctx, RunOptions{ Dir: r.path, @@ -203,7 +193,48 @@ func (r *Repo) CredentialFill(ctx context.Context, request string) (string, erro return string(res.Stdout), nil } -// run executes a git subcommand in the repository, capturing its output. +// trackRequestSourceBranch best-effort sets prBranch's upstream to the origin +// branch it was opened from, matched by commit since the head ref lacks the +// branch name. Forks, shallow clones, and ambiguous matches stay untracked. +func (r *Repo) trackRequestSourceBranch(ctx context.Context, prBranch string) { + source, err := r.originBranchAt(ctx, prBranch) + if err != nil { + log.Debugf("could not resolve source branch for %s: %v", prBranch, err) + return + } + if source == "" { + return + } + if _, err := r.run(ctx, "branch", "--set-upstream-to="+source, prBranch); err != nil { + log.Debugf("could not set upstream of %s to %s: %v", prBranch, source, err) + return + } + log.Debugf("branch %s now tracks %s", prBranch, source) +} + +// originBranchAt returns the origin branch whose tip is rev's commit, or "" when +// there is not exactly one (origin/HEAD aside). +func (r *Repo) originBranchAt(ctx context.Context, rev string) (string, error) { + res, err := r.run(ctx, "for-each-ref", + "--points-at", rev, + "--format=%(refname:short)", + "refs/remotes/origin") + if err != nil { + return "", err + } + var matches []string + for _, branch := range splitLines(res.Stdout) { + if branch == "origin" || branch == "origin/HEAD" { + continue + } + matches = append(matches, branch) + } + if len(matches) != 1 { + return "", nil + } + return matches[0], nil +} + func (r *Repo) run(ctx context.Context, args ...string) (RunResult, error) { return r.runner.Run(ctx, RunOptions{ Dir: r.path, @@ -227,7 +258,6 @@ func (r *Repo) runLogged(ctx context.Context, args ...string) error { return err } -// cloneWith executes a `git clone` with the given config, streaming output to the logs. func (r *Repo) cloneWith(ctx context.Context, repository string, c cloneConfig) error { w := log.Writer(log.LevelInfo) defer func() { _ = w.Close() }() @@ -246,7 +276,6 @@ func (r *Repo) cloneWith(ctx context.Context, repository string, c cloneConfig) return err } -// splitLines splits command output into non-empty trimmed lines. func splitLines(b []byte) []string { var lines []string for line := range strings.SplitSeq(string(b), "\n") { diff --git a/pkg/git/repo_test.go b/pkg/git/repo_test.go index 39a6d5adb..bf0cc1dac 100644 --- a/pkg/git/repo_test.go +++ b/pkg/git/repo_test.go @@ -132,6 +132,37 @@ func TestRepoCheckoutPRNonRefErrorNoFallback(t *testing.T) { assert.Equal(t, 1, len(fake.calls)) } +func TestRepoCheckoutPRTracksSourceBranch(t *testing.T) { + // for-each-ref resolves the source branch by commit; a single match is + // adopted as the local branch's upstream. + fake := &fakeRunner{stdout: []byte("origin/feature-x\n")} + repo := At("/tmp/repo", WithRunner(fake)) + + err := repo.CheckoutPR(context.Background(), testGitLabURL, "merge-requests/7/head") + assert.NilError(t, err) + assert.DeepEqual(t, []string{subFetch, originRemote, testMRRefSpec}, fake.calls[0].Args) + assert.DeepEqual(t, []string{subSwitch, testMRLocal}, fake.calls[1].Args) + assert.DeepEqual(t, []string{ + "for-each-ref", "--points-at", testMRLocal, + "--format=%(refname:short)", "refs/remotes/origin", + }, fake.calls[2].Args) + assert.DeepEqual(t, []string{ + "branch", "--set-upstream-to=origin/feature-x", testMRLocal, + }, fake.calls[3].Args) +} + +func TestRepoCheckoutPRAmbiguousSourceBranchLeftUntracked(t *testing.T) { + // More than one origin branch points at the head (e.g. the default branch + // too): the match is ambiguous, so no upstream is configured. + fake := &fakeRunner{stdout: []byte("origin/HEAD\norigin/main\norigin/feature-x\n")} + repo := At("/tmp/repo", WithRunner(fake)) + + err := repo.CheckoutPR(context.Background(), testGitLabURL, "merge-requests/7/head") + assert.NilError(t, err) + assert.Equal(t, 3, len(fake.calls)) // fetch, switch, for-each-ref — no set-upstream + assert.Equal(t, "for-each-ref", fake.calls[2].Args[0]) +} + func TestRepoLsRemote(t *testing.T) { fake := &fakeRunner{} repo := At("", WithRunner(fake))