Files

313 lines
9.3 KiB
Go

//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)
}
}
}