feat: own the branch synchronization engine

Replace the external hub sync dependency with an attributed internal implementation that fetches and updates branches independently. Allow dirty repositories while protecting affected checked-out branches, retain all-clean strict mode, add real-remote safety tests, and update the CLI, TUI, JSON output, version, and documentation.
This commit is contained in:
2026-09-06 19:41:53 +01:00
parent ca53874b48
commit a980153b3e
8 changed files with 703 additions and 170 deletions
+383
View File
@@ -0,0 +1,383 @@
package main
// The branch synchronization behavior in this file is adapted from hub's
// `hub sync` command. See THIRD_PARTY_NOTICES.md for attribution and license.
import (
"bytes"
"context"
"errors"
"fmt"
"os/exec"
"sort"
"strings"
"time"
)
type branchSyncResult struct {
Branch string `json:"branch"`
Action string `json:"action"`
Message string `json:"message,omitempty"`
}
type repositorySyncReport struct {
Message string
Branches []branchSyncResult
}
type localBranch struct {
name string
ref string
oid string
}
func syncRepository(path string, timeout time.Duration, initiallyDirty bool, ignoredNested []string) (repositorySyncReport, error) {
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
var report repositorySyncReport
remote, err := mainRemote(ctx, path)
if err != nil {
return failedReport(report, ctx, timeout, err)
}
fetchOutput, err := gitCombinedContext(ctx, path, "fetch", "--prune", "--quiet", "--progress", remote)
if err != nil {
return failedReport(report, ctx, timeout, gitCommandError("fetch "+remote, fetchOutput, err))
}
if operation := gitOperation(path); operation != "" {
return failedReport(report, ctx, timeout, fmt.Errorf("%s started during sync", operation))
}
current, checkedOut, err := worktreeState(ctx, path)
if err != nil {
return failedReport(report, ctx, timeout, err)
}
branches, err := localBranches(ctx, path)
if err != nil {
return failedReport(report, ctx, timeout, err)
}
defaultRef, defaultBranch, err := remoteDefaultBranch(ctx, path, remote)
if err != nil {
return failedReport(report, ctx, timeout, err)
}
for _, branch := range branches {
target, gone, err := upstreamForBranch(ctx, path, remote, branch.name)
if err != nil {
return failedReport(report, ctx, timeout, err)
}
if target != "" {
targetOID, err := resolveCommit(ctx, path, target)
if err != nil {
return failedReport(report, ctx, timeout, fmt.Errorf("resolve upstream for %s: %w", branch.name, err))
}
if targetOID == branch.oid {
continue
}
behind, err := isAncestor(ctx, path, branch.oid, targetOID)
if err != nil {
return failedReport(report, ctx, timeout, fmt.Errorf("compare branch %s with its upstream: %w", branch.name, err))
}
if !behind {
report.add(branch.name, "warning", "contains commits that are not in its upstream")
continue
}
current, checkedOut, err = worktreeState(ctx, path)
if err != nil {
return failedReport(report, ctx, timeout, err)
}
if branch.name == current {
dirty, err := worktreeDirty(ctx, path, ignoredNested)
if err != nil {
return failedReport(report, ctx, timeout, err)
}
if initiallyDirty || dirty {
report.add(branch.name, "protected", "checked-out branch has worktree changes")
continue
}
output, err := gitCombinedContext(ctx, path, "merge", "--ff-only", "--quiet", target)
if err != nil {
return failedReport(report, ctx, timeout, gitCommandError("fast-forward branch "+branch.name, output, err))
}
} else if checkedOut[branch.ref] {
report.add(branch.name, "protected", "checked out in another worktree")
continue
} else {
output, err := gitCombinedContext(ctx, path, "update-ref", branch.ref, targetOID, branch.oid)
if err != nil {
return failedReport(report, ctx, timeout, gitCommandError("fast-forward branch "+branch.name, output, err))
}
}
report.add(branch.name, "updated", "fast-forwarded from "+abbreviateOID(branch.oid))
continue
}
if !gone {
continue
}
if defaultRef == "" {
report.add(branch.name, "warning", "upstream was deleted; kept because the remote default branch is unknown")
continue
}
merged, err := isAncestor(ctx, path, branch.oid, defaultRef)
if err != nil {
return failedReport(report, ctx, timeout, fmt.Errorf("check whether branch %s is merged: %w", branch.name, err))
}
if !merged {
report.add(branch.name, "warning", "upstream was deleted, but the branch is not merged into "+defaultBranch)
continue
}
current, checkedOut, err = worktreeState(ctx, path)
if err != nil {
return failedReport(report, ctx, timeout, err)
}
if checkedOut[branch.ref] {
if branch.name != current {
report.add(branch.name, "protected", "upstream was deleted, but the branch is checked out in another worktree")
continue
}
dirty, err := worktreeDirty(ctx, path, ignoredNested)
if err != nil {
return failedReport(report, ctx, timeout, err)
}
if initiallyDirty || dirty {
report.add(branch.name, "protected", "upstream was deleted, but the checked-out branch has worktree changes")
continue
}
if !hasLocalBranch(branches, defaultBranch) {
report.add(branch.name, "protected", "upstream was deleted, but the local default branch does not exist")
continue
}
output, err := gitCombinedContext(ctx, path, "checkout", "--quiet", defaultBranch)
if err != nil {
return failedReport(report, ctx, timeout, gitCommandError("check out default branch "+defaultBranch, output, err))
}
current = defaultBranch
}
output, err := gitCombinedContext(ctx, path, "branch", "-D", "--", branch.name)
if err != nil {
return failedReport(report, ctx, timeout, gitCommandError("delete merged branch "+branch.name, output, err))
}
report.add(branch.name, "deleted", "upstream was deleted and the branch was merged into "+defaultBranch)
}
if len(report.Branches) == 0 {
report.Message = "Already up to date."
} else {
report.Message = formatBranchResults(report.Branches)
}
return report, nil
}
func (r *repositorySyncReport) add(branch, action, message string) {
r.Branches = append(r.Branches, branchSyncResult{Branch: branch, Action: action, Message: message})
}
func failedReport(report repositorySyncReport, ctx context.Context, timeout time.Duration, err error) (repositorySyncReport, error) {
if ctx.Err() == context.DeadlineExceeded {
err = fmt.Errorf("timed out after %s", timeout)
}
report.Message = err.Error()
return report, err
}
func mainRemote(ctx context.Context, path string) (string, error) {
output, err := gitOutputContext(ctx, path, "remote")
if err != nil {
return "", fmt.Errorf("list remotes: %w", err)
}
remotes := strings.Fields(output)
if len(remotes) == 0 {
return "", errors.New("no remotes")
}
for _, preferred := range []string{"upstream", "github", "origin"} {
for _, remote := range remotes {
if remote == preferred {
return remote, nil
}
}
}
if len(remotes) == 1 {
return remotes[0], nil
}
return "", fmt.Errorf("cannot choose a main remote from %s; name one upstream, github, or origin", strings.Join(remotes, ", "))
}
func localBranches(ctx context.Context, path string) ([]localBranch, error) {
output, err := gitBytesContext(ctx, path, "for-each-ref", "--format=%(refname)%00%(refname:short)%00%(objectname)", "refs/heads")
if err != nil {
return nil, fmt.Errorf("list local branches: %w", err)
}
var branches []localBranch
for _, line := range bytes.Split(bytes.TrimSpace(output), []byte{'\n'}) {
parts := bytes.Split(line, []byte{0})
if len(parts) != 3 {
continue
}
branches = append(branches, localBranch{ref: string(parts[0]), name: string(parts[1]), oid: string(parts[2])})
}
sort.Slice(branches, func(i, j int) bool { return branches[i].name < branches[j].name })
return branches, nil
}
func upstreamForBranch(ctx context.Context, path, remote, branch string) (target string, gone bool, err error) {
configuredRemote, _ := gitOutputContext(ctx, path, "config", "--get", "branch."+branch+".remote")
configuredRemote = strings.TrimSpace(configuredRemote)
mergeRef, _ := gitOutputContext(ctx, path, "config", "--get", "branch."+branch+".merge")
mergeRef = strings.TrimSpace(mergeRef)
if configuredRemote == remote && mergeRef != "" {
target = "refs/remotes/" + remote + "/" + strings.TrimPrefix(mergeRef, "refs/heads/")
exists, err := refExists(ctx, path, target)
if err != nil {
return "", false, err
}
if !exists {
return "", true, nil
}
return target, false, nil
}
target = "refs/remotes/" + remote + "/" + branch
exists, err := refExists(ctx, path, target)
if err != nil {
return "", false, err
}
if !exists {
return "", false, nil
}
return target, false, nil
}
func remoteDefaultBranch(ctx context.Context, path, remote string) (ref, branch string, err error) {
head := "refs/remotes/" + remote + "/HEAD"
if output, symbolicErr := gitOutputContext(ctx, path, "symbolic-ref", "--quiet", head); symbolicErr == nil {
ref = strings.TrimSpace(output)
exists, existsErr := refExists(ctx, path, ref)
if existsErr != nil {
return "", "", existsErr
}
if exists {
return ref, strings.TrimPrefix(ref, "refs/remotes/"+remote+"/"), nil
}
} else if ctx.Err() != nil {
return "", "", ctx.Err()
}
return "", "", nil
}
func checkedOutBranches(ctx context.Context, path string) (map[string]bool, error) {
output, err := gitOutputContext(ctx, path, "worktree", "list", "--porcelain")
if err != nil {
return nil, fmt.Errorf("list worktrees: %w", err)
}
branches := make(map[string]bool)
for _, line := range strings.Split(output, "\n") {
if strings.HasPrefix(line, "branch ") {
branches[strings.TrimSpace(strings.TrimPrefix(line, "branch "))] = true
}
}
return branches, nil
}
func worktreeState(ctx context.Context, path string) (string, map[string]bool, error) {
current, err := gitOutputContext(ctx, path, "symbolic-ref", "--quiet", "--short", "HEAD")
if err != nil {
return "", nil, errors.New("cannot determine the checked-out branch after fetch")
}
checkedOut, err := checkedOutBranches(ctx, path)
if err != nil {
return "", nil, err
}
return strings.TrimSpace(current), checkedOut, nil
}
func worktreeDirty(ctx context.Context, path string, ignoredNested []string) (bool, error) {
output, err := gitBytesContext(ctx, path, "status", "--porcelain=v1", "-z", "--untracked-files=all", "--ignore-submodules=none")
if err != nil {
return false, fmt.Errorf("read worktree status: %w", err)
}
for _, change := range parsePorcelain(output) {
if change.code != "??" || !belongsToNestedRepo(change.path, ignoredNested) {
return true, nil
}
}
return false, nil
}
func resolveCommit(ctx context.Context, path, ref string) (string, error) {
output, err := gitOutputContext(ctx, path, "rev-parse", "--verify", ref+"^{commit}")
return strings.TrimSpace(output), err
}
func refExists(ctx context.Context, path, ref string) (bool, error) {
cmd := exec.CommandContext(ctx, "git", "-C", path, "show-ref", "--verify", "--quiet", ref)
err := cmd.Run()
if err == nil {
return true, nil
}
var exitErr *exec.ExitError
if errors.As(err, &exitErr) && exitErr.ExitCode() == 1 {
return false, nil
}
return false, err
}
func isAncestor(ctx context.Context, path, older, newer string) (bool, error) {
cmd := exec.CommandContext(ctx, "git", "-C", path, "merge-base", "--is-ancestor", older, newer)
err := cmd.Run()
if err == nil {
return true, nil
}
var exitErr *exec.ExitError
if errors.As(err, &exitErr) && exitErr.ExitCode() == 1 {
return false, nil
}
return false, err
}
func hasLocalBranch(branches []localBranch, name string) bool {
for _, branch := range branches {
if branch.name == name {
return true
}
}
return false
}
func gitOutputContext(ctx context.Context, path string, args ...string) (string, error) {
output, err := gitBytesContext(ctx, path, args...)
return string(output), err
}
func gitBytesContext(ctx context.Context, path string, args ...string) ([]byte, error) {
cmd := exec.CommandContext(ctx, "git", append([]string{"-C", path}, args...)...)
return cmd.Output()
}
func gitCombinedContext(ctx context.Context, path string, args ...string) ([]byte, error) {
cmd := exec.CommandContext(ctx, "git", append([]string{"-C", path}, args...)...)
return cmd.CombinedOutput()
}
func gitCommandError(action string, output []byte, err error) error {
detail := strings.TrimSpace(string(output))
if detail == "" {
return fmt.Errorf("%s: %w", action, err)
}
return fmt.Errorf("%s: %s", action, firstLine(detail))
}
func formatBranchResults(results []branchSyncResult) string {
lines := make([]string, 0, len(results))
for _, result := range results {
lines = append(lines, fmt.Sprintf("%-9s %s: %s", strings.ToUpper(result.Action), result.Branch, result.Message))
}
return strings.Join(lines, "\n")
}
func abbreviateOID(oid string) string {
if len(oid) <= 7 {
return oid
}
return oid[:7]
}