# 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

```go
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
}
```

```go
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)
}
```
