Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
70 changes: 70 additions & 0 deletions cmd/modify_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,13 @@ package cmd

import (
"encoding/json"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"testing"
"time"

Expand Down Expand Up @@ -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()
Expand Down
4 changes: 2 additions & 2 deletions cmd/rebase.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Comment thread
skarim marked this conversation as resolved.
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
}
Expand Down
81 changes: 81 additions & 0 deletions cmd/rebase_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@ import (
"os/exec"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"testing"

"github.com/github/gh-stack/internal/config"
Expand Down Expand Up @@ -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.
Expand Down
18 changes: 17 additions & 1 deletion internal/git/gitops_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{}))
Expand Down
15 changes: 4 additions & 11 deletions internal/modify/state.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@ import (
"os"
"path/filepath"
"time"

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

New line needed?

"github.com/github/gh-stack/internal/stack"
)

const stateFileName = "gh-stack-modify-state"
Expand Down Expand Up @@ -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
Expand All @@ -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 {
Comment thread
skarim marked this conversation as resolved.
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)
}
Comment on lines -99 to -110

@Lukeghenco Lukeghenco Sep 30, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This seems like the biggest core change in this PR right? If I am understanding this correctly we are now writing to a temp file before we replace the pre-existing one (in stack/atomic.go) instead of what appears to be deleting and then replacing the file in the deleted code?

return nil
}

Expand Down
72 changes: 72 additions & 0 deletions internal/stack/atomic.go
Original file line number Diff line number Diff line change
@@ -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))
}
29 changes: 29 additions & 0 deletions internal/stack/atomic_unix.go
Original file line number Diff line number Diff line change
@@ -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())
}
Loading
Loading