diff --git a/internal/update/extract.go b/internal/update/extract.go index 84b86ecef..8a5d1431f 100644 --- a/internal/update/extract.go +++ b/internal/update/extract.go @@ -21,6 +21,17 @@ func extractArchive(archivePath string, destDir string) error { } func extractTarGz(archivePath string, destDir string) error { + if err := os.MkdirAll(destDir, 0o755); err != nil { + return err + } + destRoot, err := os.OpenRoot(destDir) + if err != nil { + return err + } + defer func() { + _ = destRoot.Close() + }() + file, err := os.Open(archivePath) if err != nil { return err @@ -44,37 +55,40 @@ func extractTarGz(archivePath string, destDir string) error { if err != nil { return err } - target, err := safeExtractPath(destDir, header.Name) + cleanName, err := cleanEntryPath(header.Name) if err != nil { return err } + if cleanName == "." { + continue + } switch header.Typeflag { case tar.TypeDir: - if err := os.MkdirAll(target, 0o755); err != nil { - return err + if err := destRoot.MkdirAll(cleanName, 0o755); err != nil { + return fmt.Errorf("archive directory entry %s escapes destination: %w", header.Name, err) } case tar.TypeSymlink: - if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { - return err + parent := filepath.Dir(cleanName) + if parent != "." { + if err := destRoot.MkdirAll(parent, 0o755); err != nil { + return fmt.Errorf("archive symlink parent %s escapes destination: %w", parent, err) + } } - if filepath.IsAbs(header.Linkname) { - return fmt.Errorf("absolute symlink targets are not supported: %s -> %s", header.Name, header.Linkname) + if err := validateSymlinkTarget(destRoot, parent, header.Linkname); err != nil { + return fmt.Errorf("archive symlink target escapes destination: %s -> %s: %w", header.Name, header.Linkname, err) } - // Verify that the symlink target, when resolved, does not escape destDir. - resolvedTarget := filepath.Join(filepath.Dir(target), header.Linkname) - destDirClean := filepath.Clean(destDir) - if !strings.HasPrefix(resolvedTarget, destDirClean+string(os.PathSeparator)) && resolvedTarget != destDirClean { - return fmt.Errorf("archive symlink target escapes destination: %s -> %s", header.Name, header.Linkname) - } - _ = os.Remove(target) - if err := os.Symlink(header.Linkname, target); err != nil { + _ = destRoot.Remove(cleanName) + if err := destRoot.Symlink(header.Linkname, cleanName); err != nil { return err } case tar.TypeReg: - if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { - return err + parent := filepath.Dir(cleanName) + if parent != "." { + if err := destRoot.MkdirAll(parent, 0o755); err != nil { + return fmt.Errorf("archive file parent %s escapes destination: %w", parent, err) + } } - if err := writeExtractedFile(target, tarReader, fs.FileMode(header.Mode)); err != nil { + if err := writeExtractedFile(destRoot, cleanName, tarReader, fs.FileMode(header.Mode)); err != nil { return err } default: @@ -86,6 +100,17 @@ func extractTarGz(archivePath string, destDir string) error { } func extractZip(archivePath string, destDir string) error { + if err := os.MkdirAll(destDir, 0o755); err != nil { + return err + } + destRoot, err := os.OpenRoot(destDir) + if err != nil { + return err + } + defer func() { + _ = destRoot.Close() + }() + reader, err := zip.OpenReader(archivePath) if err != nil { return err @@ -94,13 +119,16 @@ func extractZip(archivePath string, destDir string) error { _ = reader.Close() }() for _, entry := range reader.File { - target, err := safeExtractPath(destDir, entry.Name) + cleanName, err := cleanEntryPath(entry.Name) if err != nil { return err } + if cleanName == "." { + continue + } if entry.FileInfo().IsDir() { - if err := os.MkdirAll(target, 0o755); err != nil { - return err + if err := destRoot.MkdirAll(cleanName, 0o755); err != nil { + return fmt.Errorf("archive directory entry %s escapes destination: %w", entry.Name, err) } continue } @@ -117,8 +145,11 @@ func extractZip(archivePath string, destDir string) error { if !entry.Mode().IsRegular() { return fmt.Errorf("unsupported archive entry type for %s", entry.Name) } - if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { - return err + parent := filepath.Dir(cleanName) + if parent != "." { + if err := destRoot.MkdirAll(parent, 0o755); err != nil { + return fmt.Errorf("archive file parent %s escapes destination: %w", parent, err) + } } if err := func() error { entryReader, err := entry.Open() @@ -128,7 +159,7 @@ func extractZip(archivePath string, destDir string) error { defer func() { _ = entryReader.Close() }() - return writeExtractedFile(target, entryReader, entry.Mode()) + return writeExtractedFile(destRoot, cleanName, entryReader, entry.Mode()) }(); err != nil { return err } @@ -136,11 +167,12 @@ func extractZip(archivePath string, destDir string) error { return nil } -func writeExtractedFile(target string, source io.Reader, mode fs.FileMode) error { - if mode == 0 { - mode = 0o644 +func writeExtractedFile(destRoot *os.Root, name string, source io.Reader, mode fs.FileMode) error { + perm := mode.Perm() + if perm == 0 { + perm = 0o644 } - out, err := os.OpenFile(target, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, mode) + out, err := destRoot.OpenFile(name, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, perm) if err != nil { return err } @@ -152,22 +184,100 @@ func writeExtractedFile(target string, source io.Reader, mode fs.FileMode) error return closeErr } -// safeExtractPath resolves an archive entry name against destDir, rejecting -// absolute paths or entries that would escape destDir via "..". -func safeExtractPath(destDir string, name string) (string, error) { +// cleanEntryPath resolves an archive entry name, rejecting absolute paths or +// entries that would escape the destination via "..". +func cleanEntryPath(name string) (string, error) { cleanName := filepath.Clean(strings.ReplaceAll(name, "\\", "/")) if cleanName == "." { - return destDir, nil + return ".", nil } - if filepath.IsAbs(cleanName) || cleanName == ".." || strings.HasPrefix(cleanName, "../") { + if filepath.IsAbs(cleanName) || strings.HasPrefix(name, "/") || strings.HasPrefix(name, "\\") || + cleanName == ".." || strings.HasPrefix(cleanName, ".."+string(os.PathSeparator)) || strings.HasPrefix(cleanName, "../") || + filepath.VolumeName(cleanName) != "" || strings.Contains(cleanName, ":") { return "", fmt.Errorf("archive entry escapes destination: %s", name) } - target := filepath.Join(destDir, cleanName) - destDirClean := filepath.Clean(destDir) - if target != destDirClean && !strings.HasPrefix(target, destDirClean+string(os.PathSeparator)) { - return "", fmt.Errorf("archive entry escapes destination: %s", name) + return cleanName, nil +} + +func validateSymlinkTarget(root *os.Root, parentRel, linkTarget string) error { + if filepath.IsAbs(linkTarget) || strings.HasPrefix(linkTarget, "/") || strings.HasPrefix(linkTarget, "\\") || + filepath.VolumeName(linkTarget) != "" || strings.Contains(linkTarget, ":") { + return fmt.Errorf("archive symlink has absolute or invalid target: %s", linkTarget) + } + base := parentRel + if base == "" { + base = "." + } + resolvedBase, err := followUnderRoot(root, base, 0) + if err != nil { + return err + } + _, err = walkUnderRoot(root, resolvedBase, linkTarget, 0) + return err +} + +const maxRootSymlinks = 8 + +func followUnderRoot(root *os.Root, rel string, depth int) (string, error) { + if rel == "" || rel == "." { + return ".", nil + } + return walkUnderRoot(root, ".", rel, depth) +} + +func walkUnderRoot(root *os.Root, base, linkTarget string, depth int) (string, error) { + if depth > maxRootSymlinks { + return "", fmt.Errorf("archive symlink nest exceeds limit") + } + current := base + if current == "" { + current = "." + } + for _, part := range strings.Split(strings.ReplaceAll(linkTarget, "\\", "/"), "/") { + if part == "" || part == "." { + continue + } + if part == ".." { + if current == "." { + return "", fmt.Errorf("archive symlink target escapes destination") + } + current = filepath.Dir(current) + if current == "" { + current = "." + } + continue + } + next := part + if current != "." { + next = filepath.Join(current, part) + } + info, err := root.Lstat(next) + if err != nil { + if os.IsNotExist(err) { + current = next + continue + } + return "", err + } + if info.Mode()&os.ModeSymlink == 0 { + current = next + continue + } + tgt, err := root.Readlink(next) + if err != nil { + return "", err + } + if filepath.IsAbs(tgt) || strings.HasPrefix(tgt, "/") || strings.HasPrefix(tgt, "\\") || + filepath.VolumeName(tgt) != "" || strings.Contains(tgt, ":") { + return "", fmt.Errorf("archive symlink %s has absolute target: %s", next, tgt) + } + followed, err := walkUnderRoot(root, current, tgt, depth+1) + if err != nil { + return "", err + } + current = followed } - return target, nil + return current, nil } // findByBasename recursively searches root for the first regular file whose @@ -175,16 +285,25 @@ func safeExtractPath(destDir string, name string) (string, error) { // helper binaries nested under archive subdirectories (e.g. helpers/) are // still found. func findByBasename(root string, name string) (string, error) { + destRoot, err := os.OpenRoot(root) + if err != nil { + return "", err + } + defer destRoot.Close() var found string - err := filepath.WalkDir(root, func(path string, entry fs.DirEntry, walkErr error) error { + err = fs.WalkDir(destRoot.FS(), ".", func(path string, entry fs.DirEntry, walkErr error) error { if walkErr != nil { return walkErr } if found != "" { return fs.SkipAll } - if !entry.IsDir() && entry.Name() == name { - found = path + if entry.Type().IsRegular() && entry.Name() == name { + if path == "." { + found = filepath.Join(root, name) + } else { + found = filepath.Join(root, filepath.FromSlash(path)) + } } return nil }) diff --git a/internal/update/extract_test.go b/internal/update/extract_test.go index 1742866ea..08bc07875 100644 --- a/internal/update/extract_test.go +++ b/internal/update/extract_test.go @@ -4,6 +4,7 @@ import ( "archive/tar" "archive/zip" "compress/gzip" + "fmt" "os" "path/filepath" "testing" @@ -257,6 +258,38 @@ func TestExtractTarGzAllowsSafeSymlink(t *testing.T) { t.Fatalf("Write target: %v", err) } + // Cover a linked subdirectory with a file extracted through the link. + // Windows Root.Symlink needs the directory target to exist first so the + // link is created as a directory junction rather than a file symlink. + h3 := &tar.Header{ + Name: "subdir", + Typeflag: tar.TypeDir, + Mode: 0o755, + } + if err := tw.WriteHeader(h3); err != nil { + t.Fatalf("WriteHeader subdir: %v", err) + } + h4 := &tar.Header{ + Name: "sublink", + Typeflag: tar.TypeSymlink, + Linkname: "subdir", + } + if err := tw.WriteHeader(h4); err != nil { + t.Fatalf("WriteHeader sublink: %v", err) + } + h5 := &tar.Header{ + Name: "sublink/nested.txt", + Typeflag: tar.TypeReg, + Mode: 0o644, + Size: 6, + } + if err := tw.WriteHeader(h5); err != nil { + t.Fatalf("WriteHeader nested: %v", err) + } + if _, err := tw.Write([]byte("nested")); err != nil { + t.Fatalf("Write nested: %v", err) + } + tw.Close() gw.Close() file.Close() @@ -274,6 +307,14 @@ func TestExtractTarGzAllowsSafeSymlink(t *testing.T) { if gotLink != "target.txt" { t.Fatalf("link target = %q, want %q", gotLink, "target.txt") } + + nestedData, err := os.ReadFile(filepath.Join(destDir, "subdir", "nested.txt")) + if err != nil { + t.Fatalf("ReadFile nested: %v", err) + } + if string(nestedData) != "nested" { + t.Fatalf("nestedData = %q, want %q", string(nestedData), "nested") + } } func TestExtractTarGzRejectsEscapingSymlink(t *testing.T) { @@ -308,3 +349,246 @@ func TestExtractTarGzRejectsEscapingSymlink(t *testing.T) { t.Fatal("expected error extracting escaping symlink") } } + +func TestExtractTarGzRejectsChainedSymlinkEscapingFile(t *testing.T) { + if !symlinksSupported(t) { + t.Skip("symlinks not supported") + } + dir := t.TempDir() + destDir := filepath.Join(dir, "extracted") + outsideDir := filepath.Join(dir, "outside") + if err := os.MkdirAll(destDir, 0o755); err != nil { + t.Fatalf("Mkdir dest: %v", err) + } + if err := os.MkdirAll(outsideDir, 0o755); err != nil { + t.Fatalf("Mkdir outside: %v", err) + } + // chain -> mid -> outsideDir. Lexical path destDir/chain/pwned.txt stays + // under destDir; EvalSymlinks of the chain does not. + if err := os.Symlink(outsideDir, filepath.Join(destDir, "mid")); err != nil { + t.Fatalf("symlink mid: %v", err) + } + if err := os.Symlink("mid", filepath.Join(destDir, "chain")); err != nil { + t.Fatalf("symlink chain: %v", err) + } + + archivePath := filepath.Join(dir, "archive.tar.gz") + writeTestTarGz(t, archivePath, map[string]string{ + "chain/pwned.txt": "escaped", + }) + + if err := extractArchive(archivePath, destDir); err == nil { + t.Fatal("expected error extracting a file through a chained escaping symlink") + } + if _, err := os.Stat(filepath.Join(outsideDir, "pwned.txt")); !os.IsNotExist(err) { + t.Fatalf("escaped file exists outside destDir: %v", err) + } +} + +type testTarEntry struct { + name string + typeflag byte + linkname string + body string +} + +func writeTestTarGzEntries(t *testing.T, archivePath string, entries []testTarEntry) { + t.Helper() + file, err := os.Create(archivePath) + if err != nil { + t.Fatalf("Create archive: %v", err) + } + defer func() { _ = file.Close() }() + gzipWriter := gzip.NewWriter(file) + tarWriter := tar.NewWriter(gzipWriter) + for _, entry := range entries { + header := &tar.Header{ + Name: entry.name, + Typeflag: entry.typeflag, + Linkname: entry.linkname, + Mode: 0o644, + Size: int64(len(entry.body)), + } + if entry.typeflag == tar.TypeSymlink || entry.typeflag == tar.TypeDir { + header.Size = 0 + } + if err := tarWriter.WriteHeader(header); err != nil { + t.Fatalf("WriteHeader %s: %v", entry.name, err) + } + if header.Size > 0 { + if _, err := tarWriter.Write([]byte(entry.body)); err != nil { + t.Fatalf("Write %s: %v", entry.name, err) + } + } + } + if err := tarWriter.Close(); err != nil { + t.Fatalf("close tar writer: %v", err) + } + if err := gzipWriter.Close(); err != nil { + t.Fatalf("close gzip writer: %v", err) + } +} + +// A tar can plant d -> ., then d/s -> .. (physically destDir/s -> .. because +// d is already a link), then l -> d/s/missing, then a regular file named l. +// EvalSymlinks(l) is ENOENT; lexical Join would accept destDir/d/s/missing +// while open follows d and s and writes missing beside destDir. +func TestExtractTarGzRejectsDanglingSymlinkChainEscape(t *testing.T) { + if !symlinksSupported(t) { + t.Skip("symlinks not supported") + } + dir := t.TempDir() + destDir := filepath.Join(dir, "extracted") + archivePath := filepath.Join(dir, "archive.tar.gz") + writeTestTarGzEntries(t, archivePath, []testTarEntry{ + {name: "d", typeflag: tar.TypeSymlink, linkname: "."}, + {name: "d/s", typeflag: tar.TypeSymlink, linkname: ".."}, + {name: "l", typeflag: tar.TypeSymlink, linkname: "d/s/missing"}, + {name: "l", typeflag: tar.TypeReg, body: "pwned"}, + }) + extractErr := extractArchive(archivePath, destDir) + escaped := filepath.Join(dir, "missing") + _, statErr := os.Stat(escaped) + if !os.IsNotExist(statErr) { + t.Fatalf("escaped file exists outside destDir: %v", statErr) + } + if extractErr == nil { + t.Fatal("expected extractArchive to reject a dangling symlink chain that escapes destDir") + } +} + +func TestExtractTarGzAllowsSafeDanglingRelative(t *testing.T) { + if !symlinksSupported(t) { + t.Skip("symlinks not supported") + } + dir := t.TempDir() + destDir := filepath.Join(dir, "extracted") + archivePath := filepath.Join(dir, "archive.tar.gz") + writeTestTarGzEntries(t, archivePath, []testTarEntry{ + {name: "l", typeflag: tar.TypeSymlink, linkname: "not-yet-there"}, + {name: "l", typeflag: tar.TypeReg, body: "safe-content"}, + }) + if err := extractArchive(archivePath, destDir); err != nil { + t.Fatalf("extractArchive: %v", err) + } + data, err := os.ReadFile(filepath.Join(destDir, "not-yet-there")) + if err != nil { + t.Fatalf("ReadFile not-yet-there: %v", err) + } + if string(data) != "safe-content" { + t.Fatalf("not-yet-there content = %q", data) + } + if _, err := os.Stat(filepath.Join(dir, "not-yet-there")); err == nil { + t.Fatal("safe dangling target was written outside destDir") + } +} + +// Tar entry d -> . followed by d/s -> .. creates destDir/s -> .. if the +// symlink entry's link-target check uses lexical filepath.Dir instead of the +// physical parent that will create it. The symlink entry itself must be +// rejected before an outbound link can be planted. +func TestExtractTarGzRejectsSymlinkParentEscape(t *testing.T) { + if !symlinksSupported(t) { + t.Skip("symlinks not supported") + } + dir := t.TempDir() + destDir := filepath.Join(dir, "extracted") + archivePath := filepath.Join(dir, "archive.tar.gz") + writeTestTarGzEntries(t, archivePath, []testTarEntry{ + {name: "d", typeflag: tar.TypeSymlink, linkname: "."}, + {name: "d/s", typeflag: tar.TypeSymlink, linkname: ".."}, + }) + extractErr := extractArchive(archivePath, destDir) + escaped := filepath.Join(destDir, "s") + if _, err := os.Lstat(escaped); err == nil { + t.Fatalf("outbound symlink %s was created", escaped) + } + if extractErr == nil { + t.Fatal("expected extractArchive to reject symlink entry with escaping target through symlink parent") + } +} + +// An archive attempting to plant zero -> d/s/outside-file through a symlink parent +// must not create the target outside destDir and extractArchive must fail. +func TestExtractTarGzRejectsSymlinkThroughSymlinkParentOutsideTarget(t *testing.T) { + if !symlinksSupported(t) { + t.Skip("symlinks not supported") + } + dir := t.TempDir() + destDir := filepath.Join(dir, "extracted") + outsideDir := filepath.Join(dir, "outside") + if err := os.MkdirAll(outsideDir, 0o755); err != nil { + t.Fatalf("MkdirAll outside: %v", err) + } + outsideFile := filepath.Join(outsideDir, "pwned.txt") + archivePath := filepath.Join(dir, "archive.tar.gz") + writeTestTarGzEntries(t, archivePath, []testTarEntry{ + {name: "d", typeflag: tar.TypeSymlink, linkname: "."}, + {name: "d/s", typeflag: tar.TypeSymlink, linkname: ".."}, + {name: "zero", typeflag: tar.TypeSymlink, linkname: "d/s/outside/pwned.txt"}, + }) + extractErr := extractArchive(archivePath, destDir) + if extractErr == nil { + t.Fatal("expected extractArchive to fail on escaping symlink chain") + } + if _, err := os.Stat(outsideFile); !os.IsNotExist(err) { + t.Fatalf("file outside destDir was accessed/created: %v", err) + } +} + +func TestExtractTarGzRejectsIntermediateDirSymlinkEscape(t *testing.T) { + if !symlinksSupported(t) { + t.Skip("symlinks not supported") + } + dir := t.TempDir() + destDir := filepath.Join(dir, "extracted") + outside := filepath.Join(dir, "archive.tar.gz.outside") + if err := os.WriteFile(outside, []byte("secret"), 0o644); err != nil { + t.Fatal(err) + } + archivePath := filepath.Join(dir, "archive.tar.gz") + writeTestTarGzEntries(t, archivePath, []testTarEntry{ + {name: "d", typeflag: tar.TypeSymlink, linkname: "."}, + {name: "d/a", typeflag: tar.TypeDir}, + {name: "d/a/zero", typeflag: tar.TypeSymlink, linkname: "../../archive.tar.gz.outside"}, + }) + if err := extractArchive(archivePath, destDir); err == nil { + t.Fatal("expected reject of ../../ escape through intermediate dir symlink") + } + if _, err := os.Lstat(filepath.Join(destDir, "d", "a", "zero")); err == nil { + t.Fatal("escaping link must not be created") + } +} + +func TestFollowUnderRootRejectsNinthLink(t *testing.T) { + if !symlinksSupported(t) { + t.Skip("symlinks not supported") + } + dir := t.TempDir() + root, err := os.OpenRoot(dir) + if err != nil { + t.Fatal(err) + } + defer root.Close() + if err := os.Mkdir(filepath.Join(dir, "leaf"), 0o755); err != nil { + t.Fatal(err) + } + prev := "leaf" + for i := 0; i < maxRootSymlinks; i++ { + name := fmt.Sprintf("l%d", i) + if err := os.Symlink(prev, filepath.Join(dir, name)); err != nil { + t.Fatal(err) + } + prev = name + } + if _, err := followUnderRoot(root, prev, 0); err != nil { + t.Fatalf("eight links must resolve: %v", err) + } + ninth := "l8" + if err := os.Symlink(prev, filepath.Join(dir, ninth)); err != nil { + t.Fatal(err) + } + if _, err := followUnderRoot(root, ninth, 0); err == nil { + t.Fatal("ninth link must be rejected") + } +}