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
130 changes: 119 additions & 11 deletions component/updater/update_core.go
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,11 @@ type CoreUpdater struct {

var DefaultCoreUpdater = CoreUpdater{}

// preservePosixAttrs is set via init() on linux/darwin to copy owner, mode,
// xattrs and (on darwin) file flags from the existing binary onto the staged
// replacement. Nil on other platforms.
var preservePosixAttrs func(src, dst string)

func (u *CoreUpdater) CoreBaseName() string {
switch runtime.GOARCH {
case "arm":
Expand Down Expand Up @@ -147,6 +152,11 @@ func (u *CoreUpdater) Update(currentExePath string, channel string, force bool)

defer u.clean(updateDir)

err = u.prepareUpdateDir(updateDir)
if err != nil {
return fmt.Errorf("preparing update dir: %w", err)
}

err = u.download(updateDir, packagePath, packageURL)
if err != nil {
return fmt.Errorf("downloading: %w", err)
Expand All @@ -162,7 +172,11 @@ func (u *CoreUpdater) Update(currentExePath string, channel string, force bool)
return fmt.Errorf("backuping: %w", err)
}

err = u.copyFile(updateExePath, currentExePath)
if runtime.GOOS == "windows" {
err = u.copyFile(updateExePath, currentExePath)
} else {
err = u.replaceFileAtomically(updateExePath, currentExePath)
}
if err != nil {
return fmt.Errorf("replacing: %w", err)
}
Expand Down Expand Up @@ -192,6 +206,18 @@ func (u *CoreUpdater) getLatestVersion(versionURL string) (version string, err e
return content, nil
}

func (u *CoreUpdater) prepareUpdateDir(updateDir string) error {
if err := os.RemoveAll(updateDir); err != nil {
return fmt.Errorf("os.RemoveAll(%s): %w", updateDir, err)
}

if err := os.MkdirAll(updateDir, 0o755); err != nil {
return fmt.Errorf("os.MkdirAll(%s): %w", updateDir, err)
}

return nil
}

// download package file and save it to disk
func (u *CoreUpdater) download(updateDir, packagePath, packageURL string) (err error) {
ctx, cancel := context.WithTimeout(context.Background(), time.Second*90)
Expand All @@ -208,12 +234,6 @@ func (u *CoreUpdater) download(updateDir, packagePath, packageURL string) (err e
}
}()

log.Debugln("updateDir %s", updateDir)
err = os.Mkdir(updateDir, 0o755)
if err != nil {
return fmt.Errorf("mkdir error: %w", err)
}

log.Debugln("updater: saving package to file %s", packagePath)
// Create the output file
wc, err := os.OpenFile(packagePath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o755)
Expand Down Expand Up @@ -462,13 +482,101 @@ func (u *CoreUpdater) copyFile(src, dst string) (err error) {
return fmt.Errorf("io.Copy(): %w", err)
}

if runtime.GOOS == "darwin" {
err = exec.Command("/usr/bin/codesign", "--sign", "-", dst).Run()
log.Infoln("updater: copy: %s to %s", src, dst)
return nil
}

func (u *CoreUpdater) replaceFileAtomically(src, dst string) (err error) {
// Resolve symlinks so deployments that expose a stable path (e.g.
// /usr/local/bin/mihomo -> /opt/mihomo/<ver>/mihomo) replace the real
// file rather than turning the symlink into a regular file.
if resolved, linkErr := filepath.EvalSymlinks(dst); linkErr == nil {
dst = resolved
}

rc, err := os.Open(src)
if err != nil {
return fmt.Errorf("os.Open(%s): %w", src, err)
}

defer func() {
closeErr := rc.Close()
if closeErr != nil && err == nil {
err = closeErr
}
}()

srcInfo, err := rc.Stat()
if err != nil {
return fmt.Errorf("rc.Stat(): %w", err)
}

// Prefer the existing dst's mode so local chmod survives updates. Fall
// back to src's mode on first install.
mode := srcInfo.Mode()
dstExists := true
if dstInfo, statErr := os.Stat(dst); statErr == nil {
mode = dstInfo.Mode()
} else if os.IsNotExist(statErr) {
dstExists = false
}

dir := filepath.Dir(dst)
tmp, err := os.CreateTemp(dir, "."+filepath.Base(dst)+".tmp-*")
if err != nil {
return fmt.Errorf("os.CreateTemp(%s): %w", dir, err)
}

tmpPath := tmp.Name()
defer func() {
if err != nil {
log.Warnln("codesign failed: %v", err)
_ = os.Remove(tmpPath)
}
}()

defer func() {
if tmp != nil {
closeErr := tmp.Close()
if closeErr != nil && err == nil {
err = closeErr
}
}
}()

if err = tmp.Chmod(mode); err != nil {
return fmt.Errorf("tmp.Chmod(%s): %w", tmpPath, err)
}

if _, err = io.Copy(tmp, rc); err != nil {
return fmt.Errorf("io.Copy(): %w", err)
}

if err = tmp.Sync(); err != nil {
return fmt.Errorf("tmp.Sync(): %w", err)
}

if err = tmp.Close(); err != nil {
return fmt.Errorf("tmp.Close(): %w", err)
}
tmp = nil

// Carry over ownership, xattrs (notably Linux security.capability for
// cap_net_admin / cap_net_bind_service), setuid/sgid and darwin flags.
if dstExists && preservePosixAttrs != nil {
preservePosixAttrs(dst, tmpPath)
}

if runtime.GOOS == "darwin" {
signErr := exec.Command("/usr/bin/codesign", "--sign", "-", tmpPath).Run()
if signErr != nil {
log.Warnln("codesign failed: %v", signErr)
}
}

log.Infoln("updater: copy: %s to %s", src, dst)
if err = os.Rename(tmpPath, dst); err != nil {
return fmt.Errorf("os.Rename(%s, %s): %w", tmpPath, dst, err)
}

log.Infoln("updater: replace: %s to %s via %s", src, dst, tmpPath)
return nil
}
29 changes: 29 additions & 0 deletions component/updater/update_core_darwin.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
//go:build darwin

package updater

import (
"syscall"

"github.com/metacubex/mihomo/log"

"golang.org/x/sys/unix"
)

func init() {
flagsHook = preserveDarwinFlags
}

// preserveDarwinFlags carries over macOS st_flags (uchg/schg/hidden).
func preserveDarwinFlags(src, dst string) {
var st syscall.Stat_t
if err := syscall.Stat(src, &st); err != nil {
return
}
if st.Flags == 0 {
return
}
if err := unix.Chflags(dst, int(st.Flags)); err != nil {
log.Warnln("updater: chflags %s: %v", dst, err)
}
}
47 changes: 47 additions & 0 deletions component/updater/update_core_darwin_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
//go:build darwin

package updater

import (
"errors"
"os"
"path/filepath"
"syscall"
"testing"

"golang.org/x/sys/unix"
)

func TestReplaceFileAtomicallyPreservesDarwinFlags(t *testing.T) {
root := t.TempDir()
src := filepath.Join(root, "src")
dst := filepath.Join(root, "dst")

if err := os.WriteFile(src, []byte("new-core"), 0o700); err != nil {
t.Fatalf("write src: %v", err)
}

if err := os.WriteFile(dst, []byte("old-core"), 0o755); err != nil {
t.Fatalf("write dst: %v", err)
}

if err := unix.Chflags(dst, unix.UF_HIDDEN); err != nil {
if errors.Is(err, syscall.EPERM) {
t.Skipf("cannot set darwin file flags in test environment: %v", err)
}
t.Fatalf("set darwin flags on dst: %v", err)
}

if err := DefaultCoreUpdater.replaceFileAtomically(src, dst); err != nil {
t.Fatalf("replace file atomically: %v", err)
}

var st syscall.Stat_t
if err := syscall.Stat(dst, &st); err != nil {
t.Fatalf("stat dst: %v", err)
}

if st.Flags&unix.UF_HIDDEN == 0 {
t.Fatalf("expected UF_HIDDEN to be preserved, flags=%#x", st.Flags)
}
}
Loading