|
9 | 9 | "os/exec" |
10 | 10 | "path/filepath" |
11 | 11 | "strings" |
| 12 | + "sync" |
| 13 | + "sync/atomic" |
12 | 14 | "testing" |
13 | 15 |
|
14 | 16 | "github.com/github/gh-stack/internal/config" |
@@ -953,6 +955,85 @@ func TestRebase_StateRoundTrip(t *testing.T) { |
953 | 955 | assert.Equal(t, original.OntoOldBase, loaded.OntoOldBase) |
954 | 956 | } |
955 | 957 |
|
| 958 | +func TestRebase_StateReadErrors(t *testing.T) { |
| 959 | + for _, kind := range []string{"missing", "directory", "invalid JSON"} { |
| 960 | + t.Run(kind, func(t *testing.T) { |
| 961 | + dir := t.TempDir() |
| 962 | + path := filepath.Join(dir, rebaseStateFile) |
| 963 | + switch kind { |
| 964 | + case "directory": |
| 965 | + require.NoError(t, os.Mkdir(path, 0700)) |
| 966 | + case "invalid JSON": |
| 967 | + require.NoError(t, os.WriteFile(path, []byte("{incomplete"), 0600)) |
| 968 | + } |
| 969 | + got, err := loadRebaseState(dir) |
| 970 | + require.Error(t, err) |
| 971 | + assert.Nil(t, got) |
| 972 | + switch kind { |
| 973 | + case "missing": |
| 974 | + assert.ErrorIs(t, err, os.ErrNotExist) |
| 975 | + case "directory": |
| 976 | + var pathErr *os.PathError |
| 977 | + require.ErrorAs(t, err, &pathErr) |
| 978 | + assert.Equal(t, path, pathErr.Path) |
| 979 | + case "invalid JSON": |
| 980 | + var syntaxErr *json.SyntaxError |
| 981 | + assert.ErrorAs(t, err, &syntaxErr) |
| 982 | + } |
| 983 | + }) |
| 984 | + } |
| 985 | +} |
| 986 | + |
| 987 | +func TestRebase_StateConcurrentReadWrite(t *testing.T) { |
| 988 | + dir := t.TempDir() |
| 989 | + state := &rebaseState{OriginalBranch: "initial", ConflictBranch: "initial"} |
| 990 | + require.NoError(t, saveRebaseState(dir, state)) |
| 991 | + stop := make(chan struct{}) |
| 992 | + errs := make(chan error, 2) |
| 993 | + var reads atomic.Int64 |
| 994 | + var readers sync.WaitGroup |
| 995 | + for range 2 { |
| 996 | + readers.Go(func() { |
| 997 | + for { |
| 998 | + select { |
| 999 | + case <-stop: |
| 1000 | + return |
| 1001 | + default: |
| 1002 | + } |
| 1003 | + got, err := loadRebaseState(dir) |
| 1004 | + if err != nil { |
| 1005 | + errs <- err |
| 1006 | + return |
| 1007 | + } |
| 1008 | + if got == nil || got.OriginalBranch == "" || got.OriginalBranch != got.ConflictBranch { |
| 1009 | + errs <- fmt.Errorf("reader observed an incomplete rebase state") |
| 1010 | + return |
| 1011 | + } |
| 1012 | + reads.Add(1) |
| 1013 | + } |
| 1014 | + }) |
| 1015 | + } |
| 1016 | + var writeErr error |
| 1017 | + for i := range 50 { |
| 1018 | + branch := strings.Repeat(fmt.Sprintf("%04d", i), 8192) |
| 1019 | + state.OriginalBranch, state.ConflictBranch = branch, branch |
| 1020 | + if writeErr = saveRebaseState(dir, state); writeErr != nil { |
| 1021 | + break |
| 1022 | + } |
| 1023 | + } |
| 1024 | + close(stop) |
| 1025 | + readers.Wait() |
| 1026 | + close(errs) |
| 1027 | + require.NoError(t, writeErr) |
| 1028 | + for err := range errs { |
| 1029 | + require.NoError(t, err) |
| 1030 | + } |
| 1031 | + assert.Positive(t, reads.Load()) |
| 1032 | + got, err := loadRebaseState(dir) |
| 1033 | + require.NoError(t, err) |
| 1034 | + assert.Equal(t, state, got) |
| 1035 | +} |
| 1036 | + |
956 | 1037 | // TestRebase_Continue_RebasesRemainingBranches verifies the --continue success |
957 | 1038 | // path: RebaseContinue is called, remaining branches are rebased via RebaseOnto, |
958 | 1039 | // the state file is cleaned up, and the original branch is restored. |
|
0 commit comments