//go:build !windows package safeio import ( "bytes" "errors" "os" "path/filepath" "testing" "time" "golang.org/x/sys/unix" ) func TestReadCanonicalRegularRejectsNamedPipeWithoutBlocking(t *testing.T) { temporaryRoot, err := filepath.EvalSymlinks(os.TempDir()) if err != nil { t.Fatal(err) } root, err := os.MkdirTemp(temporaryRoot, "thothctl-safeio-") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = os.RemoveAll(root) }) pipe := filepath.Join(root, "secret-pipe") if err := unix.Mkfifo(pipe, 0o600); err != nil { t.Fatal(err) } if _, err := ReadCanonicalRegular(pipe, 1024); !errors.Is(err, ErrUnsafeFile) { t.Fatalf("named pipe error = %v, want ErrUnsafeFile", err) } } func TestReadCanonicalRegularRejectsReplacementDuringRead(t *testing.T) { root := t.TempDir() path := filepath.Join(root, "schema.sql") replacement := filepath.Join(root, "replacement.sql") parked := filepath.Join(root, "parked.sql") contents := bytes.Repeat([]byte("x"), 64<<20) if err := os.WriteFile(path, contents, 0o600); err != nil { t.Fatal(err) } if err := os.WriteFile(replacement, contents, 0o600); err != nil { t.Fatal(err) } var caught bool for attempt := 0; attempt < 3 && !caught; attempt++ { result := make(chan error, 1) go func() { _, err := ReadCanonicalRegular(path, int64(len(contents))) result <- err }() time.Sleep(time.Millisecond) var finalErr error for i := 0; i < 20; i++ { if err := os.Rename(path, parked); err == nil { if err := os.Rename(replacement, path); err != nil { t.Fatal(err) } time.Sleep(time.Millisecond) if err := os.Rename(parked, replacement); err != nil { t.Fatal(err) } } select { case finalErr = <-result: i = 20 default: } } if finalErr == nil { finalErr = <-result } if errors.Is(finalErr, ErrUnsafeFile) { caught = true } } if !caught { t.Fatal("replacement during read was not rejected") } } func TestCanonicalDescriptorOwnershipDoesNotLeakAcrossNestedOperations(t *testing.T) { root, err := filepath.EvalSymlinks(t.TempDir()) if err != nil { t.Fatal(err) } nested := filepath.Join(root, "one", "two") if err := os.MkdirAll(nested, 0o700); err != nil { t.Fatal(err) } input := filepath.Join(nested, "input.sql") if err := os.WriteFile(input, []byte("select 1"), 0o600); err != nil { t.Fatal(err) } fdCount := func() int { f, err := os.Open("/dev/fd") if err != nil { t.Fatal(err) } defer f.Close() names, err := f.Readdirnames(-1) if err != nil { t.Fatal(err) } return len(names) } baseline := fdCount() for i := 0; i < 20; i++ { if _, err := ReadCanonicalRegular(input, 1024); err != nil { t.Fatal(err) } if err := validateCanonicalOutputPath(filepath.Join(nested, "out-"+string(rune('a'+i))+".yaml")); err != nil { t.Fatal(err) } } if got := fdCount(); got > baseline+2 { t.Fatalf("descriptor leak after read/validate: baseline=%d got=%d", baseline, got) } for i := 0; i < 20; i++ { path := filepath.Join(nested, "write-"+string(rune('a'+i))+".yaml") if err := writeCanonicalExclusive(path, []byte("ok"), 0o600); err != nil { t.Fatal(err) } } if got := fdCount(); got > baseline+2 { t.Fatalf("descriptor leak after writes: baseline=%d got=%d", baseline, got) } }