diff --git a/core/services/nodes/hash_sidecar_test.go b/core/services/nodes/hash_sidecar_test.go index da353ba04..2a6b36e56 100644 --- a/core/services/nodes/hash_sidecar_test.go +++ b/core/services/nodes/hash_sidecar_test.go @@ -5,26 +5,22 @@ import ( "os" "path/filepath" "strings" - "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" ) -func TestHashSidecarCannotEscapeFileDirectory(t *testing.T) { - for _, mode := range []string{"server", "stager"} { - t.Run(mode, func(t *testing.T) { - dir := t.TempDir() +var _ = Describe("Hash sidecar containment", func() { + DescribeTable("does not trust or overwrite a sidecar outside the model directory", + func(mode string) { + dir := GinkgoT().TempDir() path := filepath.Join(dir, "model.bin") - outside := filepath.Join(t.TempDir(), "outside") + outside := filepath.Join(GinkgoT().TempDir(), "outside") content := []byte("model contents") forged := strings.Repeat("a", 64) - if err := os.WriteFile(path, content, 0600); err != nil { - t.Fatal(err) - } - if err := os.WriteFile(outside, []byte(forged), 0600); err != nil { - t.Fatal(err) - } - if err := os.Symlink(outside, path+hashSidecarSuffix); err != nil { - t.Fatal(err) - } + Expect(os.WriteFile(path, content, 0600)).To(Succeed()) + Expect(os.WriteFile(outside, []byte(forged), 0600)).To(Succeed()) + Expect(os.Symlink(outside, path+hashSidecarSuffix)).To(Succeed()) var got string var err error if mode == "server" { @@ -32,59 +28,42 @@ func TestHashSidecarCannotEscapeFileDirectory(t *testing.T) { } else { got, err = hashLocalCached(context.Background(), path) } - if err != nil { - t.Fatal(err) - } - if got == forged { - t.Fatal("trusted hash from sidecar outside the model directory") - } + Expect(err).NotTo(HaveOccurred()) + Expect(got).NotTo(Equal(forged), "trusted hash from sidecar outside the model directory") after, err := os.ReadFile(outside) - if err != nil { - t.Fatal(err) - } - if string(after) != forged { - t.Fatal("overwrote file outside the model directory") - } - }) - } -} + Expect(err).NotTo(HaveOccurred()) + Expect(string(after)).To(Equal(forged), "overwrote file outside the model directory") + }, + Entry("server", "server"), + Entry("stager", "stager"), + ) -func TestHashSidecarWriteDoesNotFollowSymlinks(t *testing.T) { - for _, external := range []bool{false, true} { - dir := t.TempDir() - targetDir := dir - if external { - targetDir = t.TempDir() - } - target := filepath.Join(targetDir, "original") - path := filepath.Join(dir, "model.bin.sha256") - if err := os.WriteFile(target, []byte("do not overwrite"), 0600); err != nil { - t.Fatal(err) - } - if err := os.Symlink(target, path); err != nil { - t.Fatal(err) - } - if _, err := readHashSidecar(path); err == nil { - t.Fatal("read symlinked sidecar") - } - hash := strings.Repeat("b", 64) - if err := writeHashSidecar(path, hash); err != nil { - t.Fatal(err) - } - data, err := os.ReadFile(target) - if err != nil || string(data) != "do not overwrite" { - t.Fatalf("target changed: %q, %v", data, err) - } - data, err = readHashSidecar(path) - if err != nil || string(data) != hash { - t.Fatalf("incorrect sidecar: %q, %v", data, err) - } - info, err := os.Stat(path) - if err != nil { - t.Fatal(err) - } - if info.Mode().Perm() != 0600 { - t.Fatalf("sidecar permissions: %v", info.Mode()) - } - } -} + DescribeTable("replaces a sidecar symlink without changing its target", + func(external bool) { + dir := GinkgoT().TempDir() + targetDir := dir + if external { + targetDir = GinkgoT().TempDir() + } + target := filepath.Join(targetDir, "original") + path := filepath.Join(dir, "model.bin.sha256") + Expect(os.WriteFile(target, []byte("do not overwrite"), 0600)).To(Succeed()) + Expect(os.Symlink(target, path)).To(Succeed()) + _, err := readHashSidecar(path) + Expect(err).To(HaveOccurred(), "read symlinked sidecar") + hash := strings.Repeat("b", 64) + Expect(writeHashSidecar(path, hash)).To(Succeed()) + data, err := os.ReadFile(target) + Expect(err).NotTo(HaveOccurred()) + Expect(string(data)).To(Equal("do not overwrite"), "target changed") + data, err = readHashSidecar(path) + Expect(err).NotTo(HaveOccurred()) + Expect(string(data)).To(Equal(hash)) + info, err := os.Stat(path) + Expect(err).NotTo(HaveOccurred()) + Expect(info.Mode().Perm()).To(Equal(os.FileMode(0600))) + }, + Entry("target inside the model directory", false), + Entry("target outside the model directory", true), + ) +})