diff --git a/cmd/modify_test.go b/cmd/modify_test.go index c89eb623..0b313484 100644 --- a/cmd/modify_test.go +++ b/cmd/modify_test.go @@ -2,9 +2,13 @@ package cmd import ( "encoding/json" + "fmt" "io" "os" "path/filepath" + "strings" + "sync" + "sync/atomic" "testing" "time" @@ -109,6 +113,72 @@ func TestModifyStateAtomicWrite(t *testing.T) { assert.True(t, os.IsNotExist(err), "no .tmp file should remain after successful write") } +func TestModifyStateReadError(t *testing.T) { + dir := t.TempDir() + path := modify.StatePath(dir) + require.NoError(t, os.Mkdir(path, 0700)) + got, err := modify.LoadState(dir) + require.ErrorContains(t, err, "reading modify state") + assert.Nil(t, got) + var pathErr *os.PathError + require.ErrorAs(t, err, &pathErr) + assert.Equal(t, path, pathErr.Path) +} + +func TestModifyStateConcurrentReadWrite(t *testing.T) { + dir := t.TempDir() + state := &modify.StateFile{ + SchemaVersion: 1, Phase: modify.PhaseConflict, + OriginalBranch: "initial", ConflictBranch: "initial", + Snapshot: modify.Snapshot{StackMetadata: json.RawMessage("{}")}, + } + require.NoError(t, modify.SaveState(dir, state)) + stop := make(chan struct{}) + errs := make(chan error, 2) + var reads atomic.Int64 + var readers sync.WaitGroup + for range 2 { + readers.Go(func() { + for { + select { + case <-stop: + return + default: + } + got, err := modify.LoadState(dir) + if err != nil { + errs <- err + return + } + if got == nil || got.OriginalBranch == "" || got.OriginalBranch != got.ConflictBranch { + errs <- fmt.Errorf("reader observed an incomplete modify state") + return + } + reads.Add(1) + } + }) + } + var writeErr error + for i := range 50 { + branch := strings.Repeat(fmt.Sprintf("%04d", i), 8192) + state.OriginalBranch, state.ConflictBranch = branch, branch + if writeErr = modify.SaveState(dir, state); writeErr != nil { + break + } + } + close(stop) + readers.Wait() + close(errs) + require.NoError(t, writeErr) + for err := range errs { + require.NoError(t, err) + } + assert.Positive(t, reads.Load()) + got, err := modify.LoadState(dir) + require.NoError(t, err) + assert.Equal(t, state, got) +} + func TestCheckModifyStateGuard(t *testing.T) { t.Run("no state file", func(t *testing.T) { gitDir := t.TempDir() diff --git a/cmd/rebase.go b/cmd/rebase.go index 35e2d9da..bd552949 100644 --- a/cmd/rebase.go +++ b/cmd/rebase.go @@ -536,14 +536,14 @@ func saveRebaseState(gitDir string, state *rebaseState) error { if err != nil { return fmt.Errorf("error serializing rebase state: %w", err) } - if err := os.WriteFile(filepath.Join(gitDir, rebaseStateFile), data, 0644); err != nil { + if err := stack.WriteAtomic(filepath.Join(gitDir, rebaseStateFile), data); err != nil { return fmt.Errorf("error writing rebase state: %w", err) } return nil } func loadRebaseState(gitDir string) (*rebaseState, error) { - data, err := os.ReadFile(filepath.Join(gitDir, rebaseStateFile)) + data, err := stack.ReadStateFile(filepath.Join(gitDir, rebaseStateFile)) if err != nil { return nil, err } diff --git a/cmd/rebase_test.go b/cmd/rebase_test.go index 18b30523..9dd91266 100644 --- a/cmd/rebase_test.go +++ b/cmd/rebase_test.go @@ -9,6 +9,8 @@ import ( "os/exec" "path/filepath" "strings" + "sync" + "sync/atomic" "testing" "github.com/github/gh-stack/internal/config" @@ -953,6 +955,85 @@ func TestRebase_StateRoundTrip(t *testing.T) { assert.Equal(t, original.OntoOldBase, loaded.OntoOldBase) } +func TestRebase_StateReadErrors(t *testing.T) { + for _, kind := range []string{"missing", "directory", "invalid JSON"} { + t.Run(kind, func(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, rebaseStateFile) + switch kind { + case "directory": + require.NoError(t, os.Mkdir(path, 0700)) + case "invalid JSON": + require.NoError(t, os.WriteFile(path, []byte("{incomplete"), 0600)) + } + got, err := loadRebaseState(dir) + require.Error(t, err) + assert.Nil(t, got) + switch kind { + case "missing": + assert.ErrorIs(t, err, os.ErrNotExist) + case "directory": + var pathErr *os.PathError + require.ErrorAs(t, err, &pathErr) + assert.Equal(t, path, pathErr.Path) + case "invalid JSON": + var syntaxErr *json.SyntaxError + assert.ErrorAs(t, err, &syntaxErr) + } + }) + } +} + +func TestRebase_StateConcurrentReadWrite(t *testing.T) { + dir := t.TempDir() + state := &rebaseState{OriginalBranch: "initial", ConflictBranch: "initial"} + require.NoError(t, saveRebaseState(dir, state)) + stop := make(chan struct{}) + errs := make(chan error, 2) + var reads atomic.Int64 + var readers sync.WaitGroup + for range 2 { + readers.Go(func() { + for { + select { + case <-stop: + return + default: + } + got, err := loadRebaseState(dir) + if err != nil { + errs <- err + return + } + if got == nil || got.OriginalBranch == "" || got.OriginalBranch != got.ConflictBranch { + errs <- fmt.Errorf("reader observed an incomplete rebase state") + return + } + reads.Add(1) + } + }) + } + var writeErr error + for i := range 50 { + branch := strings.Repeat(fmt.Sprintf("%04d", i), 8192) + state.OriginalBranch, state.ConflictBranch = branch, branch + if writeErr = saveRebaseState(dir, state); writeErr != nil { + break + } + } + close(stop) + readers.Wait() + close(errs) + require.NoError(t, writeErr) + for err := range errs { + require.NoError(t, err) + } + assert.Positive(t, reads.Load()) + got, err := loadRebaseState(dir) + require.NoError(t, err) + assert.Equal(t, state, got) +} + // TestRebase_Continue_RebasesRemainingBranches verifies the --continue success // path: RebaseContinue is called, remaining branches are rebased via RebaseOnto, // the state file is cleaned up, and the original branch is restored. diff --git a/internal/git/gitops_test.go b/internal/git/gitops_test.go index dd6c68b9..66b42fe4 100644 --- a/internal/git/gitops_test.go +++ b/internal/git/gitops_test.go @@ -1435,7 +1435,23 @@ func TestIntegration_WorktreeRerereAutoContinuesMultipleCommits(t *testing.T) { require.Error(t, linked.Rebase("main", RebaseOpts{})) writeFile(t, path, "one.txt", "resolved one\n") require.NoError(t, linked.StageAll()) - require.Error(t, linked.RebaseContinue(RebaseOpts{}), "the second commit has an unseen conflict") + continueErr := linked.RebaseContinue(RebaseOpts{}) + t.Cleanup(func() { + if !t.Failed() { + return + } + t.Logf("first continuation: %v", continueErr) + stateDir, err := linked.GitDir() + require.NoError(t, err) + for _, name := range []string{"message", "author-script", "amend", "stopped-sha", "msgnum"} { + data, readErr := os.ReadFile(filepath.Join(stateDir, "rebase-merge", name)) + t.Logf("rebase-merge/%s: %q (read error: %v)", name, data, readErr) + } + }) + require.Error(t, continueErr, "the second commit has an unseen conflict") + conflicts, err := linked.ConflictedFiles() + require.NoError(t, err) + require.Equal(t, []string{"two.txt"}, conflicts, "unexpected second stop: %v", continueErr) writeFile(t, path, "two.txt", "resolved two\n") require.NoError(t, linked.StageAll()) require.NoError(t, linked.RebaseContinue(RebaseOpts{})) diff --git a/internal/modify/state.go b/internal/modify/state.go index 2660cba2..2bf2dc3b 100644 --- a/internal/modify/state.go +++ b/internal/modify/state.go @@ -7,6 +7,8 @@ import ( "os" "path/filepath" "time" + + "github.com/github/gh-stack/internal/stack" ) const stateFileName = "gh-stack-modify-state" @@ -75,7 +77,7 @@ func StatePath(gitDir string) string { // LoadState reads the modify state file from the git directory. // Returns nil, nil if the file does not exist. func LoadState(gitDir string) (*StateFile, error) { - data, err := os.ReadFile(StatePath(gitDir)) + data, err := stack.ReadStateFile(StatePath(gitDir)) if err != nil { if errors.Is(err, os.ErrNotExist) { return nil, nil @@ -96,18 +98,9 @@ func SaveState(gitDir string, state *StateFile) error { if err != nil { return fmt.Errorf("marshaling modify state: %w", err) } - target := StatePath(gitDir) - tmp := target + ".tmp" - if err := os.WriteFile(tmp, data, 0644); err != nil { + if err := stack.WriteAtomic(StatePath(gitDir), data); err != nil { return fmt.Errorf("writing modify state: %w", err) } - // Remove existing target before rename for Windows compatibility - // (os.Rename fails on Windows if the target already exists). - _ = os.Remove(target) - if err := os.Rename(tmp, target); err != nil { - _ = os.Remove(tmp) - return fmt.Errorf("committing modify state: %w", err) - } return nil } diff --git a/internal/stack/atomic.go b/internal/stack/atomic.go new file mode 100644 index 00000000..cafda41e --- /dev/null +++ b/internal/stack/atomic.go @@ -0,0 +1,72 @@ +package stack + +import ( + "errors" + "fmt" + "os" + "path/filepath" +) + +// ReadStateFile reads a complete state file, allowing WriteAtomic to replace it +// while the read is in progress. On Windows, its handle permits delete sharing. +func ReadStateFile(path string) ([]byte, error) { + return readStateFile(path) +} + +// WriteAtomic publishes data at path using a fully written temporary file in +// the same directory. It preserves an existing regular file's permissions and +// uses 0644 for a new file. The parent directory must already exist. +// It does not acquire locks; callers must serialize mutations. +func WriteAtomic(path string, data []byte) error { + return writeFileAtomic(path, data, 0644, true) +} + +// writeFileAtomic publishes a fully written sibling of path. When replace is +// false, an existing destination is never overwritten, even on a racing create. +func writeFileAtomic(path string, data []byte, mode os.FileMode, replace bool) (err error) { + if replace { + info, statErr := os.Lstat(path) + switch { + case statErr == nil: + if !info.Mode().IsRegular() { + return fmt.Errorf("cannot replace non-regular file %q", path) + } + mode = info.Mode().Perm() + case !errors.Is(statErr, os.ErrNotExist): + return statErr + } + } + + f, err := os.CreateTemp(filepath.Dir(path), "."+filepath.Base(path)+"-*") + if err != nil { + return err + } + temp := f.Name() + closed := false + defer func() { + if !closed { + err = errors.Join(err, f.Close()) + } + if removeErr := os.Remove(temp); removeErr != nil && !errors.Is(removeErr, os.ErrNotExist) { + err = errors.Join(err, fmt.Errorf("removing temporary file: %w", removeErr)) + } + }() + + if err := f.Chmod(mode); err != nil { + return err + } + if _, err := f.Write(data); err != nil { + return err + } + if err := f.Sync(); err != nil { + return err + } + closed = true + if err := f.Close(); err != nil { + return err + } + if err := publishFile(temp, path, replace); err != nil { + return err + } + return syncDirectory(filepath.Dir(path)) +} diff --git a/internal/stack/atomic_unix.go b/internal/stack/atomic_unix.go new file mode 100644 index 00000000..3a2f5c2d --- /dev/null +++ b/internal/stack/atomic_unix.go @@ -0,0 +1,29 @@ +//go:build !windows + +package stack + +import ( + "errors" + "os" +) + +func readStateFile(path string) ([]byte, error) { + return os.ReadFile(path) +} + +func publishFile(temp, path string, replace bool) error { + if replace { + return os.Rename(temp, path) + } + // Linking publishes a complete file without replacing an existing backup. + // The temporary name is removed by writeFileAtomic. + return os.Link(temp, path) +} + +func syncDirectory(path string) error { + f, err := os.Open(path) + if err != nil { + return err + } + return errors.Join(f.Sync(), f.Close()) +} diff --git a/internal/stack/atomic_windows.go b/internal/stack/atomic_windows.go new file mode 100644 index 00000000..fc8319d7 --- /dev/null +++ b/internal/stack/atomic_windows.go @@ -0,0 +1,110 @@ +//go:build windows + +package stack + +import ( + "errors" + "io" + "os" + "path/filepath" + "strings" + "unsafe" + + "golang.org/x/sys/windows" +) + +func readStateFile(path string) ([]byte, error) { + name, err := windowsFilePath(path) + if err != nil { + return nil, err + } + // Let publication replace the name while readers finish with the old file. + handle, err := windows.CreateFile(&name[0], windows.GENERIC_READ, + windows.FILE_SHARE_READ|windows.FILE_SHARE_WRITE|windows.FILE_SHARE_DELETE, + nil, windows.OPEN_EXISTING, windows.FILE_ATTRIBUTE_NORMAL, 0) + if err != nil { + return nil, &os.PathError{Op: "open", Path: path, Err: err} + } + f := os.NewFile(uintptr(handle), path) + data, err := io.ReadAll(f) + return data, errors.Join(err, f.Close()) +} + +func publishFile(temp, path string, replace bool) error { + from, err := windowsFilePath(temp) + if err != nil { + return err + } + to, err := windowsFilePath(path) + if err != nil { + return err + } + // Never remove the destination first, or allow a cross-volume copy/delete. + if replace { + err = replaceFileWindows(from, to) + } else { + err = windows.MoveFileEx(&from[0], &to[0], windows.MOVEFILE_WRITE_THROUGH) + } + if err != nil { + return &os.LinkError{Op: "publish", Old: temp, New: path, Err: err} + } + return nil +} + +func replaceFileWindows(from, to []uint16) error { + handle, err := windows.CreateFile(&from[0], windows.DELETE|windows.GENERIC_WRITE, + windows.FILE_SHARE_READ|windows.FILE_SHARE_WRITE|windows.FILE_SHARE_DELETE, + nil, windows.OPEN_EXISTING, windows.FILE_ATTRIBUTE_NORMAL|windows.FILE_FLAG_WRITE_THROUGH, 0) + if err != nil { + return err + } + + // FILE_RENAME_INFO has pointer-sized padding and a trailing UTF-16 name. + type renameInfo struct { + Flags uint32 + RootDirectory windows.Handle + FileNameLength uint32 + FileName [1]uint16 + } + buffer := make([]byte, int(unsafe.Sizeof(renameInfo{}))+len(to)*2) + info := (*renameInfo)(unsafe.Pointer(&buffer[0])) + info.Flags = windows.FILE_RENAME_REPLACE_IF_EXISTS | windows.FILE_RENAME_POSIX_SEMANTICS + info.FileNameLength = uint32(len(to)-1) * 2 + copy(unsafe.Slice(&info.FileName[0], len(to)), to) + + err = windows.SetFileInformationByHandle(handle, windows.FileRenameInfoEx, &buffer[0], uint32(len(buffer))) + if err == nil { + return errors.Join(windows.FlushFileBuffers(handle), windows.CloseHandle(handle)) + } + if closeErr := windows.CloseHandle(handle); closeErr != nil { + return errors.Join(err, closeErr) + } + + // Older Windows versions/filesystems may lack POSIX rename. Only fall back + // for unsupported operations; legacy rename still reports open-reader errors. + if errors.Is(err, windows.ERROR_NOT_SUPPORTED) || errors.Is(err, windows.ERROR_INVALID_PARAMETER) || + errors.Is(err, windows.ERROR_INVALID_FUNCTION) { + return windows.MoveFileEx(&from[0], &to[0], windows.MOVEFILE_REPLACE_EXISTING|windows.MOVEFILE_WRITE_THROUGH) + } + return err +} + +func windowsFilePath(path string) ([]uint16, error) { + path, err := filepath.Abs(path) + if err != nil { + return nil, err + } + switch { + case strings.HasPrefix(path, `\\?\`), strings.HasPrefix(path, `\\.\`): + case strings.HasPrefix(path, `\\`): + path = `\\?\UNC\` + path[2:] + default: + path = `\\?\` + path + } + return windows.UTF16FromString(path) +} + +func syncDirectory(string) error { + // Windows has no directory fsync; publication flushes the file or uses WRITE_THROUGH. + return nil +} diff --git a/internal/stack/atomic_windows_test.go b/internal/stack/atomic_windows_test.go new file mode 100644 index 00000000..8ced582a --- /dev/null +++ b/internal/stack/atomic_windows_test.go @@ -0,0 +1,103 @@ +//go:build windows + +package stack + +import ( + "io" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/sys/windows" +) + +func TestAtomicPublication_HeldWindowsReader(t *testing.T) { + for _, tt := range []struct { + name string + replace bool + }{ + {"replace", true}, + {"exclusive", false}, + } { + t.Run(tt.name, func(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "state with spaces") + original, replacement := []byte("original recovery state"), []byte("new state") + require.NoError(t, WriteAtomic(path, original)) + name, err := windowsFilePath(path) + require.NoError(t, err) + handle, err := windows.CreateFile(&name[0], windows.GENERIC_READ, + windows.FILE_SHARE_READ|windows.FILE_SHARE_WRITE|windows.FILE_SHARE_DELETE, + nil, windows.OPEN_EXISTING, windows.FILE_ATTRIBUTE_NORMAL, 0) + require.NoError(t, err) + reader := os.NewFile(uintptr(handle), path) + t.Cleanup(func() { assert.NoError(t, reader.Close()) }) + + err = writeFileAtomic(path, replacement, 0644, tt.replace) + if tt.replace { + require.NoError(t, err) + } else { + require.ErrorIs(t, err, os.ErrExist) + } + old, err := io.ReadAll(reader) + require.NoError(t, err) + assert.Equal(t, original, old, "the held reader must retain the original file") + current, err := readStateFile(path) + require.NoError(t, err) + if tt.replace { + assert.Equal(t, replacement, current, "new readers must see the published file") + } else { + assert.Equal(t, original, current, "exclusive publication must not replace the target") + } + temps, err := filepath.Glob(filepath.Join(dir, ".state with spaces-*")) + require.NoError(t, err) + assert.Empty(t, temps) + }) + } +} + +func TestReadStateFile_WindowsSharing(t *testing.T) { + for _, shareDelete := range []bool{false, true} { + name := "exclusive handle" + if shareDelete { + name = "shared delete handle" + } + t.Run(name, func(t *testing.T) { + path := filepath.Join(t.TempDir(), "state") + original, replacement := []byte("original state"), []byte("new state") + require.NoError(t, WriteAtomic(path, original)) + name, err := windowsFilePath(path) + require.NoError(t, err) + var shareMode uint32 + if shareDelete { + shareMode = windows.FILE_SHARE_READ | windows.FILE_SHARE_WRITE | windows.FILE_SHARE_DELETE + } + handle, err := windows.CreateFile(&name[0], windows.GENERIC_READ|windows.DELETE, + shareMode, nil, windows.OPEN_EXISTING, windows.FILE_ATTRIBUTE_NORMAL, 0) + require.NoError(t, err) + reader := os.NewFile(uintptr(handle), path) + t.Cleanup(func() { assert.NoError(t, reader.Close()) }) + + // An existing DELETE-access handle requires new readers to share delete. + _, err = os.ReadFile(path) + require.ErrorIs(t, err, windows.ERROR_SHARING_VIOLATION) + got, err := ReadStateFile(path) + if !shareDelete { + require.ErrorIs(t, err, windows.ERROR_SHARING_VIOLATION) + return + } + require.NoError(t, err) + assert.Equal(t, original, got) + + require.NoError(t, WriteAtomic(path, replacement)) + got, err = ReadStateFile(path) + require.NoError(t, err) + assert.Equal(t, replacement, got) + old, err := io.ReadAll(reader) + require.NoError(t, err) + assert.Equal(t, original, old) + }) + } +} diff --git a/internal/stack/lock.go b/internal/stack/lock.go index 02a9dd94..bf8e643d 100644 --- a/internal/stack/lock.go +++ b/internal/stack/lock.go @@ -1,15 +1,19 @@ package stack import ( + "errors" "fmt" "os" "path/filepath" "time" ) -const lockFileName = "gh-stack.lock" +const ( + lockFileName = "gh-stack.lock" + operationLockFileName = "gh-stack-operation.lock" +) -// LockError is returned when the stack file lock cannot be acquired. +// LockError is returned when the catalog or operation lock times out. // Callers can check for this with errors.As to distinguish lock failures // from other errors. type LockError struct { @@ -29,16 +33,13 @@ type StaleError struct { func (e *StaleError) Error() string { return e.Err.Error() } func (e *StaleError) Unwrap() error { return e.Err } -// LockTimeout is how long Lock() will wait for the exclusive lock before -// giving up. With the lock held only during file writes (milliseconds), -// this timeout primarily guards against a hung process holding the lock. +// LockTimeout is how long Lock and LockOperation wait for an exclusive lock. var LockTimeout = 5 * time.Second // lockRetryInterval is the sleep between non-blocking lock attempts. const lockRetryInterval = 100 * time.Millisecond -// FileLock provides an exclusive advisory lock on the stack file to prevent -// concurrent writes between multiple gh-stack processes. +// FileLock provides an exclusive advisory catalog or operation lock. type FileLock struct { f *os.File } @@ -48,29 +49,50 @@ type FileLock struct { // // Most callers should not use Lock directly — stack.Save() acquires the lock // automatically. Use Lock only when you need to hold the lock across multiple -// operations (e.g. Load-Modify-Save as an atomic unit). +// operations (e.g. Load-Modify-SaveWithLock as an atomic unit). func Lock(gitDir string) (*FileLock, error) { - path := filepath.Join(gitDir, lockFileName) + lock, _, err := acquireLock(filepath.Join(gitDir, lockFileName), "stack", true) + return lock, err +} + +// LockOperation provides a separate operation lock in the given directory. +// Callers using it must acquire it before loading mutation state or taking the +// catalog lock. Save may be used while this lock is held. +func LockOperation(commonDir string) (*FileLock, error) { + lock, _, err := acquireLock(filepath.Join(commonDir, operationLockFileName), "stack operation", true) + return lock, err +} + +// TryLockOperation attempts to acquire the operation lock without waiting. +// A false result with no error means contention; other failures are returned. +func TryLockOperation(commonDir string) (*FileLock, bool, error) { + return acquireLock(filepath.Join(commonDir, operationLockFileName), "stack operation", false) +} + +func acquireLock(path, name string, wait bool) (*FileLock, bool, error) { f, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR, 0644) if err != nil { - return nil, fmt.Errorf("opening lock file: %w", err) + return nil, false, fmt.Errorf("opening lock file: %w", err) } deadline := time.Now().Add(LockTimeout) for { err := tryLockFile(f) if err == nil { - return &FileLock{f: f}, nil + return &FileLock{f: f}, true, nil } if !isLockBusy(err) { - // Unexpected error (e.g. bad fd) — don't retry. - f.Close() - return nil, fmt.Errorf("locking stack file: %w", err) + return nil, false, fmt.Errorf("locking %s file: %w", name, errors.Join(err, f.Close())) + } + if !wait { + if err := f.Close(); err != nil { + return nil, false, fmt.Errorf("closing lock file: %w", err) + } + return nil, false, nil } if time.Now().After(deadline) { - f.Close() - return nil, &LockError{Err: fmt.Errorf( - "timed out waiting for stack lock after %s — another gh-stack process may be running", LockTimeout)} + return nil, false, &LockError{Err: errors.Join(fmt.Errorf( + "timed out waiting for %s lock after %s — another gh-stack process may be running", name, LockTimeout), f.Close())} } time.Sleep(lockRetryInterval) } @@ -86,4 +108,5 @@ func (l *FileLock) Unlock() { } unlockFile(l.f) l.f.Close() + l.f = nil } diff --git a/internal/stack/lock_test.go b/internal/stack/lock_test.go index 6afefc60..13cca1c2 100644 --- a/internal/stack/lock_test.go +++ b/internal/stack/lock_test.go @@ -5,7 +5,10 @@ import ( "fmt" "os" "path/filepath" + "runtime" + "strings" "sync" + "sync/atomic" "testing" "time" @@ -242,3 +245,386 @@ func TestSave_DoubleSaveSucceeds(t *testing.T) { require.NoError(t, err) assert.Len(t, final.Stacks, 2) } + +func TestLockOperation_IndependentCatalogLock(t *testing.T) { + dir := t.TempDir() + operation, err := LockOperation(dir) + require.NoError(t, err) + defer operation.Unlock() + + sf, err := Load(dir) + require.NoError(t, err) + sf.AddStack(makeStack("main", "feature")) + require.NoError(t, Save(dir, sf), "Save must not retake the operation lock") + + catalog, err := Lock(dir) + require.NoError(t, err) + defer catalog.Unlock() + sf.AddStack(makeStack("main", "other")) + require.NoError(t, SaveWithLock(dir, sf, catalog)) + + start := time.Now() + contender, acquired, err := TryLockOperation(dir) + require.NoError(t, err) + assert.False(t, acquired) + assert.Nil(t, contender) + assert.Less(t, time.Since(start), time.Second) + + operation.Unlock() + operation.Unlock() + contender, acquired, err = TryLockOperation(dir) + require.NoError(t, err) + require.True(t, acquired, "catalog lock must not prevent an operation lock") + contender.Unlock() + + assert.FileExists(t, filepath.Join(dir, operationLockFileName)) + assert.FileExists(t, filepath.Join(dir, lockFileName)) +} + +func TestTryLockOperation_Errors(t *testing.T) { + for _, kind := range []string{"missing directory", "directory at lock path"} { + t.Run(kind, func(t *testing.T) { + dir := t.TempDir() + if kind == "missing directory" { + dir = filepath.Join(dir, "missing") + } else { + require.NoError(t, os.Mkdir(filepath.Join(dir, operationLockFileName), 0755)) + } + lock, acquired, err := TryLockOperation(dir) + require.Error(t, err) + assert.Nil(t, lock) + assert.False(t, acquired) + var lockErr *LockError + assert.False(t, errors.As(err, &lockErr), "real I/O failure is not contention") + }) + } + + t.Run("invalid descriptor is not contention", func(t *testing.T) { + f, err := os.CreateTemp(t.TempDir(), "lock") + require.NoError(t, err) + require.NoError(t, f.Close()) + err = tryLockFile(f) + require.Error(t, err) + assert.False(t, isLockBusy(err)) + }) +} + +func TestLockOperation_TimesOut(t *testing.T) { + dir := t.TempDir() + lock, err := LockOperation(dir) + require.NoError(t, err) + defer lock.Unlock() + originalTimeout := LockTimeout + LockTimeout = 100 * time.Millisecond + defer func() { LockTimeout = originalTimeout }() + + other, err := LockOperation(dir) + require.Error(t, err) + assert.Nil(t, other) + var lockErr *LockError + require.ErrorAs(t, err, &lockErr) + assert.Contains(t, err.Error(), "stack operation lock") +} + +func TestLockOperation_SerializesMutationSnapshots(t *testing.T) { + dir := t.TempDir() + errs := make(chan error, 4) + var wg sync.WaitGroup + for i := range 4 { + wg.Go(func() { + lock, err := LockOperation(dir) + if err != nil { + errs <- err + return + } + defer lock.Unlock() + sf, err := Load(dir) + if err != nil { + errs <- err + return + } + sf.AddStack(makeStack("main", fmt.Sprintf("branch-%d", i))) + errs <- Save(dir, sf) + }) + } + wg.Wait() + close(errs) + for err := range errs { + require.NoError(t, err) + } + sf, err := Load(dir) + require.NoError(t, err) + assert.Len(t, sf.Stacks, 4) +} + +func TestSaveNonBlocking_OperationAndCatalogGuards(t *testing.T) { + for _, kind := range []string{"uncontended", "operation lock", "catalog lock", "stale", "pending migration"} { + t.Run(kind, func(t *testing.T) { + dir := t.TempDir() + sf := &StackFile{Stacks: []Stack{makeStack("main", "original")}} + require.NoError(t, Save(dir, sf)) + refresh, err := Load(dir) + require.NoError(t, err) + refresh.Stacks[0].Branches[0].Head = "metadata-refresh" + checksum := append([]byte(nil), refresh.loadChecksum...) + + var held *FileLock + switch kind { + case "operation lock": + held, err = LockOperation(dir) + case "catalog lock": + held, err = Lock(dir) + case "stale": + sf.Stacks[0].Branches[0].Head = "critical-write" + err = Save(dir, sf) + case "pending migration": + err = os.WriteFile(filepath.Join(dir, migrationFileName), []byte("{}"), 0600) + } + require.NoError(t, err) + defer held.Unlock() + + start := time.Now() + SaveNonBlocking(dir, refresh) + assert.Less(t, time.Since(start), time.Second) + held.Unlock() + got, err := Load(dir) + require.NoError(t, err) + if kind == "uncontended" { + assert.Equal(t, "metadata-refresh", got.Stacks[0].Branches[0].Head) + assert.NotEqual(t, checksum, refresh.loadChecksum) + } else { + assert.NotEqual(t, "metadata-refresh", got.Stacks[0].Branches[0].Head) + assert.Equal(t, checksum, refresh.loadChecksum) + } + lock, acquired, err := TryLockOperation(dir) + require.NoError(t, err) + require.True(t, acquired, "a skipped refresh must release the operation lock") + lock.Unlock() + }) + } +} + +func TestSave_AtomicReaderVisibility(t *testing.T) { + dir := t.TempDir() + sf := &StackFile{Stacks: []Stack{makeStack("main", "feature")}} + require.NoError(t, Save(dir, sf)) + stop := make(chan struct{}) + errs := make(chan error, 2) + var reads atomic.Int64 + var readers sync.WaitGroup + for range 2 { + readers.Go(func() { + for { + select { + case <-stop: + return + default: + } + got, err := Load(dir) + if err != nil { + errs <- err + return + } + if len(got.Stacks) != 1 || len(got.Stacks[0].Branches) != 1 || + got.Stacks[0].Trunk.Head != got.Stacks[0].Branches[0].Base { + errs <- fmt.Errorf("reader observed an incomplete catalog: %#v", got) + return + } + reads.Add(1) + } + }) + } + var writeErr error + for i := range 50 { + head := strings.Repeat(fmt.Sprintf("%04d", i), 1024) + sf.Stacks[0].Trunk.Head = head + sf.Stacks[0].Branches[0].Base = head + if writeErr = Save(dir, sf); writeErr != nil { + break + } + } + close(stop) + readers.Wait() + close(errs) + require.NoError(t, writeErr) + for err := range errs { + require.NoError(t, err) + } + assert.Positive(t, reads.Load()) + temps, err := filepath.Glob(filepath.Join(dir, "."+stackFileName+"-*")) + require.NoError(t, err) + assert.Empty(t, temps) +} + +func TestAtomicPublication_PreservesExistingFiles(t *testing.T) { + t.Run("no overwrite on exclusive publication", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "backup") + require.NoError(t, writeFileAtomic(path, []byte("original"), 0600, false)) + err := writeFileAtomic(path, []byte("replacement"), 0600, false) + require.ErrorIs(t, err, os.ErrExist) + data, err := os.ReadFile(path) + require.NoError(t, err) + assert.Equal(t, "original", string(data)) + temps, err := filepath.Glob(filepath.Join(filepath.Dir(path), ".backup-*")) + require.NoError(t, err) + assert.Empty(t, temps) + }) + + t.Run("failed replace keeps destination", func(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "destination") + require.NoError(t, os.WriteFile(path, []byte("original"), 0600)) + require.Error(t, publishFile(filepath.Join(dir, "missing"), path, true)) + data, err := os.ReadFile(path) + require.NoError(t, err) + assert.Equal(t, "original", string(data)) + }) + + t.Run("non-regular destination is not replaced", func(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "destination") + require.NoError(t, os.Mkdir(path, 0755)) + require.Error(t, writeFileAtomic(path, []byte("replacement"), 0600, true)) + assert.DirExists(t, path) + }) + + t.Run("mode survives replacement", func(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("Unix permission bits") + } + dir := t.TempDir() + sf := &StackFile{Stacks: []Stack{makeStack("main", "feature")}} + require.NoError(t, Save(dir, sf)) + require.NoError(t, os.Chmod(stackFilePath(dir), 0640)) + sf.AddStack(makeStack("main", "other")) + require.NoError(t, Save(dir, sf)) + info, err := os.Stat(stackFilePath(dir)) + require.NoError(t, err) + assert.Equal(t, os.FileMode(0640), info.Mode().Perm()) + }) +} + +func TestReadStateFile(t *testing.T) { + t.Run("reads complete bytes", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "state") + for _, want := range [][]byte{ + {}, + {0x00, 0xff, 0x0a}, + []byte(strings.Repeat("state\x00\xff\n", 16384)), + } { + require.NoError(t, os.WriteFile(path, want, 0600)) + got, err := ReadStateFile(path) + require.NoError(t, err) + assert.Equal(t, want, got) + } + }) + + t.Run("missing file", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "missing") + _, err := ReadStateFile(path) + require.ErrorIs(t, err, os.ErrNotExist) + assert.NoFileExists(t, path) + }) + + t.Run("directory", func(t *testing.T) { + path := t.TempDir() + _, err := ReadStateFile(path) + var pathErr *os.PathError + require.ErrorAs(t, err, &pathErr) + assert.Equal(t, path, pathErr.Path) + assert.DirExists(t, path) + }) + + t.Run("unreadable file", func(t *testing.T) { + if runtime.GOOS == "windows" || os.Geteuid() == 0 { + t.Skip("requires Unix permission enforcement") + } + path := filepath.Join(t.TempDir(), "state") + require.NoError(t, os.WriteFile(path, []byte("private state"), 0600)) + require.NoError(t, os.Chmod(path, 0)) + t.Cleanup(func() { assert.NoError(t, os.Chmod(path, 0600)) }) + _, err := ReadStateFile(path) + require.ErrorIs(t, err, os.ErrPermission) + }) +} + +func TestWriteAtomic(t *testing.T) { + t.Run("creates and replaces exact bytes", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "gh-stack-rebase-state") + for _, data := range [][]byte{ + []byte(`{"phase":"rebasing","branch":"feature"}`), + []byte(`{"phase":"done"}`), + {0x00, 0xff, 0x0a}, + {}, + } { + require.NoError(t, WriteAtomic(path, data)) + got, err := os.ReadFile(path) + require.NoError(t, err) + assert.Equal(t, data, got) + } + temps, err := filepath.Glob(filepath.Join(filepath.Dir(path), ".gh-stack-rebase-state-*")) + require.NoError(t, err) + assert.Empty(t, temps) + }) + + t.Run("missing parent is reported", func(t *testing.T) { + parent := filepath.Join(t.TempDir(), "missing") + require.ErrorIs(t, WriteAtomic(filepath.Join(parent, "state"), []byte("{}")), os.ErrNotExist) + assert.NoDirExists(t, parent) + }) + + t.Run("non-regular target is preserved", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "state") + require.NoError(t, os.Mkdir(path, 0755)) + require.Error(t, WriteAtomic(path, []byte("{}"))) + assert.DirExists(t, path) + }) +} + +func TestSave_FailedPublicationPreservesChecksum(t *testing.T) { + if runtime.GOOS == "windows" || os.Geteuid() == 0 { + t.Skip("requires Unix directory permission enforcement") + } + dir := t.TempDir() + sf := &StackFile{Stacks: []Stack{makeStack("main", "original")}} + require.NoError(t, Save(dir, sf)) + checksum := append([]byte(nil), sf.loadChecksum...) + before, err := os.ReadFile(stackFilePath(dir)) + require.NoError(t, err) + info, err := os.Stat(dir) + require.NoError(t, err) + require.NoError(t, os.Chmod(dir, 0500)) + t.Cleanup(func() { assert.NoError(t, os.Chmod(dir, info.Mode().Perm())) }) + + sf.AddStack(makeStack("main", "unsaved")) + require.Error(t, Save(dir, sf)) + assert.Equal(t, checksum, sf.loadChecksum) + after, err := os.ReadFile(stackFilePath(dir)) + require.NoError(t, err) + assert.Equal(t, before, after) + require.NoError(t, os.Chmod(dir, info.Mode().Perm())) + require.NoError(t, Save(dir, sf), "a failed publication must remain retryable") +} + +func TestMigrateLegacyState_TakesCatalogLock(t *testing.T) { + dir := t.TempDir() + catalogs := []migrationCatalog{ + writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", migrationTestFile(makeStack("main", "linked"))), + } + operation, err := LockOperation(dir) + require.NoError(t, err) + defer operation.Unlock() + catalog, err := Lock(dir) + require.NoError(t, err) + defer catalog.Unlock() + originalTimeout := LockTimeout + LockTimeout = 0 + defer func() { LockTimeout = originalTimeout }() + + var lockErr *LockError + require.ErrorAs(t, MigrateLegacyState(dir), &lockErr) + assertMigrationOriginals(t, dir, catalogs, false) + catalog.Unlock() + require.NoError(t, MigrateLegacyState(dir)) + assertMigrationOriginals(t, dir, catalogs, true) +} diff --git a/internal/stack/migration.go b/internal/stack/migration.go new file mode 100644 index 00000000..a8e5ff14 --- /dev/null +++ b/internal/stack/migration.go @@ -0,0 +1,486 @@ +package stack + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "os" + "path/filepath" + "reflect" + "slices" + "strings" +) + +const ( + migrationFileName = "gh-stack-migration" + migrationBackupSuffix = ".pre-worktree-migration" + migrationVersion = 1 +) + +// MigrationConflictError reports catalog definitions or migration artifacts +// that cannot be reconciled without choosing which tracking state to keep. +type MigrationConflictError struct { + Sources []string + Branches []string + Reason string +} + +func (e *MigrationConflictError) Error() string { + branches := "" + if len(e.Branches) > 0 { + branches = fmt.Sprintf("; branches: %q", e.Branches) + } + return fmt.Sprintf("cannot migrate stack catalogs: %s; sources: %q%s; reconcile or recreate the intended tracking state before retrying", e.Reason, e.Sources, branches) +} + +// MigrationBlockedError identifies legacy recovery records whose original +// catalogs must remain available to the operation's original worktree. +type MigrationBlockedError struct { + RecoveryPaths []string +} + +func (e *MigrationBlockedError) Error() string { + return fmt.Sprintf("stack migration is blocked by recovery state at %q; finish or abort rebase/modify in the original worktree before retrying migration", e.RecoveryPaths) +} + +type migrationCatalog struct { + Path string `json:"path"` + Data []byte `json:"data"` + Mode os.FileMode `json:"mode"` +} + +// Original bytes, including the old common catalog, must survive publication +// before any named backup is created. Relative administration paths also keep +// this record usable if the repository itself moves during an interruption. +type migrationState struct { + Version int `json:"version"` + Catalogs []migrationCatalog `json:"catalogs"` +} + +// HasLegacyState reports linked-worktree catalogs or an unfinished migration. +// It inspects every retained administration directory, regardless of whether +// its worktree still exists. It never invokes Git or changes worktree state. +func HasLegacyState(commonDir string) (bool, error) { + catalogs, err := legacyCatalogs(commonDir) + if err != nil { + return false, err + } + _, _, pending, err := readMigrationFile(filepath.Join(commonDir, migrationFileName)) + if err != nil { + return false, err + } + return pending || len(catalogs) > 0, nil +} + +// MigrateLegacyState consolidates legacy catalogs into the common directory. +// Command paths do not call this helper yet; catalog locations are unchanged. +// The caller must hold LockOperation; this function takes the catalog lock. +// Only disjoint stacks and equivalent duplicates are merged. Original bytes +// are retained in *.pre-worktree-migration backups after common publication. +// Old and new gh-stack versions must not write catalogs concurrently. +func MigrateLegacyState(commonDir string) error { + lock, err := Lock(commonDir) + if err != nil { + return err + } + defer lock.Unlock() + + legacy, err := legacyCatalogs(commonDir) + if err != nil { + return err + } + journalPath := filepath.Join(commonDir, migrationFileName) + journal, _, pending, err := readMigrationFile(journalPath) + if err != nil { + return err + } + if !pending && len(legacy) == 0 { + return nil + } + + state := migrationState{Version: migrationVersion} + if pending { + state.Version = 0 + if err := decodeMigrationJSON(journal, &state); err != nil { + return fmt.Errorf("reading migration journal %q: %w", journalPath, err) + } + if state.Version != migrationVersion { + return fmt.Errorf("migration journal %q has unsupported version %d; preserve it and use the matching gh-stack version", journalPath, state.Version) + } + } else { + data, mode, exists, err := readMigrationFile(stackFilePath(commonDir)) + if err != nil { + return err + } + if exists { + state.Catalogs = append(state.Catalogs, migrationCatalog{Path: stackFileName, Data: data, Mode: mode}) + } + state.Catalogs = append(state.Catalogs, legacy...) + } + if err := validateMigrationPaths(state.Catalogs); err != nil { + return fmt.Errorf("invalid migration journal %q: %w", journalPath, err) + } + if err := checkMigrationRecovery(commonDir, state.Catalogs); err != nil { + return err + } + merged, err := mergeMigrationCatalogs(commonDir, state.Catalogs) + if err != nil { + return err + } + mergedData, err := marshalStackFile(merged) + if err != nil { + return err + } + + published, err := checkMigrationSnapshot(commonDir, state.Catalogs, legacy, mergedData) + if err != nil { + return err + } + if !pending { + data, err := json.MarshalIndent(state, "", " ") + if err != nil { + return fmt.Errorf("encoding migration journal: %w", err) + } + if err := writeFileAtomic(journalPath, data, 0600, false); err != nil { + return fmt.Errorf("publishing migration journal: %w", err) + } + } + if !published || !pending { + if err := writeStackFile(commonDir, merged); err != nil { + return err + } + } + + for _, catalog := range state.Catalogs { + path := filepath.Join(commonDir, filepath.FromSlash(catalog.Path)) + if err := backupMigrationCatalog(path, catalog); err != nil { + return err + } + if catalog.Path == stackFileName { + continue + } + data, _, exists, err := readMigrationFile(path) + if err != nil { + return err + } + if !exists { + continue // A prior attempt archived it; its backup was just checked. + } + if !bytes.Equal(data, catalog.Data) { + return migrationConflict("legacy catalog changed during migration", []string{path, journalPath}, nil) + } + if err := os.Remove(path); err != nil { + return fmt.Errorf("archiving legacy catalog %q: %w", path, err) + } + if err := syncDirectory(filepath.Dir(path)); err != nil { + return fmt.Errorf("syncing archived catalog directory: %w", err) + } + } + if err := os.Remove(journalPath); err != nil { + return fmt.Errorf("removing completed migration journal: %w", err) + } + return syncDirectory(commonDir) +} + +func legacyCatalogs(commonDir string) ([]migrationCatalog, error) { + info, err := os.Stat(commonDir) + if err != nil { + return nil, fmt.Errorf("inspecting common directory: %w", err) + } + if !info.IsDir() { + return nil, fmt.Errorf("common directory %q is not a directory", commonDir) + } + dir := filepath.Join(commonDir, "worktrees") + info, err = os.Lstat(dir) + if errors.Is(err, os.ErrNotExist) { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("inspecting worktree administration directory: %w", err) + } + if !info.IsDir() { + return nil, fmt.Errorf("worktree administration path %q is not a directory", dir) + } + entries, err := os.ReadDir(dir) + if err != nil { + return nil, fmt.Errorf("listing worktree administration directories: %w", err) + } + var catalogs []migrationCatalog + for _, entry := range entries { + path := filepath.Join(dir, entry.Name()) + info, err := entry.Info() + if err != nil { + return nil, fmt.Errorf("inspecting worktree administration path %q: %w", path, err) + } + if !info.IsDir() { + return nil, fmt.Errorf("worktree administration path %q is not a directory", path) + } + data, mode, exists, err := readMigrationFile(filepath.Join(path, stackFileName)) + if err != nil { + return nil, err + } + if exists { + catalogs = append(catalogs, migrationCatalog{ + Path: filepath.ToSlash(filepath.Join("worktrees", entry.Name(), stackFileName)), + Data: data, + Mode: mode, + }) + } + } + return catalogs, nil +} + +func readMigrationFile(path string) ([]byte, os.FileMode, bool, error) { + info, err := os.Lstat(path) + if errors.Is(err, os.ErrNotExist) { + return nil, 0, false, nil + } + if err != nil { + return nil, 0, false, fmt.Errorf("inspecting migration source %q: %w", path, err) + } + if !info.Mode().IsRegular() { + return nil, 0, false, fmt.Errorf("migration source %q is not a regular file", path) + } + data, err := readStateFile(path) + if err != nil { + return nil, 0, false, fmt.Errorf("reading migration source %q: %w", path, err) + } + return data, info.Mode().Perm(), true, nil +} + +func validateMigrationPaths(catalogs []migrationCatalog) error { + seen := make(map[string]bool) + hasLegacy := false + for _, catalog := range catalogs { + parts := strings.Split(catalog.Path, "/") + if catalog.Path != stackFileName { + if len(parts) != 3 || parts[0] != "worktrees" || parts[2] != stackFileName || + parts[1] == "" || parts[1] == "." || parts[1] == ".." || + filepath.Base(parts[1]) != parts[1] { + return fmt.Errorf("unsafe catalog path %q", catalog.Path) + } + hasLegacy = true + } + if !filepath.IsLocal(filepath.FromSlash(catalog.Path)) || seen[catalog.Path] { + return fmt.Errorf("invalid or duplicate catalog path %q", catalog.Path) + } + if catalog.Mode != catalog.Mode.Perm() { + return fmt.Errorf("invalid catalog permissions for %q", catalog.Path) + } + seen[catalog.Path] = true + } + if !hasLegacy { + return errors.New("migration contains no linked-worktree catalogs") + } + return nil +} + +func checkMigrationRecovery(commonDir string, catalogs []migrationCatalog) error { + dirs := []string{commonDir} + for _, catalog := range catalogs { + if catalog.Path != stackFileName { + dirs = append(dirs, filepath.Dir(filepath.Join(commonDir, filepath.FromSlash(catalog.Path)))) + } + } + var paths []string + for _, dir := range dirs { + for _, name := range []string{"gh-stack-rebase-state", "gh-stack-modify-state"} { + path := filepath.Join(dir, name) + _, err := os.Lstat(path) + if err == nil { + paths = append(paths, path) + } else if !errors.Is(err, os.ErrNotExist) { + return fmt.Errorf("inspecting legacy recovery state %q: %w", path, err) + } + } + } + if len(paths) > 0 { + return &MigrationBlockedError{RecoveryPaths: paths} + } + return nil +} + +func decodeMigrationJSON(data []byte, value any) error { + if !json.Valid(data) { + return errors.New("invalid JSON") + } + decoder := json.NewDecoder(bytes.NewReader(data)) + decoder.DisallowUnknownFields() + return decoder.Decode(value) +} + +func parseMigrationCatalog(catalog migrationCatalog) (*StackFile, error) { + sf, err := parseStackFile(catalog.Data) + if err != nil { + return nil, err + } + // Refuse data that the normal catalog model would silently discard. + if err := decodeMigrationJSON(catalog.Data, sf); err != nil { + return nil, err + } + var catalogFields struct { + SchemaVersion *int `json:"schemaVersion"` + Stacks *[]json.RawMessage `json:"stacks"` + } + if err := json.Unmarshal(catalog.Data, &catalogFields); err != nil { + return nil, err + } + if catalogFields.SchemaVersion == nil || *catalogFields.SchemaVersion < 0 || catalogFields.Stacks == nil { + return nil, errors.New("catalog must contain schemaVersion and a stacks array") + } + for i, s := range sf.Stacks { + var stackFields map[string]json.RawMessage + if err := json.Unmarshal((*catalogFields.Stacks)[i], &stackFields); err != nil { + return nil, err + } + if s.Trunk.Branch == "" || stackFields["branches"] == nil { + return nil, errors.New("each stack must contain a named trunk and branches") + } + owned := map[string]bool{s.Trunk.Branch: true} + for _, branch := range s.Branches { + if branch.Branch == "" { + return nil, errors.New("stack contains an unnamed branch") + } + if owned[branch.Branch] { + return nil, &MigrationConflictError{Branches: []string{branch.Branch}, Reason: "branch is repeated within one stack"} + } + owned[branch.Branch] = true + } + } + return sf, nil +} + +func mergeMigrationCatalogs(commonDir string, catalogs []migrationCatalog) (*StackFile, error) { + merged := &StackFile{SchemaVersion: schemaVersion, Stacks: []Stack{}} + var sources []string + repositorySource := "" + for _, catalog := range catalogs { + path := filepath.Join(commonDir, filepath.FromSlash(catalog.Path)) + sf, err := parseMigrationCatalog(catalog) + if err != nil { + var conflict *MigrationConflictError + if errors.As(err, &conflict) { + conflict.Sources = []string{path} + } + return nil, fmt.Errorf("parsing migration catalog %q: %w", path, err) + } + if sf.Repository != "" { + if merged.Repository != "" && merged.Repository != sf.Repository { + return nil, migrationConflict("repository identities differ", []string{repositorySource, path}, nil) + } + merged.Repository = sf.Repository + repositorySource = path + } + for _, s := range sf.Stacks { + duplicate := false + for i, other := range merged.Stacks { + if equivalentMigrationStacks(s, other) { + duplicate = true + break + } + var shared []string + for _, branch := range s.Branches { + if other.IndexOf(branch.Branch) >= 0 { + shared = append(shared, branch.Branch) + } + } + sameIdentity := (s.ID != "" && s.ID == other.ID) || (s.Number != 0 && s.Number == other.Number) + if sameIdentity { + branches := append(s.BranchNames(), other.BranchNames()...) + return nil, migrationConflict("stack identity has differing definitions", []string{sources[i], path}, branches) + } + if len(shared) > 0 { + return nil, migrationConflict("non-trunk branches belong to differing stack definitions", []string{sources[i], path}, shared) + } + } + if !duplicate { + merged.Stacks = append(merged.Stacks, s) + sources = append(sources, path) + } + } + } + return merged, nil +} + +func equivalentMigrationStacks(a, b Stack) bool { + if len(a.Branches) == 0 { + a.Branches = nil + } + if len(b.Branches) == 0 { + b.Branches = nil + } + return reflect.DeepEqual(a, b) +} + +func migrationConflict(reason string, sources, branches []string) error { + slices.Sort(branches) + branches = slices.Compact(branches) + return &MigrationConflictError{Sources: sources, Branches: branches, Reason: reason} +} + +func checkMigrationSnapshot(commonDir string, catalogs, legacy []migrationCatalog, mergedData []byte) (bool, error) { + expected := make(map[string]migrationCatalog, len(catalogs)) + for _, catalog := range catalogs { + expected[catalog.Path] = catalog + } + for _, catalog := range legacy { + if _, ok := expected[catalog.Path]; !ok { + return false, migrationConflict("a new legacy catalog appeared during migration", []string{filepath.Join(commonDir, filepath.FromSlash(catalog.Path)), filepath.Join(commonDir, migrationFileName)}, nil) + } + } + + commonPath := stackFilePath(commonDir) + current, _, exists, err := readMigrationFile(commonPath) + if err != nil { + return false, err + } + published := exists && bytes.Equal(current, mergedData) + original, hadCommon := expected[stackFileName] + if !published && (exists != hadCommon || !bytes.Equal(current, original.Data)) { + return false, migrationConflict("common catalog changed during migration", []string{commonPath, filepath.Join(commonDir, migrationFileName)}, nil) + } + + for _, catalog := range catalogs { + path := filepath.Join(commonDir, filepath.FromSlash(catalog.Path)) + backup, _, backedUp, err := readMigrationFile(path + migrationBackupSuffix) + if err != nil { + return false, err + } + if backedUp && !bytes.Equal(backup, catalog.Data) { + return false, migrationConflict("existing backup differs from the original catalog", []string{path, path + migrationBackupSuffix}, nil) + } + if catalog.Path == stackFileName { + continue + } + data, _, exists, err := readMigrationFile(path) + if err != nil { + return false, err + } + if exists && !bytes.Equal(data, catalog.Data) { + return false, migrationConflict("legacy catalog changed during migration", []string{path, filepath.Join(commonDir, migrationFileName)}, nil) + } + if !exists && (!published || !backedUp) { + return false, migrationConflict("legacy catalog disappeared without a published common catalog and matching backup", []string{path, path + migrationBackupSuffix}, nil) + } + } + return published, nil +} + +func backupMigrationCatalog(path string, catalog migrationCatalog) error { + backupPath := path + migrationBackupSuffix + data, _, exists, err := readMigrationFile(backupPath) + if err != nil { + return err + } + if exists { + if !bytes.Equal(data, catalog.Data) { + return migrationConflict("existing backup differs from the original catalog", []string{path, backupPath}, nil) + } + return nil + } + if err := writeFileAtomic(backupPath, catalog.Data, catalog.Mode, false); err != nil { + return fmt.Errorf("preserving original catalog %q: %w", backupPath, err) + } + return nil +} diff --git a/internal/stack/stack.go b/internal/stack/stack.go index c03205ab..518a5963 100644 --- a/internal/stack/stack.go +++ b/internal/stack/stack.go @@ -316,7 +316,7 @@ func stackFilePath(gitDir string) string { // Save can detect concurrent modifications. func Load(gitDir string) (*StackFile, error) { path := stackFilePath(gitDir) - data, err := os.ReadFile(path) + data, err := readStateFile(path) if err != nil { if errors.Is(err, os.ErrNotExist) { // loadChecksum stays nil — sentinel for "file absent at load time". @@ -328,6 +328,10 @@ func Load(gitDir string) (*StackFile, error) { return nil, fmt.Errorf("reading stack file: %w", err) } + return parseStackFile(data) +} + +func parseStackFile(data []byte) (*StackFile, error) { var sf StackFile if err := json.Unmarshal(data, &sf); err != nil { return nil, fmt.Errorf("parsing stack file: %w", err) @@ -345,6 +349,8 @@ func Load(gitDir string) (*StackFile, error) { // Save acquires an exclusive lock on the stack file, verifies the file hasn't // been modified since Load (optimistic concurrency), writes sf as JSON, and // releases the lock. The lock is held only for the read-compare-write window. +// Callers may hold the separate operation lock across Load/preflight/Save; +// Save itself only acquires the catalog lock. // Returns *LockError if the lock times out, or *StaleError if another process // modified the file since it was loaded. func Save(gitDir string, sf *StackFile) error { @@ -360,7 +366,8 @@ func Save(gitDir string, sf *StackFile) error { return writeStackFile(gitDir, sf) } -// SaveWithLock writes the stack file while the caller already holds the lock. +// SaveWithLock writes the stack file while the caller already holds the catalog +// lock (not merely the operation lock). // The caller is responsible for acquiring and releasing the lock. // Panics if lock is nil to catch programming errors. func SaveWithLock(gitDir string, sf *StackFile, lock *FileLock) error { @@ -370,22 +377,27 @@ func SaveWithLock(gitDir string, sf *StackFile, lock *FileLock) error { return writeStackFile(gitDir, sf) } -// SaveNonBlocking attempts to save without blocking. If another process holds -// the lock or the file was modified since Load, the save is silently skipped. -// Use this for best-effort metadata persistence (e.g. syncing PR state in view). +// SaveNonBlocking attempts a best-effort metadata refresh without waiting for +// either the operation or catalog lock. Contention, stale data and I/O errors +// skip the refresh. Use Save for critical writes, including callers already +// holding the operation lock. func SaveNonBlocking(gitDir string, sf *StackFile) { - path := filepath.Join(gitDir, lockFileName) - f, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR, 0644) - if err != nil { + operation, acquired, err := TryLockOperation(gitDir) + if err != nil || !acquired { return } - if tryLockFile(f) != nil { - f.Close() + defer operation.Unlock() + + lock, acquired, err := acquireLock(filepath.Join(gitDir, lockFileName), "stack", false) + if err != nil || !acquired { return } - lock := &FileLock{f: f} defer lock.Unlock() + // Do not change a published migration snapshot before its backups finish. + if _, err := os.Lstat(filepath.Join(gitDir, migrationFileName)); !errors.Is(err, os.ErrNotExist) { + return + } if checkStale(gitDir, sf) != nil { return } @@ -397,7 +409,7 @@ func SaveNonBlocking(gitDir string, sf *StackFile) { // by another process. The caller must hold the lock. func checkStale(gitDir string, sf *StackFile) error { path := stackFilePath(gitDir) - data, err := os.ReadFile(path) + data, err := readStateFile(path) if errors.Is(err, os.ErrNotExist) { // File absent on disk. @@ -428,16 +440,11 @@ func checkStale(gitDir string, sf *StackFile) error { } func writeStackFile(gitDir string, sf *StackFile) error { - sf.SchemaVersion = schemaVersion - if sf.Stacks == nil { - sf.Stacks = []Stack{} - } - data, err := json.MarshalIndent(sf, "", " ") + data, err := marshalStackFile(sf) if err != nil { - return fmt.Errorf("marshaling stack file: %w", err) + return err } - path := stackFilePath(gitDir) - if err := os.WriteFile(path, data, 0644); err != nil { + if err := WriteAtomic(stackFilePath(gitDir), data); err != nil { return fmt.Errorf("writing stack file: %w", err) } // Refresh checksum so a second Save on the same StackFile doesn't @@ -446,3 +453,15 @@ func writeStackFile(gitDir string, sf *StackFile) error { sf.loadChecksum = sum[:] return nil } + +func marshalStackFile(sf *StackFile) ([]byte, error) { + sf.SchemaVersion = schemaVersion + if sf.Stacks == nil { + sf.Stacks = []Stack{} + } + data, err := json.MarshalIndent(sf, "", " ") + if err != nil { + return nil, fmt.Errorf("marshaling stack file: %w", err) + } + return data, nil +} diff --git a/internal/stack/stack_test.go b/internal/stack/stack_test.go index 73ddc408..22d5a581 100644 --- a/internal/stack/stack_test.go +++ b/internal/stack/stack_test.go @@ -2,8 +2,11 @@ package stack import ( "encoding/json" + "errors" + "fmt" "os" "path/filepath" + "runtime" "testing" "github.com/stretchr/testify/assert" @@ -656,3 +659,711 @@ func TestNearestSurvivingBranch(t *testing.T) { }) } } + +func migrationTestFile(stacks ...Stack) StackFile { + if stacks == nil { + stacks = []Stack{} + } + return StackFile{SchemaVersion: schemaVersion, Repository: "github.com:owner/repo", Stacks: stacks} +} + +func migrationTestStack() Stack { + return Stack{ + ID: "stack-global-id", + Number: 17, + Trunk: BranchRef{Branch: "main", Head: "trunk-head", Base: "trunk-base"}, + Branches: []BranchRef{ + { + Branch: "feature/one", Head: "first-head", Base: "first-base", + PullRequest: &PullRequestRef{Number: 21, ID: "PR_one", URL: "https://example.com/pull/21", Merged: true}, + }, + { + Branch: "feature/two", Head: "second-head", Base: "second-base", + PullRequest: &PullRequestRef{Number: 22, ID: "PR_two", URL: "https://example.com/pull/22"}, + }, + }, + } +} + +func writeMigrationTestData(t *testing.T, dir, relative string, data []byte) migrationCatalog { + t.Helper() + path := filepath.Join(dir, filepath.FromSlash(relative)) + require.NoError(t, os.MkdirAll(filepath.Dir(path), 0755)) + require.NoError(t, os.WriteFile(path, data, 0644)) + info, err := os.Stat(path) + require.NoError(t, err) + return migrationCatalog{Path: relative, Data: data, Mode: info.Mode().Perm()} +} + +func writeMigrationTestCatalog(t *testing.T, dir, relative string, sf StackFile) migrationCatalog { + t.Helper() + data, err := json.MarshalIndent(sf, "", " ") + require.NoError(t, err) + return writeMigrationTestData(t, dir, relative, append(data, '\n')) +} + +func migrateWithTestOperationLock(t *testing.T, dir string) error { + t.Helper() + lock, err := LockOperation(dir) + require.NoError(t, err) + defer lock.Unlock() + return MigrateLegacyState(dir) +} + +func assertMigrationOriginals(t *testing.T, dir string, catalogs []migrationCatalog, migrated bool) { + t.Helper() + for _, catalog := range catalogs { + path := filepath.Join(dir, filepath.FromSlash(catalog.Path)) + if migrated { + if catalog.Path != stackFileName { + assert.NoFileExists(t, path) + } + path += migrationBackupSuffix + } else { + assert.NoFileExists(t, path+migrationBackupSuffix) + } + data, err := os.ReadFile(path) + require.NoError(t, err) + assert.Equal(t, catalog.Data, data, "original bytes at %s", path) + } + assert.NoFileExists(t, filepath.Join(dir, migrationFileName)) +} + +func TestMigrateLegacyState_Merge(t *testing.T) { + detailed := migrationTestStack() + other := makeStack("main", "other") + other.Trunk.Head = "a-different-recorded-trunk-head" + empty := makeStack("main") + emptySlice := Stack{Trunk: BranchRef{Branch: "main"}, Branches: []BranchRef{}} + tests := []struct { + name string + common []Stack + legacy [][]Stack + want []Stack + }{ + {"absent common catalog", nil, [][]Stack{{detailed}, {other}}, []Stack{detailed, other}}, + {"disjoint common and linked catalogs", []Stack{detailed}, [][]Stack{{other}}, []Stack{detailed, other}}, + {"equivalent common and linked catalogs", []Stack{detailed}, [][]Stack{{detailed}}, []Stack{detailed}}, + {"equivalent linked catalogs", nil, [][]Stack{{detailed}, {detailed}}, []Stack{detailed}}, + {"deduplicate individual stacks", []Stack{detailed}, [][]Stack{{other, detailed}, {other}}, []Stack{detailed, other}}, + {"shared trunks with differing recorded heads", nil, [][]Stack{{detailed}, {other}}, []Stack{detailed, other}}, + {"a member can be another stacks trunk", []Stack{makeStack("main", "base")}, [][]Stack{{makeStack("base", "top")}}, []Stack{makeStack("main", "base"), makeStack("base", "top")}}, + {"equivalent empty branch lists", []Stack{empty}, [][]Stack{{emptySlice}}, []Stack{empty}}, + {"empty legacy catalog", nil, [][]Stack{{}}, []Stack{}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dir := t.TempDir() + var catalogs []migrationCatalog + if tt.common != nil { + catalogs = append(catalogs, writeMigrationTestCatalog(t, dir, stackFileName, migrationTestFile(tt.common...))) + } + for i, stacks := range tt.legacy { + relative := fmt.Sprintf("worktrees/admin id %d/gh-stack", i) + catalogs = append(catalogs, writeMigrationTestCatalog(t, dir, relative, migrationTestFile(stacks...))) + } + has, err := HasLegacyState(dir) + require.NoError(t, err) + require.True(t, has) + require.NoError(t, migrateWithTestOperationLock(t, dir)) + + got, err := Load(dir) + require.NoError(t, err) + assert.Equal(t, schemaVersion, got.SchemaVersion) + assert.Equal(t, "github.com:owner/repo", got.Repository) + assert.Equal(t, tt.want, got.Stacks) + assertMigrationOriginals(t, dir, catalogs, true) + if tt.common == nil { + assert.NoFileExists(t, stackFilePath(dir)+migrationBackupSuffix) + } + first, err := os.ReadFile(stackFilePath(dir)) + require.NoError(t, err) + var fields map[string]json.RawMessage + require.NoError(t, json.Unmarshal(first, &fields)) + assert.Len(t, fields, 3, "the shared catalog retains its ordinary schema") + has, err = HasLegacyState(dir) + require.NoError(t, err) + assert.False(t, has) + require.NoError(t, migrateWithTestOperationLock(t, dir)) + second, err := os.ReadFile(stackFilePath(dir)) + require.NoError(t, err) + assert.Equal(t, first, second) + assertMigrationOriginals(t, dir, catalogs, true) + }) + } +} + +func TestMigrateLegacyState_Conflicts(t *testing.T) { + tests := []struct { + name string + change func(*Stack) + }{ + {"same ID different branches", func(s *Stack) { s.Branches = []BranchRef{{Branch: "different"}} }}, + {"same number different ID and branches", func(s *Stack) { + s.ID = "different-id" + s.Branches = []BranchRef{{Branch: "different"}} + }}, + {"same branches different identities", func(s *Stack) { s.ID, s.Number = "different-id", 18 }}, + {"missing versus known identity", func(s *Stack) { s.ID, s.Number = "", 0 }}, + {"different recorded base", func(s *Stack) { s.Branches[0].Base = "new-base" }}, + {"different recorded head", func(s *Stack) { s.Branches[0].Head = "new-head" }}, + {"different trunk head", func(s *Stack) { s.Trunk.Head = "new-trunk-head" }}, + {"different trunk base", func(s *Stack) { s.Trunk.Base = "new-trunk-base" }}, + {"different PR number", func(s *Stack) { s.Branches[0].PullRequest.Number++ }}, + {"different PR ID", func(s *Stack) { s.Branches[0].PullRequest.ID = "new-PR-id" }}, + {"different PR URL", func(s *Stack) { s.Branches[0].PullRequest.URL = "https://example.com/new" }}, + {"different merge state", func(s *Stack) { s.Branches[0].PullRequest.Merged = false }}, + {"longer stack must not win", func(s *Stack) { s.Branches = append(s.Branches, BranchRef{Branch: "extra"}) }}, + {"different order", func(s *Stack) { s.Branches[0], s.Branches[1] = s.Branches[1], s.Branches[0] }}, + } + for _, tt := range tests { + for _, withCommon := range []bool{false, true} { + t.Run(fmt.Sprintf("%s/common=%t", tt.name, withCommon), func(t *testing.T) { + dir := t.TempDir() + first := migrationTestStack() + second := migrationTestStack() + tt.change(&second) + firstPath := "worktrees/first/gh-stack" + if withCommon { + firstPath = stackFileName + } + catalogs := []migrationCatalog{ + writeMigrationTestCatalog(t, dir, firstPath, migrationTestFile(first)), + writeMigrationTestCatalog(t, dir, "worktrees/second/gh-stack", migrationTestFile(second)), + } + err := migrateWithTestOperationLock(t, dir) + var conflict *MigrationConflictError + require.ErrorAs(t, err, &conflict) + assert.ElementsMatch(t, []string{filepath.Join(dir, filepath.FromSlash(firstPath)), filepath.Join(dir, "worktrees", "second", stackFileName)}, conflict.Sources) + assert.Contains(t, conflict.Branches, "feature/one") + assert.Contains(t, err.Error(), "reconcile or recreate") + assertMigrationOriginals(t, dir, catalogs, false) + }) + } + } + + t.Run("overlapping local stacks without IDs", func(t *testing.T) { + dir := t.TempDir() + catalogs := []migrationCatalog{ + writeMigrationTestCatalog(t, dir, stackFileName, migrationTestFile(makeStack("main", "shared", "one"))), + writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", migrationTestFile(makeStack("main", "shared", "two"))), + } + var conflict *MigrationConflictError + require.ErrorAs(t, migrateWithTestOperationLock(t, dir), &conflict) + assert.Equal(t, []string{"shared"}, conflict.Branches) + assertMigrationOriginals(t, dir, catalogs, false) + }) + + t.Run("duplicate branch within one stack", func(t *testing.T) { + dir := t.TempDir() + catalogs := []migrationCatalog{ + writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", migrationTestFile(makeStack("main", "duplicate", "duplicate"))), + } + var conflict *MigrationConflictError + require.ErrorAs(t, migrateWithTestOperationLock(t, dir), &conflict) + assert.Equal(t, []string{"duplicate"}, conflict.Branches) + assert.Equal(t, []string{filepath.Join(dir, "worktrees", "linked", stackFileName)}, conflict.Sources) + assertMigrationOriginals(t, dir, catalogs, false) + }) +} + +func TestMigrateLegacyState_RepositoryIdentity(t *testing.T) { + for _, repositories := range [][2]string{ + {"github.com:owner/repo", "github.com:other/repo"}, + {"", "github.com:owner/repo"}, + {"github.com:owner/repo", ""}, + {"", ""}, + } { + t.Run(fmt.Sprintf("%q and %q", repositories[0], repositories[1]), func(t *testing.T) { + dir := t.TempDir() + first, second := migrationTestFile(makeStack("main", "one")), migrationTestFile(makeStack("main", "two")) + first.Repository, second.Repository = repositories[0], repositories[1] + catalogs := []migrationCatalog{ + writeMigrationTestCatalog(t, dir, stackFileName, first), + writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", second), + } + err := migrateWithTestOperationLock(t, dir) + if repositories[0] != "" && repositories[1] != "" && repositories[0] != repositories[1] { + var conflict *MigrationConflictError + require.ErrorAs(t, err, &conflict) + assert.Contains(t, conflict.Reason, "repository") + assertMigrationOriginals(t, dir, catalogs, false) + return + } + require.NoError(t, err) + sf, err := Load(dir) + require.NoError(t, err) + want := repositories[0] + if want == "" { + want = repositories[1] + } + assert.Equal(t, want, sf.Repository) + assertMigrationOriginals(t, dir, catalogs, true) + }) + } +} + +func TestMigrateLegacyState_InvalidCatalogs(t *testing.T) { + tests := []struct { + name string + data string + }{ + {"malformed JSON", "{not JSON"}, + {"null catalog", "null"}, + {"array catalog", "[]"}, + {"future schema", `{"schemaVersion":999,"stacks":[]}`}, + {"negative schema", `{"schemaVersion":-1,"stacks":[]}`}, + {"missing schema", `{"stacks":[]}`}, + {"missing stacks", `{"schemaVersion":1}`}, + {"null stacks", `{"schemaVersion":1,"stacks":null}`}, + {"unknown metadata", `{"schemaVersion":1,"stacks":[],"futureMetadata":true}`}, + {"unnamed trunk", `{"schemaVersion":1,"stacks":[{"trunk":{},"branches":[]}]}`}, + {"missing branches", `{"schemaVersion":1,"stacks":[{"trunk":{"branch":"main"}}]}`}, + {"unnamed branch", `{"schemaVersion":1,"stacks":[{"trunk":{"branch":"main"},"branches":[{}]}]}`}, + {"unknown branch metadata", `{"schemaVersion":1,"stacks":[{"trunk":{"branch":"main"},"branches":[{"branch":"a","unknown":"value"}]}]}`}, + {"invalid PR type", `{"schemaVersion":1,"stacks":[{"trunk":{"branch":"main"},"branches":[{"branch":"a","pullRequest":{"number":"21"}}]}]}`}, + } + for _, tt := range tests { + for _, invalidCommon := range []bool{false, true} { + t.Run(fmt.Sprintf("%s/common=%t", tt.name, invalidCommon), func(t *testing.T) { + dir := t.TempDir() + invalid, valid := "worktrees/invalid/gh-stack", stackFileName + if invalidCommon { + invalid, valid = stackFileName, "worktrees/valid/gh-stack" + } + catalogs := []migrationCatalog{ + writeMigrationTestData(t, dir, invalid, []byte(tt.data)), + writeMigrationTestCatalog(t, dir, valid, migrationTestFile(makeStack("main", "valid"))), + } + err := migrateWithTestOperationLock(t, dir) + require.Error(t, err) + assert.Contains(t, err.Error(), fmt.Sprintf("%q", filepath.Join(dir, filepath.FromSlash(invalid)))) + assertMigrationOriginals(t, dir, catalogs, false) + }) + } + } + + t.Run("older schema keeps existing load compatibility", func(t *testing.T) { + dir := t.TempDir() + catalog := writeMigrationTestData(t, dir, "worktrees/older/gh-stack", []byte(`{"schemaVersion":0,"stacks":[{"trunk":{"branch":"main"},"branches":[{"branch":"old"}]}]}`)) + require.NoError(t, migrateWithTestOperationLock(t, dir)) + sf, err := Load(dir) + require.NoError(t, err) + assert.Equal(t, schemaVersion, sf.SchemaVersion) + assert.Equal(t, []Stack{makeStack("main", "old")}, sf.Stacks) + assertMigrationOriginals(t, dir, []migrationCatalog{catalog}, true) + }) +} + +func TestMigrateLegacyState_RecoveryBlocks(t *testing.T) { + for _, location := range []string{".", "worktrees/linked"} { + for _, name := range []string{"gh-stack-rebase-state", "gh-stack-modify-state"} { + t.Run(location+"/"+name, func(t *testing.T) { + dir := t.TempDir() + catalogs := []migrationCatalog{ + writeMigrationTestCatalog(t, dir, stackFileName, migrationTestFile(makeStack("main", "common"))), + writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", migrationTestFile(makeStack("main", "linked"))), + } + recoveryPath := filepath.Join(dir, filepath.FromSlash(location), name) + require.NoError(t, os.WriteFile(recoveryPath, []byte("{even a damaged recovery record blocks migration"), 0600)) + err := migrateWithTestOperationLock(t, dir) + var blocked *MigrationBlockedError + require.ErrorAs(t, err, &blocked) + assert.Equal(t, []string{recoveryPath}, blocked.RecoveryPaths) + assert.Contains(t, err.Error(), "original worktree") + assertMigrationOriginals(t, dir, catalogs, false) + assert.FileExists(t, recoveryPath) + require.NoError(t, os.Remove(recoveryPath)) + require.NoError(t, migrateWithTestOperationLock(t, dir)) + assertMigrationOriginals(t, dir, catalogs, true) + }) + } + } + + t.Run("common journal blocks even with no common catalog", func(t *testing.T) { + dir := t.TempDir() + catalog := writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", migrationTestFile(makeStack("main", "linked"))) + require.NoError(t, os.WriteFile(filepath.Join(dir, "gh-stack-rebase-state"), []byte("{}"), 0600)) + var blocked *MigrationBlockedError + require.ErrorAs(t, migrateWithTestOperationLock(t, dir), &blocked) + assert.NoFileExists(t, stackFilePath(dir)) + assertMigrationOriginals(t, dir, []migrationCatalog{catalog}, false) + }) + + t.Run("unrelated worktree without a catalog is not touched", func(t *testing.T) { + dir := t.TempDir() + writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", migrationTestFile(makeStack("main", "linked"))) + recovery := writeMigrationTestData(t, dir, "worktrees/unrelated/gh-stack-modify-state", []byte("{}")) + require.NoError(t, migrateWithTestOperationLock(t, dir)) + data, err := os.ReadFile(filepath.Join(dir, filepath.FromSlash(recovery.Path))) + require.NoError(t, err) + assert.Equal(t, recovery.Data, data) + }) +} + +func TestHasLegacyState_RetainedAdministrationDirectories(t *testing.T) { + dir := t.TempDir() + has, err := HasLegacyState(dir) + require.NoError(t, err) + assert.False(t, has) + common := writeMigrationTestCatalog(t, dir, stackFileName, migrationTestFile(makeStack("main", "common"))) + has, err = HasLegacyState(dir) + require.NoError(t, err) + assert.False(t, has) + require.NoError(t, migrateWithTestOperationLock(t, dir)) + assertMigrationOriginals(t, dir, []migrationCatalog{common}, false) + + legacy := writeMigrationTestCatalog(t, dir, "worktrees/unrelated-admin-id/gh-stack", migrationTestFile(makeStack("main", "linked"))) + adminDir := filepath.Dir(filepath.Join(dir, filepath.FromSlash(legacy.Path))) + gitdir := []byte(filepath.Join(dir, "missing worktree with different basename", ".git") + "\n") + require.NoError(t, os.WriteFile(filepath.Join(adminDir, "gitdir"), gitdir, 0644)) + require.NoError(t, os.WriteFile(filepath.Join(adminDir, "locked"), []byte("retained worktree"), 0644)) + has, err = HasLegacyState(dir) + require.NoError(t, err) + assert.True(t, has) + require.NoError(t, migrateWithTestOperationLock(t, dir)) + assertMigrationOriginals(t, dir, []migrationCatalog{common, legacy}, true) + data, err := os.ReadFile(filepath.Join(adminDir, "gitdir")) + require.NoError(t, err) + assert.Equal(t, gitdir, data) + assert.FileExists(t, filepath.Join(adminDir, "locked")) + assert.DirExists(t, adminDir) + has, err = HasLegacyState(dir) + require.NoError(t, err) + assert.False(t, has, "archived catalogs do not trigger re-import") + + require.NoError(t, os.WriteFile(filepath.Join(dir, migrationFileName), []byte("{}"), 0600)) + has, err = HasLegacyState(dir) + require.NoError(t, err) + assert.True(t, has, "a pending journal must be detected even after every catalog was archived") +} + +func TestLegacyState_MetadataErrors(t *testing.T) { + tests := []struct { + name string + setup func(*testing.T, string) + }{ + {"worktrees is a file", func(t *testing.T, dir string) { + require.NoError(t, os.WriteFile(filepath.Join(dir, "worktrees"), []byte("not a directory"), 0600)) + }}, + {"admin entry is a file", func(t *testing.T, dir string) { + writeMigrationTestData(t, dir, "worktrees/not-a-directory", []byte("metadata")) + }}, + {"catalog is a directory", func(t *testing.T, dir string) { + require.NoError(t, os.MkdirAll(filepath.Join(dir, "worktrees", "linked", stackFileName), 0755)) + }}, + {"catalog is a dangling symlink", func(t *testing.T, dir string) { + admin := filepath.Join(dir, "worktrees", "linked") + require.NoError(t, os.MkdirAll(admin, 0755)) + if err := os.Symlink(filepath.Join(dir, "missing"), filepath.Join(admin, stackFileName)); err != nil { + if runtime.GOOS == "windows" { + t.Skipf("symlinks unavailable: %v", err) + } + require.NoError(t, err) + } + }}, + {"symlinked admin is not silently skipped", func(t *testing.T, dir string) { + target := filepath.Join(dir, "target") + require.NoError(t, os.MkdirAll(target, 0755)) + require.NoError(t, os.MkdirAll(filepath.Join(dir, "worktrees"), 0755)) + if err := os.Symlink(target, filepath.Join(dir, "worktrees", "linked")); err != nil { + if runtime.GOOS == "windows" { + t.Skipf("symlinks unavailable: %v", err) + } + require.NoError(t, err) + } + }}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dir := t.TempDir() + tt.setup(t, dir) + _, err := HasLegacyState(dir) + require.Error(t, err) + require.Error(t, migrateWithTestOperationLock(t, dir)) + assert.NoFileExists(t, stackFilePath(dir)) + assert.NoFileExists(t, filepath.Join(dir, migrationFileName)) + }) + } + + t.Run("missing common directory", func(t *testing.T) { + _, err := HasLegacyState(filepath.Join(t.TempDir(), "missing")) + require.Error(t, err) + }) + + for _, inaccessible := range []string{"admin directory", "catalog"} { + t.Run("inaccessible "+inaccessible, func(t *testing.T) { + if runtime.GOOS == "windows" || os.Geteuid() == 0 { + t.Skip("requires Unix permission enforcement") + } + dir := t.TempDir() + writeMigrationTestCatalog(t, dir, "worktrees/a-readable/gh-stack", migrationTestFile(makeStack("main", "readable"))) + catalog := writeMigrationTestCatalog(t, dir, "worktrees/z-inaccessible/gh-stack", migrationTestFile(makeStack("main", "hidden"))) + path := filepath.Join(dir, filepath.FromSlash(catalog.Path)) + if inaccessible == "admin directory" { + path = filepath.Dir(path) + } + info, err := os.Stat(path) + require.NoError(t, err) + require.NoError(t, os.Chmod(path, 0)) + t.Cleanup(func() { assert.NoError(t, os.Chmod(path, info.Mode().Perm())) }) + _, err = HasLegacyState(dir) + require.Error(t, err, "a readable catalog must not hide a later access failure") + require.Error(t, migrateWithTestOperationLock(t, dir)) + assert.NoFileExists(t, stackFilePath(dir)) + }) + } +} + +func TestMigrateLegacyState_BackupCollisions(t *testing.T) { + for _, relative := range []string{stackFileName, "worktrees/linked/gh-stack"} { + for _, kind := range []string{"equivalent bytes", "different bytes", "directory", "symlink"} { + t.Run(relative+"/"+kind, func(t *testing.T) { + dir := t.TempDir() + catalogs := []migrationCatalog{ + writeMigrationTestCatalog(t, dir, stackFileName, migrationTestFile(makeStack("main", "common"))), + writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", migrationTestFile(makeStack("main", "linked"))), + } + path := filepath.Join(dir, filepath.FromSlash(relative)) + migrationBackupSuffix + original := catalogs[0].Data + if relative != stackFileName { + original = catalogs[1].Data + } + switch kind { + case "equivalent bytes": + require.NoError(t, os.WriteFile(path, original, 0600)) + case "different bytes": + require.NoError(t, os.WriteFile(path, []byte("do not overwrite this backup"), 0600)) + case "directory": + require.NoError(t, os.Mkdir(path, 0755)) + case "symlink": + if err := os.Symlink(filepath.Join(dir, filepath.FromSlash(relative)), path); err != nil { + if runtime.GOOS == "windows" { + t.Skipf("symlinks unavailable: %v", err) + } + require.NoError(t, err) + } + } + err := migrateWithTestOperationLock(t, dir) + if kind == "equivalent bytes" { + require.NoError(t, err) + assertMigrationOriginals(t, dir, catalogs, true) + return + } + require.Error(t, err) + assert.Contains(t, err.Error(), fmt.Sprintf("%q", path)) + if kind == "different bytes" { + var conflict *MigrationConflictError + require.ErrorAs(t, err, &conflict) + data, err := os.ReadFile(path) + require.NoError(t, err) + assert.Equal(t, "do not overwrite this backup", string(data)) + } else if kind == "directory" { + assert.DirExists(t, path) + } else { + info, err := os.Lstat(path) + require.NoError(t, err) + assert.NotZero(t, info.Mode()&os.ModeSymlink) + } + for _, catalog := range catalogs { + data, err := os.ReadFile(filepath.Join(dir, filepath.FromSlash(catalog.Path))) + require.NoError(t, err) + assert.Equal(t, catalog.Data, data) + } + assert.NoFileExists(t, filepath.Join(dir, migrationFileName)) + }) + } + } +} + +func TestMigrateLegacyState_InterruptedMigration(t *testing.T) { + tests := []struct { + name string + common bool + published bool + backups int + archived int + }{ + {"before publication", true, false, 0, 0}, + {"after publication before backups", true, true, 0, 0}, + {"after common backup", true, true, 1, 0}, + {"after linked backup before archive", true, true, 2, 0}, + {"after partial archival", true, true, 2, 1}, + {"after all archival before journal removal", true, true, 3, 2}, + {"absent common before publication", false, false, 0, 0}, + {"absent common after publication", false, true, 0, 0}, + {"absent common partially archived", false, true, 1, 1}, + {"absent common fully archived", false, true, 2, 2}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dir := t.TempDir() + var catalogs []migrationCatalog + want := migrationTestFile() + if tt.common { + common := migrationTestStack() + catalogs = append(catalogs, writeMigrationTestCatalog(t, dir, stackFileName, migrationTestFile(common))) + want.Stacks = append(want.Stacks, common) + } + for i := range 2 { + s := makeStack("main", fmt.Sprintf("linked-%d", i)) + catalogs = append(catalogs, writeMigrationTestCatalog(t, dir, fmt.Sprintf("worktrees/linked-%d/gh-stack", i), migrationTestFile(s))) + want.Stacks = append(want.Stacks, s) + } + journal, err := json.Marshal(migrationState{Version: migrationVersion, Catalogs: catalogs}) + require.NoError(t, err) + require.NoError(t, os.WriteFile(filepath.Join(dir, migrationFileName), journal, 0600)) + merged, err := json.MarshalIndent(want, "", " ") + require.NoError(t, err) + if tt.published { + require.NoError(t, os.WriteFile(stackFilePath(dir), merged, 0644)) + } + for _, catalog := range catalogs[:tt.backups] { + require.NoError(t, os.WriteFile(filepath.Join(dir, filepath.FromSlash(catalog.Path))+migrationBackupSuffix, catalog.Data, catalog.Mode)) + } + archived := 0 + for _, catalog := range catalogs { + if archived == tt.archived { + break + } + if catalog.Path != stackFileName { + require.NoError(t, os.Remove(filepath.Join(dir, filepath.FromSlash(catalog.Path)))) + archived++ + } + } + has, err := HasLegacyState(dir) + require.NoError(t, err) + require.True(t, has) + require.NoError(t, migrateWithTestOperationLock(t, dir)) + got, err := Load(dir) + require.NoError(t, err) + assert.Equal(t, want.Stacks, got.Stacks) + assert.Equal(t, want.Repository, got.Repository) + assertMigrationOriginals(t, dir, catalogs, true) + require.NoError(t, migrateWithTestOperationLock(t, dir)) + assertMigrationOriginals(t, dir, catalogs, true) + }) + } +} + +func TestMigrateLegacyState_InterruptedMigrationRefusesChanges(t *testing.T) { + for _, kind := range []string{"common modified", "legacy modified", "legacy disappeared", "new catalog", "backup conflict", "archive before publication", "recovery appeared"} { + t.Run(kind, func(t *testing.T) { + dir := t.TempDir() + catalogs := []migrationCatalog{ + writeMigrationTestCatalog(t, dir, stackFileName, migrationTestFile(makeStack("main", "common"))), + writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", migrationTestFile(makeStack("main", "linked"))), + } + journal, err := json.Marshal(migrationState{Version: migrationVersion, Catalogs: catalogs}) + require.NoError(t, err) + journalPath := filepath.Join(dir, migrationFileName) + require.NoError(t, os.WriteFile(journalPath, journal, 0600)) + merged, err := json.MarshalIndent(migrationTestFile(makeStack("main", "common"), makeStack("main", "linked")), "", " ") + require.NoError(t, err) + if kind != "archive before publication" { + require.NoError(t, os.WriteFile(stackFilePath(dir), merged, 0644)) + } + legacyPath := filepath.Join(dir, "worktrees", "linked", stackFileName) + switch kind { + case "common modified": + require.NoError(t, os.WriteFile(stackFilePath(dir), []byte("externally changed common catalog"), 0644)) + case "legacy modified": + require.NoError(t, os.WriteFile(legacyPath, []byte("externally changed legacy catalog"), 0644)) + case "legacy disappeared": + require.NoError(t, os.Remove(legacyPath)) + case "new catalog": + writeMigrationTestCatalog(t, dir, "worktrees/new/gh-stack", migrationTestFile(makeStack("main", "new"))) + case "backup conflict": + require.NoError(t, os.WriteFile(legacyPath+migrationBackupSuffix, []byte("existing unrelated backup"), 0644)) + case "archive before publication": + require.NoError(t, os.WriteFile(legacyPath+migrationBackupSuffix, catalogs[1].Data, 0644)) + require.NoError(t, os.Remove(legacyPath)) + case "recovery appeared": + require.NoError(t, os.WriteFile(filepath.Join(dir, "gh-stack-rebase-state"), []byte("{}"), 0600)) + } + before, err := os.ReadFile(stackFilePath(dir)) + require.NoError(t, err) + err = migrateWithTestOperationLock(t, dir) + if kind == "recovery appeared" { + var blocked *MigrationBlockedError + require.ErrorAs(t, err, &blocked) + } else { + var conflict *MigrationConflictError + require.ErrorAs(t, err, &conflict) + } + after, err := os.ReadFile(stackFilePath(dir)) + require.NoError(t, err) + assert.Equal(t, before, after) + afterJournal, err := os.ReadFile(journalPath) + require.NoError(t, err) + assert.Equal(t, journal, afterJournal, "recovery snapshots must remain available") + assert.NoFileExists(t, stackFilePath(dir)+migrationBackupSuffix) + }) + } +} + +func TestMigrateLegacyState_InvalidJournals(t *testing.T) { + for _, kind := range []string{"invalid JSON", "null", "new version", "missing version", "missing catalogs", "unsafe path", "duplicate path", "invalid mode"} { + t.Run(kind, func(t *testing.T) { + dir := t.TempDir() + catalog := writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", migrationTestFile(makeStack("main", "linked"))) + state := migrationState{Version: migrationVersion, Catalogs: []migrationCatalog{catalog}} + switch kind { + case "new version": + state.Version++ + case "missing version": + state.Version = 0 + case "missing catalogs": + state.Catalogs = nil + case "unsafe path": + state.Catalogs[0].Path = "../gh-stack" + case "duplicate path": + state.Catalogs = append(state.Catalogs, catalog) + case "invalid mode": + state.Catalogs[0].Mode = os.ModeSymlink + } + journal, err := json.Marshal(state) + require.NoError(t, err) + if kind == "invalid JSON" { + journal = []byte("{incomplete") + } else if kind == "null" { + journal = []byte("null") + } else if kind == "missing version" { + var fields map[string]json.RawMessage + require.NoError(t, json.Unmarshal(journal, &fields)) + delete(fields, "version") + journal, err = json.Marshal(fields) + require.NoError(t, err) + } + journalPath := filepath.Join(dir, migrationFileName) + require.NoError(t, os.WriteFile(journalPath, journal, 0600)) + require.Error(t, migrateWithTestOperationLock(t, dir)) + data, err := os.ReadFile(filepath.Join(dir, filepath.FromSlash(catalog.Path))) + require.NoError(t, err) + assert.Equal(t, catalog.Data, data) + assert.FileExists(t, journalPath) + assert.NoFileExists(t, stackFilePath(dir)) + assert.NoFileExists(t, filepath.Join(dir, filepath.FromSlash(catalog.Path))+migrationBackupSuffix) + }) + } +} + +func TestMigrateLegacyState_PreservesStaleDetection(t *testing.T) { + dir := t.TempDir() + writeMigrationTestCatalog(t, dir, stackFileName, migrationTestFile(makeStack("main", "common"))) + writeMigrationTestCatalog(t, dir, "worktrees/linked/gh-stack", migrationTestFile(makeStack("main", "linked"))) + before, err := Load(dir) + require.NoError(t, err) + operation, err := LockOperation(dir) + require.NoError(t, err) + defer operation.Unlock() + require.NoError(t, MigrateLegacyState(dir)) + err = Save(dir, before) + var stale *StaleError + require.True(t, errors.As(err, &stale)) + after, err := Load(dir) + require.NoError(t, err) + after.AddStack(makeStack("main", "new")) + require.NoError(t, Save(dir, after)) + require.NoError(t, Save(dir, after), "atomic publication must refresh the load checksum") +}