//go:build !windows package backup import ( "crypto/rand" "encoding/hex" "errors" "io" "os" "path/filepath" "strings" "github.com/aritmolab/thothii/tools/tht/internal/safeio" "golang.org/x/sys/unix" ) type restoreTargetIdentity struct { exists bool device uint64 inode uint64 } func replaceRestoreFile(target string, contents []byte, mode os.FileMode) error { if safeio.ValidateCanonicalPath(target) != nil || mode&os.ModeType != 0 || mode.Perm() == 0 { return safeio.ErrUnsafeFile } components := strings.Split(strings.TrimPrefix(target, string(os.PathSeparator)), string(os.PathSeparator)) if len(components) < 2 || components[0] == "" || components[len(components)-1] == "" { return safeio.ErrUnsafeFile } directory, err := unix.Open(string(os.PathSeparator), unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY, 0) if err != nil { return safeio.ErrUnsafeFile } defer unix.Close(directory) for _, component := range components[:len(components)-1] { next, openErr := unix.Openat(directory, component, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0) if openErr != nil { return safeio.ErrUnsafeFile } unix.Close(directory) directory = next } name := components[len(components)-1] identity, err := inspectRestoreTargetAt(directory, name) if err != nil { return safeio.ErrUnsafeFile } temporary, err := writeRestoreTemporaryAt(directory, contents, mode.Perm()) if err != nil { return safeio.ErrUnsafeFile } defer func() { _ = unix.Unlinkat(directory, temporary, 0) }() current, err := inspectRestoreTargetAt(directory, name) if err != nil || current != identity { return safeio.ErrUnsafeFile } if err := unix.Renameat(directory, temporary, directory, name); err != nil { return safeio.ErrUnsafeFile } temporary = "" if err := unix.Fsync(directory); err != nil { return safeio.ErrUnsafeFile } return nil } func inspectRestoreTargetAt(directory int, name string) (restoreTargetIdentity, error) { var status unix.Stat_t err := unix.Fstatat(directory, name, &status, unix.AT_SYMLINK_NOFOLLOW) if errors.Is(err, unix.ENOENT) { return restoreTargetIdentity{}, nil } if err != nil || status.Mode&unix.S_IFMT != unix.S_IFREG || status.Nlink != 1 { return restoreTargetIdentity{}, safeio.ErrUnsafeFile } return restoreTargetIdentity{exists: true, device: uint64(status.Dev), inode: status.Ino}, nil } func writeRestoreTemporaryAt(directory int, contents []byte, mode os.FileMode) (string, error) { for attempt := 0; attempt < 16; attempt++ { random := make([]byte, 8) if _, err := rand.Read(random); err != nil { return "", err } name := ".tht-restore-" + hex.EncodeToString(random) + ".tmp" descriptor, err := unix.Openat(directory, name, unix.O_WRONLY|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, uint32(mode)) if errors.Is(err, unix.EEXIST) { continue } if err != nil { return "", err } file := os.NewFile(uintptr(descriptor), filepath.Base(name)) if file == nil { unix.Close(descriptor) return "", safeio.ErrUnsafeFile } if err := file.Chmod(mode); err == nil { var written int written, err = file.Write(contents) if err == nil && written != len(contents) { err = io.ErrShortWrite } } if err == nil { err = file.Sync() } closeErr := file.Close() if err == nil { err = closeErr } if err != nil { _ = unix.Unlinkat(directory, name, 0) return "", err } return name, nil } return "", safeio.ErrUnsafeFile }