diff --git a/server/lib/ziputil/ziputil.go b/server/lib/ziputil/ziputil.go index 980097bb..3a857f09 100644 --- a/server/lib/ziputil/ziputil.go +++ b/server/lib/ziputil/ziputil.go @@ -53,6 +53,17 @@ func ZipDir(sourceDir string) ([]byte, error) { return nil } + if info.Mode()&os.ModeSymlink != 0 { + target, err := os.Readlink(path) + if err != nil { + return fmt.Errorf("read symlink %s: %w", path, err) + } + if _, err := writer.Write([]byte(target)); err != nil { + return fmt.Errorf("write symlink %s: %w", path, err) + } + return nil + } + // Only include regular files. Skip sockets, devices, FIFOs, etc. if !info.Mode().IsRegular() { return nil @@ -93,13 +104,31 @@ func Unzip(zipFilePath, destDir string) error { if err := os.MkdirAll(destDir, 0755); err != nil { return fmt.Errorf("failed to create destination directory: %w", err) } + cleanDestDir, err := filepath.Abs(destDir) + if err != nil { + return fmt.Errorf("failed to resolve destination directory: %w", err) + } + cleanDestDir, err = filepath.EvalSymlinks(cleanDestDir) + if err != nil { + return fmt.Errorf("failed to evaluate destination directory: %w", err) + } + // Extract each file for _, file := range reader.File { + entryPath := filepath.FromSlash(file.Name) + // Create the full destination path - destPath := filepath.Join(destDir, file.Name) + destPath := filepath.Join(cleanDestDir, entryPath) // Check for directory traversal vulnerabilities - if !strings.HasPrefix(destPath, filepath.Clean(destDir)+string(os.PathSeparator)) { + if destPath == cleanDestDir || !isPathWithinDir(cleanDestDir, destPath) { + return fmt.Errorf("illegal file path: %s", file.Name) + } + resolvedParentPath, err := resolvePathWithSymlinks(cleanDestDir, filepath.Dir(entryPath)) + if err != nil { + return fmt.Errorf("failed to resolve destination path %s: %w", file.Name, err) + } + if !isPathWithinDir(cleanDestDir, resolvedParentPath) { return fmt.Errorf("illegal file path: %s", file.Name) } @@ -123,6 +152,36 @@ func Unzip(zipFilePath, destDir string) error { } defer fileReader.Close() + if file.Mode()&os.ModeSymlink != 0 { + target, err := io.ReadAll(fileReader) + if err != nil { + return fmt.Errorf("failed to read symlink target: %w", err) + } + targetPath := string(target) + if !filepath.IsAbs(targetPath) { + resolvedTarget, err := resolvePathWithSymlinks(resolvedParentPath, targetPath) + if err != nil { + return fmt.Errorf("failed to resolve symlink target: %w", err) + } + if !isPathWithinDir(cleanDestDir, resolvedTarget) { + return fmt.Errorf("illegal symlink target: %s -> %s", file.Name, targetPath) + } + } + if err := os.Remove(destPath); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("failed to remove existing symlink path: %w", err) + } + if err := os.Symlink(targetPath, destPath); err != nil { + return fmt.Errorf("failed to create symlink: %w", err) + } + continue + } + + if info, err := os.Lstat(destPath); err == nil && info.Mode()&os.ModeSymlink != 0 { + if err := os.Remove(destPath); err != nil { + return fmt.Errorf("failed to remove existing symlink: %w", err) + } + } + // Create the destination file destFile, err := os.OpenFile(destPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, file.Mode()) if err != nil { @@ -138,3 +197,33 @@ func Unzip(zipFilePath, destDir string) error { return nil } + +func isPathWithinDir(dir, path string) bool { + return path == dir || strings.HasPrefix(path, dir+string(os.PathSeparator)) +} + +func resolvePathWithSymlinks(baseDir, relativePath string) (string, error) { + currentPath := filepath.Clean(baseDir) + for _, part := range strings.Split(filepath.FromSlash(relativePath), string(os.PathSeparator)) { + switch part { + case "", ".": + continue + case "..": + currentPath = filepath.Dir(currentPath) + continue + } + + nextPath := filepath.Join(currentPath, part) + resolvedPath, err := filepath.EvalSymlinks(nextPath) + if err == nil { + currentPath = resolvedPath + continue + } + if !os.IsNotExist(err) { + return "", fmt.Errorf("evaluate symlinks for %s: %w", nextPath, err) + } + currentPath = nextPath + } + + return filepath.Clean(currentPath), nil +} diff --git a/server/lib/ziputil/ziputil_test.go b/server/lib/ziputil/ziputil_test.go index 832ada59..0f357f62 100644 --- a/server/lib/ziputil/ziputil_test.go +++ b/server/lib/ziputil/ziputil_test.go @@ -10,6 +10,165 @@ import ( "github.com/stretchr/testify/require" ) +func TestZipDirPreservesSymlinks(t *testing.T) { + sourceDir := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(sourceDir, "target.txt"), []byte("target contents"), 0644)) + require.NoError(t, os.Symlink("target.txt", filepath.Join(sourceDir, "link.txt"))) + + zipContent, err := ZipDir(sourceDir) + require.NoError(t, err) + + zipFile, err := os.CreateTemp(t.TempDir(), "archive-*.zip") + require.NoError(t, err) + _, err = zipFile.Write(zipContent) + require.NoError(t, err) + require.NoError(t, zipFile.Close()) + + destDir := t.TempDir() + require.NoError(t, Unzip(zipFile.Name(), destDir)) + + linkPath := filepath.Join(destDir, "link.txt") + info, err := os.Lstat(linkPath) + require.NoError(t, err) + assert.True(t, info.Mode()&os.ModeSymlink != 0) + target, err := os.Readlink(linkPath) + require.NoError(t, err) + assert.Equal(t, "target.txt", target) +} + +func TestUnzipRejectsEscapingSymlink(t *testing.T) { + zipPath := createSymlinkZip(t, "../outside.txt") + + err := Unzip(zipPath, t.TempDir()) + require.Error(t, err) + assert.Contains(t, err.Error(), "illegal symlink target") +} + +func TestUnzipRejectsSymlinkChainEscape(t *testing.T) { + zipPath := createSymlinkChainEscapeZip(t) + destDir := filepath.Join(t.TempDir(), "extract") + + err := Unzip(zipPath, destDir) + require.Error(t, err) + assert.Contains(t, err.Error(), "illegal symlink target") +} + +func TestUnzipPreservesAbsoluteSymlink(t *testing.T) { + target := filepath.Join(t.TempDir(), "target.txt") + zipPath := createSymlinkZip(t, target) + destDir := t.TempDir() + + require.NoError(t, Unzip(zipPath, destDir)) + actualTarget, err := os.Readlink(filepath.Join(destDir, "link.txt")) + require.NoError(t, err) + assert.Equal(t, target, actualTarget) +} + +func TestUnzipRejectsRootEntry(t *testing.T) { + zipPath := createNamedSymlinkZip(t, ".", "target.txt") + destDir := t.TempDir() + + err := Unzip(zipPath, destDir) + require.Error(t, err) + assert.Contains(t, err.Error(), "illegal file path") +} + +func TestUnzipRejectsEntryUnderExistingSymlink(t *testing.T) { + zipPath := createFileZip(t, "link/file.txt") + destDir := t.TempDir() + outsideDir := t.TempDir() + require.NoError(t, os.Symlink(outsideDir, filepath.Join(destDir, "link"))) + + err := Unzip(zipPath, destDir) + require.Error(t, err) + assert.Contains(t, err.Error(), "illegal file path") + require.NoFileExists(t, filepath.Join(outsideDir, "file.txt")) +} + +func TestUnzipOverwritesFileWithSymlink(t *testing.T) { + zipPath := createSymlinkZip(t, "target.txt") + destDir := t.TempDir() + linkPath := filepath.Join(destDir, "link.txt") + require.NoError(t, os.WriteFile(linkPath, []byte("old contents"), 0644)) + + require.NoError(t, Unzip(zipPath, destDir)) + + info, err := os.Lstat(linkPath) + require.NoError(t, err) + assert.True(t, info.Mode()&os.ModeSymlink != 0) + target, err := os.Readlink(linkPath) + require.NoError(t, err) + assert.Equal(t, "target.txt", target) +} + +func createSymlinkZip(t *testing.T, target string) string { + t.Helper() + return createNamedSymlinkZip(t, "link.txt", target) +} + +func createNamedSymlinkZip(t *testing.T, name, target string) string { + t.Helper() + + zipPath := filepath.Join(t.TempDir(), "symlink.zip") + zipFile, err := os.Create(zipPath) + require.NoError(t, err) + + zipWriter := zip.NewWriter(zipFile) + header := &zip.FileHeader{Name: name, Method: zip.Store} + header.SetMode(os.ModeSymlink | 0777) + writer, err := zipWriter.CreateHeader(header) + require.NoError(t, err) + _, err = writer.Write([]byte(target)) + require.NoError(t, err) + require.NoError(t, zipWriter.Close()) + require.NoError(t, zipFile.Close()) + + return zipPath +} + +func createFileZip(t *testing.T, name string) string { + t.Helper() + + zipPath := filepath.Join(t.TempDir(), "file.zip") + zipFile, err := os.Create(zipPath) + require.NoError(t, err) + zipWriter := zip.NewWriter(zipFile) + writer, err := zipWriter.Create(name) + require.NoError(t, err) + _, err = writer.Write([]byte("contents")) + require.NoError(t, err) + require.NoError(t, zipWriter.Close()) + require.NoError(t, zipFile.Close()) + return zipPath +} + +func createSymlinkChainEscapeZip(t *testing.T) string { + t.Helper() + + zipPath := filepath.Join(t.TempDir(), "chain-escape.zip") + zipFile, err := os.Create(zipPath) + require.NoError(t, err) + + zipWriter := zip.NewWriter(zipFile) + linkHeader := &zip.FileHeader{Name: "link", Method: zip.Store} + linkHeader.SetMode(os.ModeSymlink | 0777) + linkWriter, err := zipWriter.CreateHeader(linkHeader) + require.NoError(t, err) + _, err = linkWriter.Write([]byte(".")) + require.NoError(t, err) + + escapeHeader := &zip.FileHeader{Name: "escape", Method: zip.Store} + escapeHeader.SetMode(os.ModeSymlink | 0777) + escapeWriter, err := zipWriter.CreateHeader(escapeHeader) + require.NoError(t, err) + _, err = escapeWriter.Write([]byte("link/..")) + require.NoError(t, err) + + require.NoError(t, zipWriter.Close()) + require.NoError(t, zipFile.Close()) + return zipPath +} + func TestUnzipFile(t *testing.T) { // Create a temporary directory for test files sourceDir, err := os.MkdirTemp("", "zip-source-*")