package safeio import ( "errors" "os" "path/filepath" "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 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 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) } }