sync checkpoint policy refs · Entire
sync checkpoint policy refs
cd61c28·
pfleidi·3w ago·2 files·+307 added/-0 removed
Fetch, fast-forward, and strictly push the checkpoint policy ref through the configured checkpoint remote.
Sessions
adbacd1779b7View transcript
Changes
2
cmd/entire/cli/checkpointpolicy
Aremote.go+163
Aremote_test.go+144
package checkpointpolicy
import (
"context"
"errors"
"fmt"
"strings"
"github.com/entireio/cli/cmd/entire/cli/checkpoint/remote"
"github.com/entireio/cli/cmd/entire/cli/paths"
"github.com/go-git/go-git/v6"
"github.com/go-git/go-git/v6/plumbing"
"github.com/go-git/go-git/v6/plumbing/object"
)
const defaultBaseRemote = "origin"
const fetchRefName = plumbing.ReferenceName("refs/entire/policies/checkpoint-fetch")
var errStopTraversal = errors.New("stop traversal")
type Target struct {
Remote string
Label string
Dir string
}
type RemoteState struct {
Exists bool
Hash plumbing.Hash
}
func ResolveTarget(ctx context.Context, baseRemote string) (Target, error) {
if baseRemote == "" {
baseRemote = defaultBaseRemote
}
dir, err := paths.WorktreeRoot(ctx)
if err != nil {
return Target{}, err
}
target, dedicated, err := remote.PushURL(ctx, baseRemote)
if err != nil {
return Target{}, err
}
label := baseRemote
if dedicated {
label = "checkpoint remote"
}
return Target{Remote: target, Label: label, Dir: dir}, nil
}
func CheckRemote(ctx context.Context, target Target) (RemoteState, error) {
output, err := remote.LsRemoteInDir(ctx, target.Dir, target.Remote, RefName.String())
if err != nil {
return RemoteState{}, err
}
fields := strings.Fields(string(output))
if len(fields) == 0 {
return RemoteState{}, nil
}
if len(fields[0]) != 40 {
return RemoteState{}, fmt.Errorf("invalid remote checkpoint policy hash %q", fields[0])
}
return RemoteState{Exists: true, Hash: plumbing.NewHash(fields[0])}, nil
}
func Sync(ctx context.Context, repo *git.Repository, target Target) (State, error) {
local, err := ReadLocal(ctx, repo)
if err != nil {
return State{}, err
}
remoteState, err := CheckRemote(ctx, target)
if err != nil {
return State{}, err
}
if !remoteState.Exists {
return local, nil
}
if local.Hash == remoteState.Hash {
local.Source = SourceRemote
local.RemoteHash = remoteState.Hash
return local, nil
}
fetched, err := fetchRemotePolicy(ctx, repo, target)
if err != nil {
return State{}, err
}
fetched.RemoteHash = remoteState.Hash
defer func() {
_ = repo.Storer.RemoveReference(fetchRefName)
}()
if local.Hash.IsZero() || isAncestorOf(ctx, repo, local.Hash, fetched.Hash) {
if err := SetRef(repo, RefName, fetched.Hash); err != nil {
return State{}, err
}
fetched.Source = SourceRemote
return fetched, nil
}
local.Source = SourceLocalDiverged
local.RemoteHash = remoteState.Hash
local.Warning = fmt.Sprintf("local checkpoint policy %s diverges from remote %s", local.Hash, remoteState.Hash)
return local, nil
}
func Push(ctx context.Context, target Target) error {
refspec := RefName.String() + ":" + RefName.String()
result, err := remote.PushWithOptions(ctx, remote.PushOptions{
Remote: target.Remote,
RefSpecs: []string{refspec},
Dir: target.Dir,
})
if err != nil {
output := strings.TrimSpace(result.Output)
if output == "" {
return fmt.Errorf("push checkpoint policy: %w", err)
}
return fmt.Errorf("push checkpoint policy: %s: %w", output, err)
}
return nil
}
func fetchRemotePolicy(ctx context.Context, repo *git.Repository, target Target) (State, error) {
refspec := fmt.Sprintf("+%s:%s", RefName, fetchRefName)
if _, err := remote.Fetch(ctx, remote.FetchOptions{
Remote: target.Remote,
RefSpecs: []string{refspec},
NoTags: true,
NoFilter: true,
Dir: target.Dir,
}); err != nil {
return State{}, err
}
return ReadFromRef(ctx, repo, fetchRefName, SourceRemote)
}
func isAncestorOf(ctx context.Context, repo *git.Repository, ancestor, target plumbing.Hash) bool {
if ancestor == target {
return true
}
iter, err := repo.Log(&git.LogOptions{From: target})
if err != nil {
return false
}
defer iter.Close()
found := false
_ = iter.ForEach(func(commit *object.Commit) error {
if err := ctx.Err(); err != nil {
return err
}
if commit.Hash == ancestor {
found = true
return errStopTraversal
}
return nil
})
return found
}
package checkpointpolicy_test
import (
"context"
"os/exec"
"path/filepath"
"testing"
"github.com/entireio/cli/cmd/entire/cli/checkpointpolicy"
"github.com/entireio/cli/cmd/entire/cli/testutil"
"github.com/go-git/go-git/v6"
"github.com/go-git/go-git/v6/plumbing"
"github.com/stretchr/testify/require"
)
func TestSyncRemotePolicyDefaultsWhenRemoteMissing(t *testing.T) {
t.Parallel()
localDir, repo, bareDir := initPolicyRemoteFixture(t)
got, err := checkpointpolicy.Sync(t.Context(), repo, checkpointpolicy.Target{Remote: bareDir, Dir: localDir})
require.NoError(t, err)
require.Equal(t, checkpointpolicy.SourceDefaults, got.Source)
require.Equal(t, checkpointpolicy.DefaultPolicy(), got.Policy)
require.True(t, got.Hash.IsZero())
require.True(t, got.RemoteHash.IsZero())
}
func TestSyncRemotePolicyFetchesAndPromotesMissingLocalRef(t *testing.T) {
t.Parallel()
remoteDir, remoteRepo, bareDir := initPolicyRemoteFixture(t)
remoteHash, err := checkpointpolicy.WriteLocal(t.Context(), remoteRepo, plumbing.ZeroHash, checkpointpolicy.DefaultPolicy())
require.NoError(t, err)
pushPolicyRefWithGit(t, remoteDir, bareDir)
localDir, localRepo := initPolicyRepoWithDir(t)
got, err := checkpointpolicy.Sync(t.Context(), localRepo, checkpointpolicy.Target{Remote: bareDir, Dir: localDir})
require.NoError(t, err)
require.Equal(t, checkpointpolicy.SourceRemote, got.Source)
require.Equal(t, remoteHash, got.Hash)
require.Equal(t, remoteHash, got.RemoteHash)
localState, err := checkpointpolicy.ReadLocal(t.Context(), localRepo)
require.NoError(t, err)
require.Equal(t, remoteHash, localState.Hash)
}
func TestSyncRemotePolicyDoesNotLeaveTempRefWhenSHAAlreadyMatches(t *testing.T) {
t.Parallel()
localDir, repo, bareDir := initPolicyRemoteFixture(t)
localHash, err := checkpointpolicy.WriteLocal(t.Context(), repo, plumbing.ZeroHash, checkpointpolicy.DefaultPolicy())
require.NoError(t, err)
pushPolicyRefWithGit(t, localDir, bareDir)
got, err := checkpointpolicy.Sync(t.Context(), repo, checkpointpolicy.Target{Remote: bareDir, Dir: localDir})
require.NoError(t, err)
require.Equal(t, checkpointpolicy.SourceRemote, got.Source)
require.Equal(t, localHash, got.Hash)
require.Equal(t, localHash, got.RemoteHash)
requireNoRef(t, repo, "refs/entire/policies/checkpoint-fetch")
}
func TestSyncRemotePolicyKeepsDivergedLocalRef(t *testing.T) {
t.Parallel()
remoteDir, remoteRepo, bareDir := initPolicyRemoteFixture(t)
baseHash, err := checkpointpolicy.WriteLocal(t.Context(), remoteRepo, plumbing.ZeroHash, checkpointpolicy.DefaultPolicy())
require.NoError(t, err)
pushPolicyRefWithGit(t, remoteDir, bareDir)
localDir, localRepo := initPolicyRepoWithDir(t)
_, err = checkpointpolicy.Sync(t.Context(), localRepo, checkpointpolicy.Target{Remote: bareDir, Dir: localDir})
require.NoError(t, err)
localHash, err := checkpointpolicy.WriteLocal(t.Context(), localRepo, baseHash, checkpointpolicy.DefaultPolicy())
require.NoError(t, err)
remoteHash, err := checkpointpolicy.WriteLocal(t.Context(), remoteRepo, baseHash, checkpointpolicy.Policy{
CheckpointVersion: "refs-v1",
CheckpointMinVersion: "refs-v1",
})
require.NoError(t, err)
pushPolicyRefWithGit(t, remoteDir, bareDir)
got, err := checkpointpolicy.Sync(t.Context(), localRepo, checkpointpolicy.Target{Remote: bareDir, Dir: localDir})
require.NoError(t, err)
require.Equal(t, checkpointpolicy.SourceLocalDiverged, got.Source)
require.Equal(t, localHash, got.Hash)
require.Equal(t, remoteHash, got.RemoteHash)
localState, err := checkpointpolicy.ReadLocal(t.Context(), localRepo)
require.NoError(t, err)
require.Equal(t, localHash, localState.Hash)
requireNoRef(t, localRepo, "refs/entire/policies/checkpoint-fetch")
}
func TestPushPolicyRejectsNonFastForward(t *testing.T) {
t.Parallel()
firstDir, firstRepo, bareDir := initPolicyRemoteFixture(t)
_, err := checkpointpolicy.WriteLocal(t.Context(), firstRepo, plumbing.ZeroHash, checkpointpolicy.DefaultPolicy())
require.NoError(t, err)
pushPolicyRefWithGit(t, firstDir, bareDir)
secondDir, secondRepo := initPolicyRepoWithDir(t)
_, err = checkpointpolicy.WriteLocal(t.Context(), secondRepo, plumbing.ZeroHash, checkpointpolicy.Policy{
CheckpointVersion: "refs-v1",
CheckpointMinVersion: "refs-v1",
})
require.NoError(t, err)
err = checkpointpolicy.Push(t.Context(), checkpointpolicy.Target{Remote: bareDir, Dir: secondDir})
require.ErrorContains(t, err, "push checkpoint policy")
}
func initPolicyRemoteFixture(t *testing.T) (string, *git.Repository, string) {
t.Helper()
localDir, repo := initPolicyRepoWithDir(t)
bareDir := filepath.Join(t.TempDir(), "remote.git")
_, err := git.PlainInit(bareDir, true)
require.NoError(t, err)
return localDir, repo, bareDir
}
func initPolicyRepoWithDir(t *testing.T) (string, *git.Repository) {
t.Helper()
dir := t.TempDir()
testutil.InitRepo(t, dir)
repo, err := git.PlainOpen(dir)
require.NoError(t, err)
return dir, repo
}
func pushPolicyRefWithGit(t *testing.T, dir, remote string) {
t.Helper()
refspec := checkpointpolicy.RefName.String() + ":" + checkpointpolicy.RefName.String()
cmd := exec.CommandContext(context.Background(), "git", "push", remote, refspec)
cmd.Dir = dir
cmd.Env = testutil.GitIsolatedEnv()
output, err := cmd.CombinedOutput()
require.NoError(t, err, string(output))
}
func requireNoRef(t *testing.T, repo *git.Repository, refName plumbing.ReferenceName) {
t.Helper()
_, err := repo.Reference(refName, true)
require.ErrorIs(t, err, plumbing.ErrReferenceNotFound)
}