ghupdate: add checksum verification and extraction path containment guard

This commit is contained in:
henrygd
2026-08-14 19:23:08 -04:00
parent d50c09176f
commit 052489cada
6 changed files with 195 additions and 15 deletions
+43
View File
@@ -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
}
+57
View File
@@ -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)
}
})
}
}
+19 -11
View File
@@ -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()
+16 -4
View File
@@ -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 {
+59
View File
@@ -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)
}
}
+1
View File
@@ -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"`
}