Files
ThothII/tools/thothctl/internal/safeio/files_test.go
T

160 lines
5.0 KiB
Go

package safeio
import (
"errors"
"os"
"path/filepath"
"strings"
"testing"
"github.com/aritmolab/thothii/tools/thothctl/internal/testsupport"
)
func TestReadCanonicalRegularRejectsFinalAndParentSymlinks(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) })
realDirectory := filepath.Join(root, "real")
if err := os.Mkdir(realDirectory, 0o700); err != nil {
t.Fatal(err)
}
realFile := filepath.Join(realDirectory, "secret")
if err := os.WriteFile(realFile, []byte("secret"), 0o600); err != nil {
t.Fatal(err)
}
parentLink := filepath.Join(root, "parent-link")
testsupport.SymlinkOrSkip(t, realDirectory, parentLink)
if _, err := ReadCanonicalRegular(filepath.Join(parentLink, "secret"), 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("parent symlink error = %v, want ErrUnsafeFile", err)
}
finalLink := filepath.Join(root, "final-link")
testsupport.SymlinkOrSkip(t, realFile, finalLink)
if _, err := ReadCanonicalRegular(finalLink, 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("final symlink error = %v, want ErrUnsafeFile", err)
}
}
func TestReadCanonicalUTF8RejectsNonUTF8AndBounds(t *testing.T) {
path := filepath.Join(t.TempDir(), "input.sql")
if err := os.WriteFile(path, []byte("\xff"), 0o600); err != nil {
t.Fatal(err)
}
if _, err := ReadCanonicalUTF8(path, 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("ReadCanonicalUTF8 invalid UTF-8 = %v", err)
}
if err := os.WriteFile(path, []byte("12345"), 0o600); err != nil {
t.Fatal(err)
}
if _, err := ReadCanonicalUTF8(path, 4); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("ReadCanonicalUTF8 oversized = %v", err)
}
}
func TestReadCanonicalUTF8RejectsSQLLargerThanOneMiB(t *testing.T) {
path := filepath.Join(t.TempDir(), "schema.sql")
if err := os.WriteFile(path, make([]byte, (1<<20)+1), 0o600); err != nil {
t.Fatal(err)
}
if _, err := ReadCanonicalUTF8(path, 1<<20); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("ReadCanonicalUTF8 1 MiB + 1 SQL = %v, want ErrUnsafeFile", err)
}
}
func TestWriteCanonicalExclusiveRejectsExistingAndCreatesPrivateFile(t *testing.T) {
root, err := filepath.EvalSymlinks(t.TempDir())
if err != nil {
t.Fatal(err)
}
path := filepath.Join(root, "candidate.yaml")
if err := WriteCanonicalExclusive(path, []byte("ok"), 0o600); err != nil {
t.Fatal(err)
}
contents, err := os.ReadFile(path)
if err != nil || string(contents) != "ok" {
t.Fatalf("output = %q, %v", contents, err)
}
if err := WriteCanonicalExclusive(path, []byte("replace"), 0o600); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("replacement = %v", err)
}
}
func TestWriteCanonicalExclusiveCleansPrivateStageWhenPublicationRacesExistingLeaf(t *testing.T) {
root, err := filepath.EvalSymlinks(t.TempDir())
if err != nil {
t.Fatal(err)
}
path := filepath.Join(root, "candidate.yaml")
if err := os.WriteFile(path, []byte("attacker"), 0o600); err != nil {
t.Fatal(err)
}
if err := WriteCanonicalExclusive(path, []byte("candidate"), 0o600); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("write error=%v", err)
}
entries, err := os.ReadDir(root)
if err != nil {
t.Fatal(err)
}
for _, entry := range entries {
if strings.HasPrefix(entry.Name(), ".thothctl-candidate-") {
t.Fatalf("private stage leaked: %s", entry.Name())
}
}
}
func TestReadCanonicalRegularRejectsHardlinkAndDirectory(t *testing.T) {
root, err := filepath.EvalSymlinks(t.TempDir())
if err != nil {
t.Fatal(err)
}
original := filepath.Join(root, "original.sql")
if err := os.WriteFile(original, []byte("select 1"), 0o600); err != nil {
t.Fatal(err)
}
hardlink := filepath.Join(root, "hardlink.sql")
if err := os.Link(original, hardlink); err != nil {
t.Fatal(err)
}
if _, err := ReadCanonicalRegular(hardlink, 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("hardlink error = %v, want ErrUnsafeFile", err)
}
directory := filepath.Join(root, "directory.sql")
if err := os.Mkdir(directory, 0o700); err != nil {
t.Fatal(err)
}
if _, err := ReadCanonicalRegular(directory, 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("directory error = %v, want ErrUnsafeFile", err)
}
}
func TestValidateCanonicalOutputPathRejectsExistingDirectoryAndSymlink(t *testing.T) {
root, err := filepath.EvalSymlinks(t.TempDir())
if err != nil {
t.Fatal(err)
}
directory := filepath.Join(root, "existing.yaml")
if err := os.Mkdir(directory, 0o700); err != nil {
t.Fatal(err)
}
if err := ValidateCanonicalOutputPath(directory); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("existing directory error = %v, want ErrUnsafeFile", err)
}
target := filepath.Join(root, "target.yaml")
if err := os.WriteFile(target, []byte("target"), 0o600); err != nil {
t.Fatal(err)
}
link := filepath.Join(root, "link.yaml")
testsupport.SymlinkOrSkip(t, target, link)
if err := ValidateCanonicalOutputPath(link); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("output symlink error = %v, want ErrUnsafeFile", err)
}
}