diff --git a/go.mod b/go.mod index d0da3cf3..d7bf6a16 100644 --- a/go.mod +++ b/go.mod @@ -10,6 +10,7 @@ require ( github.com/go-chi/render v1.0.3 github.com/golang-jwt/jwt/v5 v5.3.0 github.com/gomarkdown/markdown v0.0.0-20250810172220-2e2c11897d1a + github.com/google/renameio v1.0.1 github.com/spf13/cobra v1.10.1 github.com/stretchr/testify v1.11.1 ) diff --git a/go.sum b/go.sum index d2f6bbef..ed710877 100644 --- a/go.sum +++ b/go.sum @@ -17,6 +17,8 @@ github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9v github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/gomarkdown/markdown v0.0.0-20250810172220-2e2c11897d1a h1:l7A0loSszR5zHd/qK53ZIHMO8b3bBSmENnQ6eKnUT0A= github.com/gomarkdown/markdown v0.0.0-20250810172220-2e2c11897d1a/go.mod h1:JDGcbDT52eL4fju3sZ4TeHGsQwhG9nbDV21aMyhwPoA= +github.com/google/renameio v1.0.1 h1:Lh/jXZmvZxb0BBeSY5VKEfidcbcbenKjZFzM/q0fSeU= +github.com/google/renameio v1.0.1/go.mod h1:t/HQoYBZSsWSNK35C6CO/TpPLDVWvxOHboWUAweKUpk= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= diff --git a/server/cmd/upgrade.go b/server/cmd/upgrade.go index f75dd5d4..41ce08fa 100644 --- a/server/cmd/upgrade.go +++ b/server/cmd/upgrade.go @@ -10,6 +10,7 @@ import ( "path/filepath" "runtime" + "github.com/google/renameio" "github.com/spf13/cobra" ) @@ -94,7 +95,7 @@ func upgrade(urlPrefix string) error { fmt.Printf("Now going to replace the existing silverbullet binary in %s\n", installDir) // Extract the zip file - err = extractZip(zipPath, installDir) + err = extractZip(tmpDir, zipPath, installDir) if err != nil { return fmt.Errorf("failed to extract zip: %w", err) } @@ -113,7 +114,7 @@ func upgrade(urlPrefix string) error { return nil } -func extractZip(src, dest string) error { +func extractZip(tmpDir, src, dest string) error { reader, err := zip.OpenReader(src) if err != nil { return err @@ -121,36 +122,49 @@ func extractZip(src, dest string) error { defer reader.Close() for _, file := range reader.File { + path := filepath.Join(dest, file.Name) + + // Handle directories first, report errors here too + if file.FileInfo().IsDir() { + if err := os.MkdirAll(path, file.FileInfo().Mode()); err != nil { + return err + } + continue + } + rc, err := file.Open() if err != nil { return err } defer rc.Close() - path := filepath.Join(dest, file.Name) - - // Check for directory - if file.FileInfo().IsDir() { - os.MkdirAll(path, file.FileInfo().Mode()) - continue - } - // Create parent directories if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { return err } - // Create the file - outFile, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, file.FileInfo().Mode()) + // Create the file as a temporary file first + tempFile, err := renameio.TempFile(tmpDir, path) if err != nil { return err } - defer outFile.Close() + defer tempFile.Cleanup() - _, err = io.Copy(outFile, rc) + // Extract to the temporary file + _, err = io.Copy(tempFile, rc) if err != nil { return err } + + // Properly set the mode of the target file + if err := tempFile.Chmod(file.FileInfo().Mode()); err != nil { + return err + } + + // Atomically create the target file + if err := tempFile.CloseAtomicallyReplace(); err != nil { + return err + } } return nil