//go:build windows package safeio import ( "errors" "fmt" "os" "path/filepath" "sync" "testing" "time" ) func createWindowsPrivateTestFile(t *testing.T, path string, contents []byte) { t.Helper() file, err := CreateCanonicalNewPrivateFile(path) if err != nil { t.Fatal(err) } if _, err := file.Write(contents); err != nil { _ = file.Close() t.Fatal(err) } if err := file.Close(); err != nil { t.Fatal(err) } } func TestRemoveCanonicalPrivateClaimRetainsParentDuringDeletion(t *testing.T) { parent := filepath.Join(t.TempDir(), "claims") if err := os.Mkdir(parent, 0o700); err != nil { t.Fatal(err) } if err := ProtectPrivateDirectory(parent); err != nil { t.Fatal(err) } source := filepath.Join(parent, "state.json") claim := filepath.Join(parent, "state.claim") createWindowsPrivateTestFile(t, source, []byte("state")) if claimed, err := ClaimCanonicalPrivateRegular(source, claim); err != nil || !claimed { t.Fatalf("ClaimCanonicalPrivateRegular() = claimed %v, err %v", claimed, err) } outside := filepath.Join(t.TempDir(), "outside") if err := os.Mkdir(outside, 0o700); err != nil { t.Fatal(err) } if err := ProtectPrivateDirectory(outside); err != nil { t.Fatal(err) } sentinel := filepath.Join(outside, "sentinel") createWindowsPrivateTestFile(t, sentinel, []byte("outside-safe")) attemptedSwap := false restoreHook := SetPrivateDirectoryTestHookForTest(func(stage string) { if stage != "after-canonical-private-claim-parent-open" || attemptedSwap { return } attemptedSwap = true if err := os.Rename(parent, parent+"-moved"); err == nil { t.Fatal("claim parent rename succeeded while removal retained its handle") } }) defer restoreHook() removed, err := RemoveCanonicalPrivateClaim(source, claim) if err != nil || !removed || !attemptedSwap { t.Fatalf("RemoveCanonicalPrivateClaim() = removed %v, attempted %v, err %v", removed, attemptedSwap, err) } for _, path := range []string{source, claim} { if _, err := os.Stat(path); !errors.Is(err, os.ErrNotExist) { t.Fatalf("removed path %q still exists: %v", filepath.Base(path), err) } } if contents, err := os.ReadFile(sentinel); err != nil || string(contents) != "outside-safe" { t.Fatalf("outside sentinel = %q, err %v", contents, err) } } func TestRemoveCanonicalPrivateClaimPreservesOrphan(t *testing.T) { parent := filepath.Join(t.TempDir(), "claims") if err := os.Mkdir(parent, 0o700); err != nil { t.Fatal(err) } if err := ProtectPrivateDirectory(parent); err != nil { t.Fatal(err) } source := filepath.Join(parent, "state.json") claim := filepath.Join(parent, "state.claim") createWindowsPrivateTestFile(t, claim, []byte("orphan")) removed, err := RemoveCanonicalPrivateClaim(source, claim) if removed || err != nil { t.Fatalf("RemoveCanonicalPrivateClaim() = removed %v, err %v, want false/nil", removed, err) } if _, err := os.Stat(source); !errors.Is(err, os.ErrNotExist) { t.Fatalf("orphan source unexpectedly exists: %v", err) } if contents, err := os.ReadFile(claim); err != nil || string(contents) != "orphan" { t.Fatalf("orphan claim = %q, err %v", contents, err) } } func TestCanonicalPrivateClaimWaitsForRetainedRemoveOperation(t *testing.T) { parent := filepath.Join(t.TempDir(), "claims") if err := os.Mkdir(parent, 0o700); err != nil { t.Fatal(err) } if err := ProtectPrivateDirectory(parent); err != nil { t.Fatal(err) } source := filepath.Join(parent, "state.json") claim := filepath.Join(parent, "state.claim") createWindowsPrivateTestFile(t, source, []byte("state")) if claimed, err := ClaimCanonicalPrivateRegular(source, claim); err != nil || !claimed { t.Fatalf("ClaimCanonicalPrivateRegular() = claimed %v, err %v", claimed, err) } removeOpened := make(chan struct{}) releaseRemove := make(chan struct{}) var releaseOnce sync.Once release := func() { releaseOnce.Do(func() { close(releaseRemove) }) } defer release() restoreHook := SetPrivateDirectoryTestHookForTest(func(stage string) { if stage != "after-canonical-private-claim-parent-open" { return } select { case <-removeOpened: default: close(removeOpened) } <-releaseRemove }) defer restoreHook() type result struct { changed bool err error } removeResult := make(chan result, 1) go func() { removed, err := RemoveCanonicalPrivateClaim(source, claim) removeResult <- result{changed: removed, err: err} }() <-removeOpened claimResult := make(chan result, 1) go func() { claimed, err := ClaimCanonicalPrivateRegular(source, claim) claimResult <- result{changed: claimed, err: err} }() select { case got := <-claimResult: t.Fatalf("concurrent claim returned before retained removal completed: claimed %v, err %v", got.changed, got.err) case <-time.After(250 * time.Millisecond): } release() if got := <-removeResult; got.err != nil || !got.changed { t.Fatalf("RemoveCanonicalPrivateClaim() = removed %v, err %v", got.changed, got.err) } if got := <-claimResult; got.err != nil || got.changed { t.Fatalf("concurrent ClaimCanonicalPrivateRegular() = claimed %v, err %v, want false/nil", got.changed, got.err) } } func TestRemoveCanonicalPrivateClaimRejectsMismatchedTwoLinkFiles(t *testing.T) { parent := filepath.Join(t.TempDir(), "claims") if err := os.Mkdir(parent, 0o700); err != nil { t.Fatal(err) } if err := ProtectPrivateDirectory(parent); err != nil { t.Fatal(err) } source := filepath.Join(parent, "state.json") claim := filepath.Join(parent, "state.claim") createWindowsPrivateTestFile(t, source, []byte("source")) createWindowsPrivateTestFile(t, claim, []byte("claim")) sourceAuxiliary := filepath.Join(parent, "state.json.aux") if err := os.Link(source, sourceAuxiliary); err != nil { t.Fatal(err) } claimAuxiliary := filepath.Join(parent, "state.claim.aux") if err := os.Link(claim, claimAuxiliary); err != nil { t.Fatal(err) } removed, err := RemoveCanonicalPrivateClaim(source, claim) if removed || !errors.Is(err, ErrUnsafeFile) { t.Fatalf("RemoveCanonicalPrivateClaim() = removed %v, err %v, want false/ErrUnsafeFile", removed, err) } for path, want := range map[string]string{ source: "source", claim: "claim", sourceAuxiliary: "source", claimAuxiliary: "claim", } { contents, readErr := os.ReadFile(path) if readErr != nil || string(contents) != want { t.Fatalf("mismatched pair path %q = %q, err %v", filepath.Base(path), contents, readErr) } } } func TestClaimCanonicalPrivateRegularHasOneConcurrentWinner(t *testing.T) { parent := filepath.Join(t.TempDir(), "claims") if err := os.Mkdir(parent, 0o700); err != nil { t.Fatal(err) } if err := ProtectPrivateDirectory(parent); err != nil { t.Fatal(err) } for iteration := range 16 { source := filepath.Join(parent, fmt.Sprintf("state-%02d.json", iteration)) claim := filepath.Join(parent, fmt.Sprintf("state-%02d.claim", iteration)) createWindowsPrivateTestFile(t, source, []byte("state")) type result struct { claimed bool err error } results := make(chan result, 2) var group sync.WaitGroup for range 2 { group.Add(1) go func() { defer group.Done() claimed, err := ClaimCanonicalPrivateRegular(source, claim) results <- result{claimed: claimed, err: err} }() } group.Wait() close(results) winners := 0 for got := range results { if got.err != nil { t.Fatalf("iteration %d concurrent claim error = %v", iteration, got.err) } if got.claimed { winners++ } } if winners != 1 { t.Fatalf("iteration %d winning claims = %d, want 1", iteration, winners) } if removed, err := RemoveCanonicalPrivateClaim(source, claim); err != nil || !removed { t.Fatalf("iteration %d claim cleanup = removed %v, err %v", iteration, removed, err) } } } func TestCanonicalPrivateClaimConsumeHasOneConcurrentWinner(t *testing.T) { parent := filepath.Join(t.TempDir(), "claims") if err := os.Mkdir(parent, 0o700); err != nil { t.Fatal(err) } if err := ProtectPrivateDirectory(parent); err != nil { t.Fatal(err) } for iteration := range 16 { source := filepath.Join(parent, fmt.Sprintf("consume-%02d.json", iteration)) claim := filepath.Join(parent, fmt.Sprintf("consume-%02d.claim", iteration)) createWindowsPrivateTestFile(t, source, []byte("state")) type result struct { found bool stage string err error } results := make(chan result, 2) var group sync.WaitGroup for range 2 { group.Add(1) go func() { defer group.Done() claimed, err := ClaimCanonicalPrivateRegular(source, claim) if err != nil || !claimed { results <- result{stage: "claim", err: err} return } contents, found, err := ReadCanonicalPrivateClaim(source, claim, 32) if err != nil || !found || string(contents) != "state" { results <- result{stage: "read", err: err} return } removed, err := RemoveCanonicalPrivateClaim(source, claim) if err != nil || !removed { results <- result{stage: "remove", err: err} return } results <- result{found: true, stage: "complete"} }() } group.Wait() close(results) winners := 0 for got := range results { if got.err != nil { t.Fatalf("iteration %d concurrent consume %s error = %v", iteration, got.stage, got.err) } if got.found { winners++ } } if winners != 1 { t.Fatalf("iteration %d winning consumes = %d, want 1", iteration, winners) } } }