add hidden checkpoint policy command · Entire

Add Hidden Checkpoint Policy Command

4518d7b→main · pfleidi · 3w ago · 16 files · +1,284 added/-8 removed

Store the repo checkpoint policy in a dedicated git ref and expose a hidden command for inspecting and updating it.

The command syncs with the checkpoint remote before updates, rejects unsupported versions, and pushes only the policy ref.

Sessions

Changes

16

1
2
3
4
5
6
7
18 unmodified lines

26
27
28
28
29
30
31
32
12 unmodified lines

45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
51
52
53
61
62
63
64
65
66
67
68
69
70
71
58
59
60
72
73
74
75
76
77
78
79
80

package checkpointpolicy

import (
    "cmp"
    "fmt"
    "strconv"
    "strings"
18 unmodified lines

}

family := CheckpointFamily(familyRaw)
    if _, ok := knownFamilies[family]; !ok {
    if _, ok := familyRanks[family]; !ok {
        return CheckpointFormat{}, fmt.Errorf("unknown checkpoint family %q", familyRaw)
    }

12 unmodified lines

return fmt.Sprintf("%s-v%d", f.Family, f.Major)
}

func Compare(a, b CheckpointFormat) int {
    aRank := familyRanks[a.Family]
    bRank := familyRanks[b.Family]
    if aRank != bRank {
        return cmp.Compare(aRank, bRank)
    }
    return cmp.Compare(a.Major, b.Major)
}

func CanRead(format CheckpointFormat) bool {
    return readFormats[format]
}

var knownFamilies = map[CheckpointFamily]bool{
    CheckpointFamilyBranch: true,
    CheckpointFamilyRefs:   true,
}

func CanWrite(format CheckpointFormat) bool {
    return writeFormats[format]
}

var familyRanks = map[CheckpointFamily]int{
    CheckpointFamilyBranch: 0,
    CheckpointFamilyRefs:   1,
}

var branchV1Format = CheckpointFormat{Family: CheckpointFamilyBranch, Major: 1}

var readFormats = map[CheckpointFormat]bool{
    branchV1Format: true,
}
var (
    readFormats = map[CheckpointFormat]bool{
        branchV1Format: true,
    }

writeFormats = map[CheckpointFormat]bool{
        branchV1Format: true,
    }
)

Mcmd/entire/cli/checkpointpolicy/format.go +27/-7

37 unmodified lines

38
39
40
41
41
42
43
44
2 unmodified lines

47
48
49
50
51
52
53
54
55
56

37 unmodified lines

}
}

func TestCanReadFormat(t *testing.T) {
func TestSupportedFormats(t *testing.T) {
    t.Parallel()

branchV1, err := checkpointpolicy.ParseFormat(checkpoint.CheckpointVersionBranchV1)
2 unmodified lines

require.NoError(t, err)

require.True(t, checkpointpolicy.CanRead(branchV1))
    require.True(t, checkpointpolicy.CanWrite(branchV1))
    require.Equal(t, checkpoint.CheckpointVersionBranchV1, branchV1.String())

require.False(t, checkpointpolicy.CanRead(refsV1))
    require.False(t, checkpointpolicy.CanWrite(refsV1))
    require.Negative(t, checkpointpolicy.Compare(branchV1, refsV1))
}

Mcmd/entire/cli/checkpointpolicy/format_test.go +6/-1

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41

package checkpointpolicy

import (
    "fmt"

"github.com/entireio/cli/cmd/entire/cli/checkpoint"
)

type Policy struct {
    CheckpointVersion    string `json:"checkpoint_version"`
    CheckpointMinVersion string `json:"checkpoint_min_version"`
}

func DefaultPolicy() Policy {
    return Policy{
        CheckpointVersion:    checkpoint.CheckpointVersionBranchV1,
        CheckpointMinVersion: checkpoint.CheckpointVersionBranchV1,
    }
}

func Normalize(policy Policy) Policy {
    if policy.CheckpointVersion == "" {
        policy.CheckpointVersion = checkpoint.CheckpointVersionBranchV1
    }
    if policy.CheckpointMinVersion == "" {
        policy.CheckpointMinVersion = checkpoint.CheckpointVersionBranchV1
    }
    return policy
}

func ValidatePolicy(policy Policy) error {
    policy = Normalize(policy)

version, err := ParseFormat(policy.CheckpointVersion)
    if err != nil {
        return fmt.Errorf("checkpoint_version: %w", err)
    }
    if !CanWrite(version) {
        return fmt.Errorf("checkpoint_version %q is not write-supported by this Entire CLI", policy.CheckpointVersion)
    }

minVersion, err := ParseFormat(policy.CheckpointMinVersion)
    if err != nil {
        return fmt.Errorf("checkpoint_min_version: %w", err)
    }
    if !CanRead(minVersion) {
        return fmt.Errorf("checkpoint_min_version %q is not read-supported by this Entire CLI", policy.CheckpointMinVersion)
    }
    if Compare(minVersion, version) > 0 {
        return fmt.Errorf("checkpoint_min_version %q is newer than checkpoint_version %q", policy.CheckpointMinVersion, policy.CheckpointVersion)
    }

return nil
}

Acmd/entire/cli/checkpointpolicy/policy.go +54

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41

package checkpointpolicy_test

import (
    "testing"

"github.com/entireio/cli/cmd/entire/cli/checkpoint"
    "github.com/entireio/cli/cmd/entire/cli/checkpointpolicy"
    "github.com/stretchr/testify/require"
)

func TestDefaultPolicy(t *testing.T) {
    t.Parallel()
    got := checkpointpolicy.DefaultPolicy()
    require.Equal(t, checkpoint.CheckpointVersionBranchV1, got.CheckpointVersion)
    require.Equal(t, checkpoint.CheckpointVersionBranchV1, got.CheckpointMinVersion)
}

func TestValidatePolicy(t *testing.T) {
    t.Parallel()
    tests := []struct {
        name    string
        policy  checkpointpolicy.Policy
        wantErr string
    }{
        {name: "default", policy: checkpointpolicy.DefaultPolicy()},
        {name: "unknown current", policy: checkpointpolicy.Policy{CheckpointVersion: "future-v1", CheckpointMinVersion: "branch-v1"}, wantErr: "unknown checkpoint family"},
        {name: "unsupported current", policy: checkpointpolicy.Policy{CheckpointVersion: "refs-v1", CheckpointMinVersion: "branch-v1"}, wantErr: "not write-supported"},
        {name: "unsupported minimum", policy: checkpointpolicy.Policy{CheckpointVersion: "branch-v1", CheckpointMinVersion: "refs-v1"}, wantErr: "not read-supported"},
    }
    for _, tt := range tests {
        t.Run(tt.name, func(t *testing.T) {
            t.Parallel()
            err := checkpointpolicy.ValidatePolicy(tt.policy)
            if tt.wantErr == "" {
                require.NoError(t, err)
                return
            }
            require.ErrorContains(t, err, tt.wantErr)
        })
    }
}

Acmd/entire/cli/checkpointpolicy/policy_test.go +41

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41

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 (
    sha1HexSize   = 40
    sha256HexSize = 64
)

const fetchRefName = plumbing.ReferenceName("refs/entire/policies/checkpoint-fetch")

var errStopTraversal = errors.New("stop traversal")

type Target struct {
    Remote string
    Dir    string
}

type RemoteState struct {
    Exists bool
    Hash   plumbing.Hash
}

func ResolveTarget(ctx context.Context) (Target, error) {
    dir, err := paths.WorktreeRoot(ctx)
    if err != nil {
        return Target{}, fmt.Errorf("resolve worktree root: %w", err)
    }
    target, err := remote.FetchURL(ctx, remote.FetchURLOptions{WorktreeRoot: dir})
    if err != nil {
        return Target{}, fmt.Errorf("resolve checkpoint remote URL: %w", err)
    }
    return Target{Remote: target, 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{}, fmt.Errorf("check remote checkpoint policy ref: %w", err)
    }
    fields := strings.Fields(string(output))
    if len(fields) == 0 {
        return RemoteState{}, nil
    }
    hash, err := parseRemotePolicyHash(fields[0])
    if err != nil {
        return RemoteState{}, err
    }
    return RemoteState{Exists: true, Hash: hash}, nil
}

func Sync(ctx context.Context, repo *git.Repository, target Target) (State, error) {
    local, err := ReadLocal(ctx, repo)
    if err != nil {
        return State{}, err
    }

baseline, remoteFound, err := remoteBaseline(ctx, repo, target, local)
    if err != nil {
        return State{}, err
    }
    if !remoteFound || local.Hash == baseline.Hash {
        return baseline, nil
    }

if local.Hash.IsZero() || isAncestorOf(ctx, repo, local.Hash, baseline.Hash) {
        if err := SetRef(repo, RefName, baseline.Hash); err != nil {
            return State{}, err
        }
        baseline.Source = SourceRemote
        return baseline, nil
    }

local.Source = SourceLocalDiverged
    local.RemoteHash = baseline.RemoteHash
    return local, nil
}

func remoteBaseline(ctx context.Context, repo *git.Repository, target Target, local State) (State, bool, error) {
    remoteState, err := CheckRemote(ctx, target)
    if err != nil {
        return State{}, false, err
    }
    if !remoteState.Exists {
        return local, false, nil
    }
    if local.Hash == remoteState.Hash {
        local.Source = SourceRemote
        local.RemoteHash = remoteState.Hash
        return local, true, nil
    }

fetched, err := fetchRemotePolicy(ctx, repo, target)
    if err != nil {
        return State{}, false, err
    }
    fetched.RemoteHash = remoteState.Hash
    defer removeFetchRef(repo)
    return fetched, true, nil
}

func parseRemotePolicyHash(raw string) (plumbing.Hash, error) {
    if !isSupportedRemotePolicyHashLength(raw) {
        return plumbing.ZeroHash, fmt.Errorf("invalid remote checkpoint policy hash %q", raw)
    }
    hash, ok := plumbing.FromHex(raw)
    if !ok {
        return plumbing.ZeroHash, fmt.Errorf("invalid remote checkpoint policy hash %q", raw)
    }
    return hash, nil
}

func isSupportedRemotePolicyHashLength(raw string) bool {
    return len(raw) == sha1HexSize || len(raw) == sha256HexSize
}

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{}, fmt.Errorf("fetch checkpoint policy ref: %w", err)
    }
    return ReadFromRef(ctx, repo, fetchRefName, SourceRemote)
}

func removeFetchRef(repo *git.Repository) {
    if err := repo.Storer.RemoveReference(fetchRefName); err != nil {
        return
    }
}

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
    err = iter.ForEach(func(commit *object.Commit) error {
        if err := ctx.Err(); err != nil {
            return fmt.Errorf("traverse checkpoint policy ancestry: %w", err)
        }
        if commit.Hash == ancestor {
            found = true
            return errStopTraversal
        }
        return nil
    })
    if err != nil && !errors.Is(err, errStopTraversal) {
        return false
    }
    return found
}