Skip to content
34 changes: 2 additions & 32 deletions cmd/entire/cli/doctor.go
Original file line number Diff line number Diff line change
Expand Up @@ -434,23 +434,16 @@ func promptSessionAction(ss stuckSession) (string, error) {
}

// discardSession removes session state and cleans up the shadow branch.
func discardSession(ctx context.Context, ss stuckSession, _ *git.Repository, errW io.Writer) error {
func discardSession(ctx context.Context, ss stuckSession, repo *git.Repository, errW io.Writer) error {
// Clear session state file
if err := strategy.ClearSessionStateWithProgress(ctx, ss.State.SessionID, errW, strategy.SessionLockNoticeDelay); err != nil {
return fmt.Errorf("failed to clear session state: %w", err)
}

// Delete shadow branch if it exists and no other sessions need it
if ss.HasShadowBranch {
if shouldDelete, err := canDeleteShadowBranch(ctx, ss.ShadowBranch, ss.State.SessionID); err != nil {
if _, err := strategy.DeleteShadowBranchIfUnused(ctx, repo, ss.ShadowBranch, ss.State.SessionID); err != nil {
fmt.Fprintf(errW, "Warning: could not check other sessions for shadow branch: %v\n", err)
} else if shouldDelete {
if err := strategy.DeleteBranchCLI(ctx, ss.ShadowBranch); err != nil {
// Branch already gone is not an error — keeps discard idempotent
if !errors.Is(err, strategy.ErrBranchNotFound) {
return fmt.Errorf("failed to delete shadow branch: %w", err)
}
}
}
}

Expand Down Expand Up @@ -1597,26 +1590,3 @@ func writeCodexTrackedHooksRemedy(w io.Writer) {
fmt.Fprintln(w, " .codex/hooks.json is tracked — commit it and make sure the root worktree has it")
fmt.Fprintln(w, " (merge to the default branch, or check that branch out there).")
}

// canDeleteShadowBranch checks if a shadow branch can be safely deleted.
// Returns true if no other sessions (besides excludeSessionID) need this branch.
func canDeleteShadowBranch(ctx context.Context, shadowBranch, excludeSessionID string) (bool, error) {
states, err := strategy.ListSessionStates(ctx)
if err != nil {
return false, fmt.Errorf("failed to list session states: %w", err)
}

for _, state := range states {
if state.SessionID == excludeSessionID {
continue
}
// Task records never live on the shadow branch, so only SaveStep
// checkpoints pin it alive.
otherShadow := checkpoint.ShadowBranchNameForCommit(state.BaseCommit, state.WorktreeID)
if otherShadow == shadowBranch && state.StepCount > 0 {
return false, nil
}
}

return true, nil
}
32 changes: 32 additions & 0 deletions cmd/entire/cli/doctor_cleanup_inventory_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
package cli

import (
"os"
"path/filepath"
"testing"
"time"

"github.com/entireio/cli/cmd/entire/cli/checkpoint"
"github.com/entireio/cli/cmd/entire/cli/session"
"github.com/entireio/cli/cmd/entire/cli/strategy"
"github.com/entireio/cli/cmd/entire/cli/testutil"
"github.com/stretchr/testify/require"
)

func TestCanDeleteShadowBranch_RejectsIncompleteInventory(t *testing.T) {
dir := t.TempDir()
testutil.InitRepo(t, dir)
t.Chdir(dir)
state := &session.State{SessionID: "protected", BaseCommit: "abc123456789", StartedAt: time.Now(), StepCount: 1}
require.NoError(t, strategy.SaveSessionState(t.Context(), state))
name := checkpoint.ShadowBranchNameForCommit(state.BaseCommit, "")
file := filepath.Join(dir, ".git", session.SessionStateDirName, "protected.json")
require.NoError(t, os.WriteFile(file, []byte(`{"session_id":`), 0o600))
allowed, err := strategy.CanDeleteShadowBranch(t.Context(), name, "other")
require.Error(t, err)
require.False(t, allowed)
require.NoError(t, strategy.SaveSessionState(t.Context(), state))
allowed, err = strategy.CanDeleteShadowBranch(t.Context(), name, "other")
require.NoError(t, err)
require.False(t, allowed)
}
41 changes: 41 additions & 0 deletions cmd/entire/cli/gitrepo/ref_absence.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,41 @@
package gitrepo

import (
"errors"
"fmt"
"os"

"github.com/go-git/go-git/v6"
"github.com/go-git/go-git/v6/plumbing"
gitfilesystem "github.com/go-git/go-git/v6/storage/filesystem"
)

// ReferenceIsAbsent requires both loose storage and the reference reader to
// establish absence. A read failure must not authorize destructive cleanup.
func ReferenceIsAbsent(repo *git.Repository, refName plumbing.ReferenceName) (bool, error) {
if err := refName.Validate(); err != nil {
return false, fmt.Errorf("validate reference: %w", err)
}
// go-git falls back to packed refs after any loose-ref read error, so its
// not-found result alone cannot distinguish a directory from absence.
if storage, ok := repo.Storer.(*gitfilesystem.Storage); ok {
info, statErr := storage.Filesystem().Lstat(refName.String())
if statErr == nil {
if info.IsDir() {
return false, fmt.Errorf("ref %s is a directory", refName)
}
return false, nil
}
if !errors.Is(statErr, os.ErrNotExist) {
return false, fmt.Errorf("inspect ref %s: %w", refName, statErr)
}
}
_, err := repo.Reference(refName, false)
if errors.Is(err, plumbing.ErrReferenceNotFound) {
return true, nil
}
if err != nil {
return false, fmt.Errorf("read ref %s: %w", refName, err)
}
return false, nil
}
47 changes: 47 additions & 0 deletions cmd/entire/cli/gitrepo/ref_absence_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
package gitrepo

import (
"os"
"path/filepath"
"testing"

"github.com/entireio/cli/cmd/entire/cli/testutil/gitenv"
"github.com/go-git/go-git/v6/plumbing"
"github.com/stretchr/testify/require"
)

func TestReferenceIsAbsent(t *testing.T) {
t.Parallel()
for _, backend := range refCASBackends() {
t.Run(backend.name, func(t *testing.T) {
t.Parallel()
dir, _, _ := backend.init(t)
repo, err := OpenPath(dir)
require.NoError(t, err)
defer repo.Close()
absent, err := ReferenceIsAbsent(repo, plumbing.NewBranchReferenceName("main"))
require.NoError(t, err)
require.False(t, absent)
absent, err = ReferenceIsAbsent(repo, plumbing.NewBranchReferenceName("missing"))
require.NoError(t, err)
require.True(t, absent)
gitenv.Run(t, dir, "pack-refs", "--all")
absent, err = ReferenceIsAbsent(repo, plumbing.NewBranchReferenceName("main"))
require.NoError(t, err)
require.False(t, absent)
})
}
}

func TestReferenceIsAbsent_RefDirectoryIsNotAbsence(t *testing.T) {
t.Parallel()
dir, _, _ := initFilesRefCASRepo(t)
repo, err := OpenPath(dir)
require.NoError(t, err)
defer repo.Close()
ref := plumbing.NewBranchReferenceName("broken")
require.NoError(t, os.Mkdir(filepath.Join(dir, ".git", ref.String()), 0o700))
absent, err := ReferenceIsAbsent(repo, ref)
require.Error(t, err)
require.False(t, absent)
}
61 changes: 37 additions & 24 deletions cmd/entire/cli/gitrepo/ref_cas.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,14 +7,12 @@ import (
"errors"
"fmt"
"io"
"os"
"os/exec"
"strings"
"sync/atomic"
"time"

"github.com/go-git/go-git/v6/plumbing"
gitfilesystem "github.com/go-git/go-git/v6/storage/filesystem"
)

const refCASWaitDelay = 3 * time.Second
Expand All @@ -26,6 +24,8 @@ var (
ErrRefLocked = errors.New("git reference lock is unavailable")
// ErrRefSymbolic means the requested CAS target is a symbolic reference.
ErrRefSymbolic = errors.New("git reference is symbolic")
// ErrRefCASAbort means a prepared guarded update did not abort cleanly.
ErrRefCASAbort = errors.New("git reference transaction abort failed")
)

// CompareAndSwapRef atomically updates a direct ref through native Git and
Expand All @@ -36,6 +36,32 @@ func CompareAndSwapRef(
repoRoot string,
refName plumbing.ReferenceName,
newHash, expectedHash plumbing.Hash,
) error {
return compareAndSwapRef(ctx, repoRoot, refName, newHash, expectedHash, nil)
}

// CompareAndSwapRefGuarded atomically updates a direct ref only when guard
// succeeds. Git holds the prepared ref lock while guard runs, but the ref is
// not changed until after guard returns nil.
func CompareAndSwapRefGuarded(
ctx context.Context,
repoRoot string,
refName plumbing.ReferenceName,
newHash, expectedHash plumbing.Hash,
guard func() error,
) error {
if guard == nil {
return errors.New("compare-and-swap ref guard is nil")
}
return compareAndSwapRef(ctx, repoRoot, refName, newHash, expectedHash, guard)
}

func compareAndSwapRef(
ctx context.Context,
repoRoot string,
refName plumbing.ReferenceName,
newHash, expectedHash plumbing.Hash,
guard func() error,
) error {
tx, err := prepareRefCAS(ctx, repoRoot, refName, newHash, expectedHash)
if err != nil {
Expand All @@ -52,6 +78,14 @@ func CompareAndSwapRef(
tx.abort(),
)
}
if guard != nil {
if err := guard(); err != nil {
if abortErr := tx.abort(); abortErr != nil {
return errors.Join(err, fmt.Errorf("abort guarded ref update: %w", errors.Join(ErrRefCASAbort, abortErr)))
}
return err
}
}
return tx.commit()
}

Expand Down Expand Up @@ -234,28 +268,7 @@ func refIsAbsent(repoRoot string, refName plumbing.ReferenceName) (bool, error)
return false, fmt.Errorf("open repository to verify missing ref: %w", err)
}
defer repo.Close()
// go-git falls back to packed refs after any loose-ref read error, so its
// not-found result alone cannot distinguish a directory from absence.
if storage, ok := repo.Storer.(*gitfilesystem.Storage); ok {
info, statErr := storage.Filesystem().Lstat(refName.String())
if statErr == nil {
if info.IsDir() {
return false, fmt.Errorf("ref %s is a directory", refName)
}
return false, nil
}
if !errors.Is(statErr, os.ErrNotExist) {
return false, fmt.Errorf("inspect ref %s: %w", refName, statErr)
}
}
_, err = repo.Reference(refName, false)
if errors.Is(err, plumbing.ErrReferenceNotFound) {
return true, nil
}
if err != nil {
return false, fmt.Errorf("read ref %s: %w", refName, err)
}
return false, nil
return ReferenceIsAbsent(repo, refName)
}

func symbolicRefTarget(ctx context.Context, repoRoot string, refName plumbing.ReferenceName) (string, bool, error) {
Expand Down
37 changes: 37 additions & 0 deletions cmd/entire/cli/gitrepo/ref_cas_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package gitrepo

import (
"context"
"errors"
"os"
"os/exec"
"path/filepath"
Expand Down Expand Up @@ -90,6 +91,42 @@ func TestCompareAndSwapRef_RejectsSymbolicRef(t *testing.T) {
}
}

func TestCompareAndSwapRefGuarded_RejectsBeforeCommit(t *testing.T) {
t.Parallel()
for _, tt := range refCASBackends() {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
repoDir, initial, replacement := tt.init(t)
refName := plumbing.ReferenceName("refs/entire/guarded-delete")
gitenv.Run(t, repoDir, "update-ref", refName.String(), initial)
unsafe := errors.New("ref is still protected")

err := CompareAndSwapRefGuarded(
t.Context(),
repoDir,
refName,
plumbing.ZeroHash,
plumbing.NewHash(initial),
func() error {
cmd := exec.Command("git", "update-ref", refName.String(), replacement, initial) //nolint:noctx // must run while the guarded transaction is open
cmd.Dir = repoDir
cmd.Env = gitenv.Isolated()
output, updateErr := cmd.CombinedOutput()
require.Error(t, updateErr, "the guard must run after Git locks the ref")
require.Contains(t, strings.ToLower(string(output)), "lock")
return unsafe
},
)

require.ErrorIs(t, err, unsafe)
require.Equal(t, initial, strings.TrimSpace(gitenv.Run(t, repoDir, "rev-parse", refName.String())))

gitenv.Run(t, repoDir, "update-ref", refName.String(), replacement, initial)
require.Equal(t, replacement, strings.TrimSpace(gitenv.Run(t, repoDir, "rev-parse", refName.String())))
})
}
}

func TestPreparedRefCASPreventsConcurrentSymbolicConversion(t *testing.T) {
t.Parallel()
for _, tt := range refCASBackends() {
Expand Down
Loading
Loading