Files
ThothII/tools/tht/internal/backup/restore_file_unix.go
T

148 lines
4.1 KiB
Go

//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
uid uint32
gid uint32
}
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
}
uid, gid, err := restoreTargetOwnerAt(directory, identity)
if err != nil {
return safeio.ErrUnsafeFile
}
temporary, err := writeRestoreTemporaryAt(directory, contents, mode.Perm(), uid, gid)
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,
uid: status.Uid,
gid: status.Gid,
}, nil
}
func restoreTargetOwnerAt(directory int, identity restoreTargetIdentity) (uint32, uint32, error) {
if identity.exists {
return identity.uid, identity.gid, nil
}
var status unix.Stat_t
if err := unix.Fstat(directory, &status); err != nil || status.Mode&unix.S_IFMT != unix.S_IFDIR {
return 0, 0, safeio.ErrUnsafeFile
}
return status.Uid, status.Gid, nil
}
func writeRestoreTemporaryAt(directory int, contents []byte, mode os.FileMode, uid, gid uint32) (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.Chown(int(uid), int(gid)); err == nil {
err = file.Chmod(mode)
}
if 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
}