133 lines
3.2 KiB
Go
133 lines
3.2 KiB
Go
//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)
|
|
}
|
|
}
|