diff --git a/pkg/tar/tar.go b/pkg/tar/tar.go index 85cf028..8fc9a1e 100644 --- a/pkg/tar/tar.go +++ b/pkg/tar/tar.go @@ -35,7 +35,8 @@ type PathWhitelistMap map[string]struct{} // ExtractTar extracts a tarball (from a tar.Reader) into the given directory // if pwl is not nil, only the paths in the map are extracted. -func ExtractTar(tr *tar.Reader, dir string, pwl PathWhitelistMap) error { +// If overwrite is true, existing files will be overwritten. +func ExtractTar(tr *tar.Reader, dir string, overwrite bool, pwl PathWhitelistMap) error { um := syscall.Umask(0) defer syscall.Umask(um) for { @@ -50,7 +51,7 @@ func ExtractTar(tr *tar.Reader, dir string, pwl PathWhitelistMap) error { continue } } - err = ExtractFile(tr, hdr, dir) + err = ExtractFile(tr, hdr, dir, overwrite) if err != nil { return fmt.Errorf("error extracting tarball: %v", err) } @@ -60,12 +61,30 @@ func ExtractTar(tr *tar.Reader, dir string, pwl PathWhitelistMap) error { } } -// ExtractFile extracts the file described by hdr fom the given tarball into -// the provided directory -func ExtractFile(tr *tar.Reader, hdr *tar.Header, dir string) error { +// ExtractFile extracts the file described by hdr from the given tarball into +// the provided directory. +// If overwrite is true, existing files will be overwritten. +func ExtractFile(tr *tar.Reader, hdr *tar.Header, dir string, overwrite bool) error { p := filepath.Join(dir, hdr.Name) fi := hdr.FileInfo() typ := hdr.Typeflag + if overwrite { + info, err := os.Lstat(p) + switch { + case os.IsNotExist(err): + case err == nil: + // If the old and new paths are both dirs do nothing or + // RemoveAll will remove all dir's contents + if !info.IsDir() || typ != tar.TypeDir { + err := os.RemoveAll(p) + if err != nil { + return err + } + } + default: + return err + } + } // Create parent dir if it doesn't exists if err := os.MkdirAll(filepath.Dir(p), DEFAULT_DIR_MODE); err != nil { @@ -73,9 +92,6 @@ func ExtractFile(tr *tar.Reader, hdr *tar.Header, dir string) error { } switch { case typ == tar.TypeReg || typ == tar.TypeRegA: - if err := os.MkdirAll(filepath.Dir(p), DEFAULT_DIR_MODE); err != nil { - return err - } f, err := os.OpenFile(p, os.O_CREATE|os.O_RDWR, fi.Mode()) if err != nil { return err diff --git a/pkg/tar/tar_test.go b/pkg/tar/tar_test.go index 626bb7a..7edcb35 100644 --- a/pkg/tar/tar_test.go +++ b/pkg/tar/tar_test.go @@ -16,6 +16,7 @@ package tar import ( "archive/tar" + "fmt" "io" "io/ioutil" "os" @@ -36,6 +37,14 @@ func newTestTar(entries []*testTarEntry) (string, error) { defer t.Close() tw := tar.NewWriter(t) for _, entry := range entries { + // Add default mode + if entry.header.Mode == 0 { + if entry.header.Typeflag == tar.TypeDir { + entry.header.Mode = 0755 + } else { + entry.header.Mode = 0644 + } + } if err := tw.WriteHeader(entry.header); err != nil { return "", err } @@ -49,6 +58,101 @@ func newTestTar(entries []*testTarEntry) (string, error) { return t.Name(), nil } +type fileInfo struct { + path string + typeflag byte + size int64 + contents string + mode os.FileMode +} + +func fileInfoSliceToMap(slice []*fileInfo) map[string]*fileInfo { + fim := make(map[string]*fileInfo, len(slice)) + for _, fi := range slice { + fim[fi.path] = fi + } + return fim +} + +func checkExpectedFiles(dir string, expectedFiles map[string]*fileInfo) error { + files := make(map[string]*fileInfo) + err := filepath.Walk(dir, func(path string, info os.FileInfo, err error) error { + fm := info.Mode() + if path == dir { + return nil + } + relpath, err := filepath.Rel(dir, path) + if err != nil { + return err + } + switch { + case fm.IsRegular(): + files[relpath] = &fileInfo{path: relpath, typeflag: tar.TypeReg, size: info.Size(), mode: info.Mode().Perm()} + case info.IsDir(): + files[relpath] = &fileInfo{path: relpath, typeflag: tar.TypeDir, mode: info.Mode().Perm()} + case fm&os.ModeSymlink != 0: + files[relpath] = &fileInfo{path: relpath, typeflag: tar.TypeSymlink, mode: info.Mode()} + default: + return fmt.Errorf("file mode not handled: %v", fm) + } + + return nil + }) + if err != nil { + return err + } + + // Set defaults for not specified expected file mode + for _, ef := range expectedFiles { + if ef.mode == 0 { + if ef.typeflag == tar.TypeDir { + ef.mode = 0755 + } else { + ef.mode = 0644 + } + } + } + + for _, ef := range expectedFiles { + _, ok := files[ef.path] + if !ok { + return fmt.Errorf("Expected file %q not in files", ef.path) + } + + } + + for _, file := range files { + ef, ok := expectedFiles[file.path] + if !ok { + return fmt.Errorf("file %q not in expectedFiles", file.path) + } + if ef.typeflag != file.typeflag { + return fmt.Errorf("file %q: file type differs: wanted: %d, got: %d", file.path, ef.typeflag, file.typeflag) + } + if ef.typeflag == tar.TypeReg { + if ef.size != file.size { + return fmt.Errorf("file %q: size differs: wanted %d, wanted: %d", file.path, ef.size, file.size) + } + if ef.contents != "" { + buf, err := ioutil.ReadFile(filepath.Join(dir, file.path)) + if err != nil { + return fmt.Errorf("unexpected error: %v", err) + } + if string(buf) != ef.contents { + return fmt.Errorf("unexpected contents, wanted: %s, got: %s", ef.contents, buf) + } + + } + } + // Check modes but ignore symlinks + if ef.mode != file.mode && ef.typeflag != tar.TypeSymlink { + return fmt.Errorf("file %q: mode differs: wanted %#o, got: %#o", file.path, ef.mode, file.mode) + } + + } + return nil +} + func TestExtractTarInsecureSymlink(t *testing.T) { entries := []*testTarEntry{ { @@ -102,7 +206,7 @@ func TestExtractTarInsecureSymlink(t *testing.T) { t.Errorf("unexpected error: %v", err) } defer os.RemoveAll(tmpdir) - err = ExtractTar(tr, tmpdir, nil) + err = ExtractTar(tr, tmpdir, false, nil) if _, ok := err.(insecureLinkError); !ok { t.Errorf("expected insecureSymlinkError error") } @@ -189,7 +293,7 @@ func TestExtractTarFolders(t *testing.T) { t.Errorf("unexpected error: %v", err) } defer os.RemoveAll(tmpdir) - err = ExtractTar(tr, tmpdir, nil) + err = ExtractTar(tr, tmpdir, false, nil) if err != nil { t.Errorf("unexpected error: %v", err) } @@ -334,7 +438,7 @@ func TestExtractTarPWL(t *testing.T) { pwl := make(PathWhitelistMap) pwl["folder/foo.txt"] = struct{}{} - err = ExtractTar(tr, tmpdir, pwl) + err = ExtractTar(tr, tmpdir, false, pwl) if err != nil { t.Errorf("unexpected error: %v", err) } @@ -346,3 +450,192 @@ func TestExtractTarPWL(t *testing.T) { t.Errorf("unexpected number of files found: %d, wanted 1", len(matches)) } } + +func TestExtractTarOverwrite(t *testing.T) { + tmpdir, err := ioutil.TempDir("", "rocket-temp-dir") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer os.RemoveAll(tmpdir) + + entries := []*testTarEntry{ + { + contents: "hello", + header: &tar.Header{ + Name: "hello.txt", + Size: 5, + }, + }, + { + header: &tar.Header{ + Name: "afolder", + Typeflag: tar.TypeDir, + }, + }, + { + contents: "hello", + header: &tar.Header{ + Name: "afolder/hello.txt", + Size: 5, + }, + }, + { + contents: "hello", + header: &tar.Header{ + Name: "afile", + Size: 5, + }, + }, + { + header: &tar.Header{ + Name: "folder01", + Typeflag: tar.TypeDir, + }, + }, + { + contents: "hello", + header: &tar.Header{ + Name: "folder01/file01", + Size: 5, + }, + }, + { + contents: "hello", + header: &tar.Header{ + Name: "filesymlinked", + Size: 5, + }, + }, + { + header: &tar.Header{ + Name: "linktofile", + Linkname: "filesymlinked", + Typeflag: tar.TypeSymlink, + }, + }, + + { + header: &tar.Header{ + Name: "dirsymlinked", + Typeflag: tar.TypeDir, + }, + }, + { + header: &tar.Header{ + Name: "linktodir", + Linkname: "dirsymlinked", + Typeflag: tar.TypeSymlink, + }, + }, + } + + testTarPath, err := newTestTar(entries) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer os.Remove(testTarPath) + containerTar, err := os.Open(testTarPath) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + tr := tar.NewReader(containerTar) + err = ExtractTar(tr, tmpdir, false, nil) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + // Now overwrite: + // a file with a new file + // a dir with a file + entries = []*testTarEntry{ + { + contents: "newhello", + header: &tar.Header{ + Name: "hello.txt", + Size: 8, + }, + }, + // Now this is a file + { + contents: "nowafile", + header: &tar.Header{ + Name: "afolder", + Typeflag: tar.TypeReg, + Size: 8, + }, + }, + // Now this is a dir + { + header: &tar.Header{ + Name: "afile", + Typeflag: tar.TypeDir, + }, + }, + // Overwrite symlink to a file with a regular file + // the linked file shouldn't be removed + { + contents: "filereplacingsymlink", + header: &tar.Header{ + Name: "linktofile", + Typeflag: tar.TypeReg, + Size: 20, + }, + }, + // Overwrite symlink to a dir with a regular file + // the linked directory and all its contents shouldn't be + // removed + { + contents: "filereplacingsymlink", + header: &tar.Header{ + Name: "linktodir", + Typeflag: tar.TypeReg, + Size: 20, + }, + }, + // folder01 already exists and shouldn't be removed (keeping folder01/file01) + { + header: &tar.Header{ + Name: "folder01", + Typeflag: tar.TypeDir, + Mode: int64(0755), + }, + }, + { + contents: "hello", + header: &tar.Header{ + Name: "folder01/file02", + Size: 5, + Mode: int64(0644), + }, + }, + } + testTarPath, err = newTestTar(entries) + if err != nil { + t.Errorf("unexpected error: %v", err) + } + defer os.Remove(testTarPath) + containerTar, err = os.Open(testTarPath) + if err != nil { + t.Errorf("unexpected error: %v", err) + } + tr = tar.NewReader(containerTar) + err = ExtractTar(tr, tmpdir, true, nil) + + expectedFiles := []*fileInfo{ + &fileInfo{path: "hello.txt", typeflag: tar.TypeReg, size: 8, contents: "newhello"}, + &fileInfo{path: "linktofile", typeflag: tar.TypeReg, size: 20}, + &fileInfo{path: "linktodir", typeflag: tar.TypeReg, size: 20}, + &fileInfo{path: "afolder", typeflag: tar.TypeReg, size: 8}, + &fileInfo{path: "dirsymlinked", typeflag: tar.TypeDir}, + &fileInfo{path: "afile", typeflag: tar.TypeDir}, + &fileInfo{path: "filesymlinked", typeflag: tar.TypeReg, size: 5}, + &fileInfo{path: "folder01", typeflag: tar.TypeDir}, + &fileInfo{path: "folder01/file01", typeflag: tar.TypeReg, size: 5}, + &fileInfo{path: "folder01/file02", typeflag: tar.TypeReg, size: 5}, + } + + err = checkExpectedFiles(tmpdir, fileInfoSliceToMap(expectedFiles)) + if err != nil { + t.Errorf("unexpected error: %v", err) + } +} diff --git a/stage0/run.go b/stage0/run.go index 59ddfcb..dfbabc2 100644 --- a/stage0/run.go +++ b/stage0/run.go @@ -230,7 +230,7 @@ func untarRootfs(r io.Reader, dir string) error { return fmt.Errorf("error creating stage1 rootfs directory: %v", err) } - if err := ptar.ExtractTar(tr, dir, nil); err != nil { + if err := ptar.ExtractTar(tr, dir, false, nil); err != nil { return fmt.Errorf("error extracting rootfs: %v", err) } return nil @@ -304,7 +304,7 @@ func setupImage(cfg Config, img types.Hash, dir string) (*schema.ImageManifest, hash := sha512.New() r := io.TeeReader(rs, hash) - if err := ptar.ExtractTar(tar.NewReader(r), ad, nil); err != nil { + if err := ptar.ExtractTar(tar.NewReader(r), ad, false, nil); err != nil { return nil, fmt.Errorf("error extracting ACI: %v", err) }