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

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