package safeio import ( "errors" "fmt" "os" "path/filepath" "testing" "github.com/aritmolab/thothii/tools/tht/internal/testsupport" ) func TestReadCanonicalRegularRejectsFinalAndParentSymlinks(t *testing.T) { temporaryRoot, err := filepath.EvalSymlinks(os.TempDir()) if err != nil { t.Fatal(err) } root, err := os.MkdirTemp(temporaryRoot, "tht-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 TestReadCanonicalUTF8RejectsNonUTF8AndHardlinks(t *testing.T) { temporaryRoot, err := filepath.EvalSymlinks(os.TempDir()) if err != nil { t.Fatal(err) } root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = os.RemoveAll(root) }) nonUTF8 := filepath.Join(root, "annotations.yaml") if err := os.WriteFile(nonUTF8, []byte{0xff, 0xfe, 0xfd}, 0o600); err != nil { t.Fatal(err) } if _, err := ReadCanonicalUTF8(nonUTF8, 1024); !errors.Is(err, ErrUnsafeFile) { t.Fatalf("ReadCanonicalUTF8(nonUTF8) error = %v, want ErrUnsafeFile", err) } target := filepath.Join(root, "regular.txt") if err := os.WriteFile(target, []byte("linked"), 0o600); err != nil { t.Fatal(err) } link := filepath.Join(root, "hardlink.txt") if err := os.Link(target, link); err != nil { t.Fatal(err) } if _, err := ReadCanonicalUTF8(link, 1024); !errors.Is(err, ErrUnsafeFile) { t.Fatalf("ReadCanonicalUTF8(hardlink) error = %v, want ErrUnsafeFile", err) } } func TestWriteCanonicalNewFileRejectsExistingTargets(t *testing.T) { temporaryRoot, err := filepath.EvalSymlinks(os.TempDir()) if err != nil { t.Fatal(err) } root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = os.RemoveAll(root) }) path := filepath.Join(root, "artifact.yaml") if err := os.WriteFile(path, []byte("existing"), 0o600); err != nil { t.Fatal(err) } if err := WriteCanonicalNewFile(path, []byte("new"), 0o600); !errors.Is(err, ErrUnsafeFile) { t.Fatalf("WriteCanonicalNewFile(existing) error = %v, want ErrUnsafeFile", err) } } func TestListCanonicalPrivateDirectoryBoundsAndSortsValidatedEntries(t *testing.T) { temporaryRoot, err := filepath.EvalSymlinks(os.TempDir()) if err != nil { t.Fatal(err) } root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-list-") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = os.RemoveAll(root) }) if err := ProtectPrivateDirectory(root); err != nil { t.Fatal(err) } for index := 255; index >= 0; index-- { path := filepath.Join(root, fmt.Sprintf("%064x.json", index)) if err := WriteCanonicalNewFile(path, []byte("record"), 0o600); err != nil { t.Fatalf("WriteCanonicalNewFile(%d) error = %v", index, err) } } entries, err := ListCanonicalPrivateDirectory(root, 256) if err != nil || len(entries) != 256 || entries[0].Name != fmt.Sprintf("%064x.json", 0) || entries[255].Name != fmt.Sprintf("%064x.json", 255) { t.Fatalf("bounded ordered entries = %#v error = %v", entries, err) } if err := WriteCanonicalNewFile(filepath.Join(root, fmt.Sprintf("%064x.json", 256)), []byte("record"), 0o600); err != nil { t.Fatal(err) } if _, err := ListCanonicalPrivateDirectory(root, 256); !errors.Is(err, ErrUnsafeFile) { t.Fatalf("257-entry listing error = %v, want ErrUnsafeFile", err) } }