From 052489cadad583db7f3e730c18cd207cfc12709e Mon Sep 17 00:00:00 2001 From: henrygd Date: Fri, 14 Aug 2026 17:12:11 -0400 Subject: [PATCH] ghupdate: add checksum verification and extraction path containment guard --- internal/ghupdate/checksum.go | 43 ++++++++++++++++++++++ internal/ghupdate/checksum_test.go | 57 +++++++++++++++++++++++++++++ internal/ghupdate/extract.go | 30 +++++++++------ internal/ghupdate/ghupdate.go | 20 ++++++++-- internal/ghupdate/ghupdate_test.go | 59 ++++++++++++++++++++++++++++++ internal/ghupdate/release.go | 1 + 6 files changed, 195 insertions(+), 15 deletions(-) create mode 100644 internal/ghupdate/checksum.go create mode 100644 internal/ghupdate/checksum_test.go diff --git a/internal/ghupdate/checksum.go b/internal/ghupdate/checksum.go new file mode 100644 index 00000000..4af61057 --- /dev/null +++ b/internal/ghupdate/checksum.go @@ -0,0 +1,43 @@ +package ghupdate + +import ( + "bytes" + "crypto/sha256" + "encoding/hex" + "fmt" + "io" + "os" + "strings" +) + +func verifyAssetChecksum(path, digest string) error { + algorithm, expectedHex, ok := strings.Cut(digest, ":") + if !ok || algorithm == "" || expectedHex == "" { + return fmt.Errorf("invalid release digest %q", digest) + } + if !strings.EqualFold(algorithm, "sha256") { + return fmt.Errorf("unsupported release digest algorithm %q", algorithm) + } + + expected, err := hex.DecodeString(expectedHex) + if err != nil || len(expected) != sha256.Size { + return fmt.Errorf("invalid SHA-256 release digest %q", digest) + } + + file, err := os.Open(path) + if err != nil { + return fmt.Errorf("failed to open release for checksum verification: %w", err) + } + defer file.Close() + + hash := sha256.New() + if _, err := io.Copy(hash, file); err != nil { + return fmt.Errorf("failed to calculate release checksum: %w", err) + } + actual := hash.Sum(nil) + if !bytes.Equal(actual, expected) { + return fmt.Errorf("release checksum mismatch: expected %s, got %s", expectedHex, hex.EncodeToString(actual)) + } + + return nil +} diff --git a/internal/ghupdate/checksum_test.go b/internal/ghupdate/checksum_test.go new file mode 100644 index 00000000..177ca4f4 --- /dev/null +++ b/internal/ghupdate/checksum_test.go @@ -0,0 +1,57 @@ +package ghupdate + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func TestVerifyAssetChecksum(t *testing.T) { + path := filepath.Join(t.TempDir(), "asset") + if err := os.WriteFile(path, []byte("beszel release asset"), 0600); err != nil { + t.Fatal(err) + } + + tests := []struct { + name string + digest string + wantErr string + }{ + { + name: "valid", + digest: "sha256:2316f86af2c3af2f0ef595ad5359cdd19b329d4829a5425460ca9ffbf92671ab", + }, + { + name: "mismatch", + digest: "sha256:0316f86af2c3af2f0ef595ad5359cdd19b329d4829a5425460ca9ffbf92671ab", + wantErr: "checksum mismatch", + }, + { + name: "malformed", + digest: "sha256:not-a-checksum", + wantErr: "invalid SHA-256", + }, + { + name: "missing", + wantErr: "invalid release asset digest", + }, + { + name: "unsupported algorithm", + digest: "sha512:2316f86af2c3af2f0ef595ad5359cdd19b329d4829a5425460ca9ffbf92671ab", + wantErr: "unsupported", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := verifyAssetChecksum(path, tt.digest) + if tt.wantErr == "" && err != nil { + t.Fatalf("expected checksum to verify, got %v", err) + } + if tt.wantErr != "" && (err == nil || !strings.Contains(err.Error(), tt.wantErr)) { + t.Fatalf("expected error containing %q, got %v", tt.wantErr, err) + } + }) + } +} diff --git a/internal/ghupdate/extract.go b/internal/ghupdate/extract.go index 38da6bb5..0d3b8a37 100644 --- a/internal/ghupdate/extract.go +++ b/internal/ghupdate/extract.go @@ -46,18 +46,23 @@ func extractTarGz(srcPath, destDir string) error { return err } + path, err := archivePath(destDir, header.Name) + if err != nil { + return err + } + if header.Typeflag == tar.TypeDir { - if err := os.MkdirAll(filepath.Join(destDir, header.Name), 0755); err != nil { + if err := os.MkdirAll(path, 0755); err != nil { return err } continue } - if err := os.MkdirAll(filepath.Dir(filepath.Join(destDir, header.Name)), 0755); err != nil { + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { return err } - outFile, err := os.Create(filepath.Join(destDir, header.Name)) + outFile, err := os.Create(path) if err != nil { return err } @@ -72,6 +77,14 @@ func extractTarGz(srcPath, destDir string) error { return nil } +// archivePath returns a path within destDir, rejecting path traversal entries. +func archivePath(destDir, name string) (string, error) { + if !filepath.IsLocal(name) { + return "", fmt.Errorf("invalid file path: %q", name) + } + return filepath.Join(destDir, name), nil +} + // extractZip extracts the zip archive at "src" to "dest". // // Note that only dirs and regular files will be extracted. @@ -84,9 +97,6 @@ func extractZip(src, dest string) error { } defer zr.Close() - // normalize dest path to check later for Zip Slip - dest = filepath.Clean(dest) + string(os.PathSeparator) - for _, f := range zr.File { err := extractFile(f, dest) if err != nil { @@ -100,11 +110,9 @@ func extractZip(src, dest string) error { // extractFile extracts the provided zipFile into "basePath/zipFileName" path, // creating all the necessary path directories. func extractFile(zipFile *zip.File, basePath string) error { - path := filepath.Join(basePath, zipFile.Name) - - // check for Zip Slip - if !strings.HasPrefix(path, basePath) { - return fmt.Errorf("invalid file path: %s", path) + path, err := archivePath(basePath, zipFile.Name) + if err != nil { + return err } r, err := zipFile.Open() diff --git a/internal/ghupdate/ghupdate.go b/internal/ghupdate/ghupdate.go index deb61b01..f5baf0d6 100644 --- a/internal/ghupdate/ghupdate.go +++ b/internal/ghupdate/ghupdate.go @@ -135,21 +135,33 @@ func (p *updater) update() (updated bool, err error) { return false, err } - releaseDir := filepath.Join(p.config.DataDir, ".beszel_update") + if err := os.MkdirAll(p.config.DataDir, 0755); err != nil { + return false, fmt.Errorf("failed to create update data directory: %w", err) + } + releaseDir, err := os.MkdirTemp(p.config.DataDir, ".beszel_update-") + if err != nil { + return false, fmt.Errorf("failed to create update directory: %w", err) + } defer os.RemoveAll(releaseDir) ColorPrintf(ColorYellow, "Downloading %s...", asset.Name) // download the release asset - assetPath := filepath.Join(releaseDir, asset.Name) + assetPath, err := archivePath(releaseDir, asset.Name) + if err != nil { + return false, err + } if err := downloadFile(p.config.Context, p.config.HttpClient, asset.DownloadUrl, assetPath, p.config.UseMirror); err != nil { return false, err } + ColorPrint(ColorYellow, "Verifying checksum...") + if err := verifyAssetChecksum(assetPath, asset.Digest); err != nil { + return false, err + } ColorPrintf(ColorYellow, "Extracting %s...", asset.Name) - extractDir := filepath.Join(releaseDir, "extracted_"+asset.Name) - defer os.RemoveAll(extractDir) + extractDir := filepath.Join(releaseDir, "extracted") // Extract the archive (automatically detects format) if err := extract(assetPath, extractDir); err != nil { diff --git a/internal/ghupdate/ghupdate_test.go b/internal/ghupdate/ghupdate_test.go index a93b1029..f23b5750 100644 --- a/internal/ghupdate/ghupdate_test.go +++ b/internal/ghupdate/ghupdate_test.go @@ -1,6 +1,9 @@ package ghupdate import ( + "archive/tar" + "compress/gzip" + "os" "path/filepath" "testing" ) @@ -43,3 +46,59 @@ func TestExtractFailure(t *testing.T) { t.Fatal("Expected Extract to fail due to missing tar.gz file") } } + +func TestArchivePath(t *testing.T) { + destDir := t.TempDir() + for _, name := range []string{ + "", + "..", + filepath.Join("..", "file"), + filepath.Join("dir", "..", "..", "file"), + string(os.PathSeparator) + filepath.Join("tmp", "file"), + } { + if _, err := archivePath(destDir, name); err == nil { + t.Errorf("expected %q to be rejected", name) + } + } + + name := filepath.Join("dir", "file") + if path, err := archivePath(destDir, name); err != nil || path != filepath.Join(destDir, name) { + t.Errorf("archivePath(%q) = %q, %v", name, path, err) + } +} + +func TestExtractTarGzRejectsPathTraversal(t *testing.T) { + testDir := t.TempDir() + archivePath := filepath.Join(testDir, "malicious.tar.gz") + destDir := filepath.Join(testDir, "extract") + escapedPath := filepath.Join(testDir, "escaped") + + archive, err := os.Create(archivePath) + if err != nil { + t.Fatal(err) + } + gz := gzip.NewWriter(archive) + tw := tar.NewWriter(gz) + if err := tw.WriteHeader(&tar.Header{Name: "../escaped", Mode: 0600, Size: 1}); err != nil { + t.Fatal(err) + } + if _, err := tw.Write([]byte("x")); err != nil { + t.Fatal(err) + } + if err := tw.Close(); err != nil { + t.Fatal(err) + } + if err := gz.Close(); err != nil { + t.Fatal(err) + } + if err := archive.Close(); err != nil { + t.Fatal(err) + } + + if err := extract(archivePath, destDir); err == nil { + t.Fatal("expected path traversal archive to be rejected") + } + if _, err := os.Stat(escapedPath); !os.IsNotExist(err) { + t.Fatalf("path traversal wrote %s", escapedPath) + } +} diff --git a/internal/ghupdate/release.go b/internal/ghupdate/release.go index 2cd84fcd..640d46a1 100644 --- a/internal/ghupdate/release.go +++ b/internal/ghupdate/release.go @@ -8,6 +8,7 @@ import ( type releaseAsset struct { Name string `json:"name"` DownloadUrl string `json:"browser_download_url"` + Digest string `json:"digest"` Id int `json:"id"` Size int `json:"size"` }