Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 7 additions & 7 deletions pkg/extract/extract_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -96,7 +96,7 @@ func TestExtract_NormalArchive(t *testing.T) {
func TestExtract_PathTraversalBlocked(t *testing.T) {
t.Parallel()
buf := newTarGz(t, []tarEntry{
{name: "../../etc/passwd", body: "malicious"},
{name: etcPasswd, body: "malicious"},
})

dest := t.TempDir()
Expand All @@ -114,7 +114,7 @@ func TestExtract_SymlinkTraversalBlocked(t *testing.T) {
buf := newTarGz(t, []tarEntry{
{
name: "evil-link",
linkTarget: "../../etc/passwd",
linkTarget: etcPasswd,
symlink: true,
},
})
Expand All @@ -134,7 +134,7 @@ func TestExtract_HardLinkTraversalBlocked(t *testing.T) {
buf := newTarGz(t, []tarEntry{
{
name: "evil-link",
linkTarget: "../../etc/passwd",
linkTarget: etcPasswd,
symlink: false,
},
})
Expand All @@ -152,10 +152,10 @@ func TestExtract_HardLinkTraversalBlocked(t *testing.T) {
func TestExtract_ValidSymlinkAllowed(t *testing.T) {
t.Parallel()
buf := newTarGz(t, []tarEntry{
{name: "target.txt", body: "content"},
{name: targetFileName, body: "content"},
{
name: "link.txt",
linkTarget: "target.txt",
linkTarget: targetFileName,
symlink: true,
},
})
Expand All @@ -170,7 +170,7 @@ func TestExtract_ValidSymlinkAllowed(t *testing.T) {
if err != nil {
t.Fatalf("readlink: %v", err)
}
if target != "target.txt" {
t.Fatalf("symlink target = %q, want %q", target, "target.txt")
if target != targetFileName {
t.Fatalf("symlink target = %q, want %q", target, targetFileName)
}
}
131 changes: 131 additions & 0 deletions pkg/extract/path_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,131 @@
package extract

import (
"archive/tar"
"strings"
"testing"
)

const (
targetFileName = "target.txt"
etcPasswd = "../../etc/passwd" // #nosec G101
)

func TestWithinDir(t *testing.T) {
t.Parallel()
const dest = "/tmp/dest"
cases := []struct {
name string
resolved string
want bool
}{
{"nested inside", dest + "/sub/file", true},
{"exactly dest", dest, true},
{"deep inside", dest + "/a/b/c", true},
{"sibling prefix outside", dest + "_evil/file", false},
{"unrelated outside", "/tmp/other/file", false},
{"parent dir", "/tmp", false},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
if got := withinDir(tt.resolved, dest); got != tt.want {
t.Errorf("withinDir(%q, %q) = %v, want %v", tt.resolved, dest, got, tt.want)
}
})
}
}

func TestResolveLinkTarget(t *testing.T) {
t.Parallel()
cases := []struct {
name string
linkname string
outFile string
want string
}{
{"absolute target preserved", "/etc/passwd", "/tmp/dest/link", "/etc/passwd"},
{"relative parent", "../target", "/tmp/dest/sub/link", "/tmp/dest/target"},
{"relative sibling", targetFileName, "/tmp/dest/link", "/tmp/dest/target.txt"},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
if got := resolveLinkTarget(tt.linkname, tt.outFile); got != tt.want {
t.Errorf(
"resolveLinkTarget(%q, %q) = %q, want %q",
tt.linkname,
tt.outFile,
got,
tt.want,
)
}
})
}
}

func TestResolveRelativePath_StripLevels(t *testing.T) {
t.Parallel()
cases := []struct {
name string
header *tar.Header
strip int
want string
}{
{"no strip keeps path", &tar.Header{Name: "dir/file"}, 0, "/dir/file"},
{"strip one level", &tar.Header{Name: "a/b/c"}, 1, "/b/c"},
{"strip two levels", &tar.Header{Name: "a/b/c"}, 2, "/c"},
{"strip beyond depth keeps remainder", &tar.Header{Name: "a"}, 2, "/a"},
{"strip to final segment", &tar.Header{Name: "a/b"}, 1, "/b"},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
opts := &Options{StripLevels: tt.strip}
if got := resolveRelativePath(tt.header, opts); got != tt.want {
t.Errorf(
"resolveRelativePath(%q, strip=%d) = %q, want %q",
tt.header.Name,
tt.strip,
got,
tt.want,
)
}
})
}
}

func TestValidateLinkTarget(t *testing.T) {
t.Parallel()
const dest = "/tmp/dest"
cases := []struct {
name string
typeflag byte
linkname string
wantError string
}{
{"symlink inside allowed", tar.TypeSymlink, targetFileName, ""},
{"symlink outside blocked", tar.TypeSymlink, etcPasswd, "symlink traversal"},
{"hard link outside blocked", tar.TypeLink, etcPasswd, "hard link traversal"},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
header := &tar.Header{Typeflag: tt.typeflag, Name: "link", Linkname: tt.linkname}
outFile := dest + "/link"
err := validateLinkTarget(header, outFile, dest)
if tt.wantError == "" {
if err != nil {
t.Errorf("expected no error, got %v", err)
}
return
}
if err == nil {
t.Fatal("expected error, got nil")
}
if !strings.Contains(err.Error(), tt.wantError) {
t.Errorf("error %q does not mention %q", err, tt.wantError)
}
})
}
}
Loading