diff --git a/pkg/utils/tar.go b/pkg/utils/tar.go index 5286074c..02d6cbd0 100644 --- a/pkg/utils/tar.go +++ b/pkg/utils/tar.go @@ -10,6 +10,19 @@ import ( ) func Untar(dst string, r io.Reader) error { + root, err := filepath.Abs(dst) + if err != nil { + return err + } + if err := os.MkdirAll(root, 0755); err != nil { + return err + } + rootFS, err := os.OpenRoot(root) + if err != nil { + return err + } + defer rootFS.Close() + tr := tar.NewReader(r) madeDir := map[string]bool{} for { @@ -28,8 +41,12 @@ func Untar(dst string, r io.Reader) error { continue } + if !filepath.IsLocal(header.Name) { + return fmt.Errorf("tar entry %q resolves outside destination", header.Name) + } + // the target location where the dir/file should be created - target := filepath.Join(dst, header.Name) + target := filepath.Clean(header.Name) // the following switch could also be done using fi.Mode(), not sure if there // a benefit of using one vs. the other. // fi := header.FileInfo() @@ -38,20 +55,19 @@ func Untar(dst string, r io.Reader) error { switch header.Typeflag { // if its a dir and it doesn't exist create it case tar.TypeDir: - if err := makeDir(target, madeDir); err != nil { + if err := makeDir(rootFS, target, madeDir); err != nil { return err } // if it's a file create it case tar.TypeReg: - if err := makeDir(filepath.Dir(target), madeDir); err != nil { + if err := makeDir(rootFS, filepath.Dir(target), madeDir); err != nil { return err } - f, err := os.OpenFile(target, os.O_CREATE|os.O_RDWR, os.FileMode(header.Mode)) + f, err := rootFS.OpenFile(target, os.O_CREATE|os.O_RDWR, os.FileMode(header.Mode)) if err != nil { - fmt.Println("could not open file") - return err + return fmt.Errorf("extract tar entry %q: %w", header.Name, err) } // copy over contents @@ -66,15 +82,13 @@ func Untar(dst string, r io.Reader) error { } } -func makeDir(target string, made map[string]bool) error { +func makeDir(root *os.Root, target string, made map[string]bool) error { if made[target] { return nil } - if _, err := os.Stat(target); err != nil { - if err := os.MkdirAll(target, 0755); err != nil { - return err - } + if err := root.MkdirAll(target, 0755); err != nil { + return err } made[target] = true diff --git a/pkg/utils/tar_test.go b/pkg/utils/tar_test.go new file mode 100644 index 00000000..8faf272f --- /dev/null +++ b/pkg/utils/tar_test.go @@ -0,0 +1,69 @@ +package utils_test + +import ( + "archive/tar" + "bytes" + "os" + "path/filepath" + "testing" + + "github.com/pluralsh/plural-cli/pkg/utils" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestUntar(t *testing.T) { + t.Run("extracts files inside destination", func(t *testing.T) { + archive := createTar(t, "nested/file.txt", "content") + dst := t.TempDir() + + require.NoError(t, utils.Untar(dst, archive)) + + content, err := os.ReadFile(filepath.Join(dst, "nested", "file.txt")) + require.NoError(t, err) + assert.Equal(t, "content", string(content)) + }) + + t.Run("rejects files outside destination", func(t *testing.T) { + archive := createTar(t, "../outside.txt", "content") + parent := t.TempDir() + dst := filepath.Join(parent, "destination") + + err := utils.Untar(dst, archive) + + require.Error(t, err) + assert.ErrorContains(t, err, "resolves outside destination") + _, err = os.Stat(filepath.Join(parent, "outside.txt")) + assert.ErrorIs(t, err, os.ErrNotExist) + }) + + t.Run("rejects symlinks outside destination", func(t *testing.T) { + archive := createTar(t, "link/file.txt", "content") + dst := t.TempDir() + outside := t.TempDir() + require.NoError(t, os.Symlink(outside, filepath.Join(dst, "link"))) + + err := utils.Untar(dst, archive) + + require.Error(t, err) + _, err = os.Stat(filepath.Join(outside, "file.txt")) + assert.ErrorIs(t, err, os.ErrNotExist) + }) +} + +func createTar(t *testing.T, name, content string) *bytes.Reader { + t.Helper() + + var archive bytes.Buffer + w := tar.NewWriter(&archive) + require.NoError(t, w.WriteHeader(&tar.Header{ + Name: name, + Mode: 0o600, + Size: int64(len(content)), + })) + _, err := w.Write([]byte(content)) + require.NoError(t, err) + require.NoError(t, w.Close()) + + return bytes.NewReader(archive.Bytes()) +}