| 1 | package bootstrap |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "errors" |
| 6 | "fmt" |
| 7 | "os" |
| 8 | "testing" |
| 9 | "time" |
| 10 | |
| 11 | "github.com/pkg/sftp" |
| 12 | |
| 13 | "reasonix/internal/remote" |
| 14 | "reasonix/internal/remote/sftpfs" |
| 15 | ) |
| 16 | |
| 17 | type releaseRaceFS struct { |
| 18 | *sftpfs.FS |
| 19 | mkdir func(context.Context, string) error |
| 20 | stat func(context.Context, string) (sftpfs.Entry, error) |
| 21 | } |
| 22 | |
| 23 | func (f releaseRaceFS) MkdirExclusive(ctx context.Context, path string) error { |
| 24 | return f.mkdir(ctx, path) |
| 25 | } |
| 26 | |
| 27 | func (f releaseRaceFS) Stat(ctx context.Context, path string) (sftpfs.Entry, error) { |
| 28 | if f.stat != nil { |
| 29 | return f.stat(ctx, path) |
| 30 | } |
| 31 | return f.FS.Stat(ctx, path) |
| 32 | } |
| 33 | |
| 34 | func TestServeLockAcquiresAfterObservedOwnerRelease(t *testing.T) { |
| 35 | skipOnWindows(t) |
| 36 | root := t.TempDir() |
| 37 | conn := newFakeConn(t, root, func(string) (remote.ExecResult, error) { return ok("") }) |
| 38 | paths := pathsFor(root, root) |
| 39 | owner, err := acquireServeLock(context.Background(), conn.fs, paths, time.Now) |
| 40 | if err != nil { |
| 41 | t.Fatal(err) |
| 42 | } |
| 43 | t.Cleanup(owner.release) |
| 44 | calls := 0 |
| 45 | wrapped := releaseRaceFS{FS: conn.fs, mkdir: func(ctx context.Context, path string) error { |
| 46 | calls++ |
| 47 | err := conn.fs.MkdirExclusive(ctx, path) |
| 48 | if calls == 1 { |
| 49 | if err == nil { |
| 50 | t.Fatal("first owner was not exclusive") |
| 51 | } |
| 52 | owner.release() // deterministically release after mkdir failed, before Stat |
| 53 | } |
| 54 | return err |
| 55 | }} |
| 56 | next, err := acquireServeLock(context.Background(), wrapped, paths, time.Now) |
| 57 | if err != nil { |
| 58 | t.Fatal(err) |
| 59 | } |
| 60 | defer next.release() |
| 61 | if calls != 2 || next.owner == owner.owner { |
| 62 | t.Fatalf("acquisition calls=%d, replacement owner unique=%v", calls, next.owner != owner.owner) |
| 63 | } |
| 64 | } |
| 65 | |
| 66 | func TestServeLockDoesNotRetryNonMissingObservations(t *testing.T) { |
| 67 | skipOnWindows(t) |
| 68 | for _, statErr := range []error{os.ErrPermission, errors.New("stat disconnected"), nil} { |
| 69 | root := t.TempDir() |
| 70 | conn := newFakeConn(t, root, func(string) (remote.ExecResult, error) { return ok("") }) |
| 71 | calls := 0 |
| 72 | wrapped := releaseRaceFS{FS: conn.fs, |
| 73 | mkdir: func(context.Context, string) error { calls++; return os.ErrExist }, |
| 74 | stat: func(context.Context, string) (sftpfs.Entry, error) { |
| 75 | return sftpfs.Entry{IsDir: false}, statErr |
| 76 | }, |
| 77 | } |
| 78 | _, err := acquireServeLock(context.Background(), wrapped, pathsFor(root, root), time.Now) |
| 79 | if !errors.Is(err, os.ErrExist) || calls != 1 { |
| 80 | t.Fatalf("stat=%v error=%v calls=%d; expected immediate creation failure", statErr, err, calls) |
| 81 | } |
| 82 | } |
| 83 | } |
| 84 | |
| 85 | func TestServeLockMissingObservationRetriesAreBounded(t *testing.T) { |
| 86 | skipOnWindows(t) |
| 87 | for _, tc := range []struct { |
| 88 | name string |
| 89 | err error |
| 90 | cancel bool |
| 91 | wantCalls int |
| 92 | }{ |
| 93 | {"generic-failure", &sftp.StatusError{Code: uint32(sftp.ErrSSHFxFailure)}, false, 2}, |
| 94 | {"exists", os.ErrExist, false, 2}, |
| 95 | {"permission", os.ErrPermission, false, 1}, |
| 96 | {"transport", errors.New("transport disconnected"), false, 1}, |
| 97 | {"cancelled", os.ErrExist, true, 1}, |
| 98 | } { |
| 99 | t.Run(tc.name, func(t *testing.T) { |
| 100 | root := t.TempDir() |
| 101 | conn := newFakeConn(t, root, func(string) (remote.ExecResult, error) { return ok("") }) |
| 102 | ctx, cancel := context.WithCancel(context.Background()) |
| 103 | defer cancel() |
| 104 | calls := 0 |
| 105 | wrapped := releaseRaceFS{FS: conn.fs, mkdir: func(context.Context, string) error { |
| 106 | calls++ |
| 107 | if tc.cancel { |
| 108 | cancel() |
| 109 | } |
| 110 | return fmt.Errorf("mkdir: %w", tc.err) |
| 111 | }} |
| 112 | _, err := acquireServeLock(ctx, wrapped, pathsFor(root, root), time.Now) |
| 113 | wantErr := tc.err |
| 114 | if tc.cancel { |
| 115 | wantErr = context.Canceled |
| 116 | } |
| 117 | if !errors.Is(err, wantErr) || calls != tc.wantCalls { |
| 118 | t.Fatalf("error=%v calls=%d, want %v/%d", err, calls, wantErr, tc.wantCalls) |
| 119 | } |
| 120 | }) |
| 121 | } |
| 122 | } |
| 123 |