Files
ThothII/tools/tht/internal/safeio/replace_unix.go
T

115 lines
2.9 KiB
Go

//go:build !windows
package safeio
import (
"errors"
"io"
"os"
"strings"
"golang.org/x/sys/unix"
)
func replaceCanonicalRegular(path string, contents []byte) error {
directory, target, err := openCanonicalParentDirectory(path)
if err != nil {
return ErrUnsafeFile
}
defer unix.Close(directory)
if err := requireSingleRegularAt(directory, target); err != nil {
return ErrUnsafeFile
}
temporary, err := writePrivateTemporaryAt(directory, contents)
if err != nil {
return ErrUnsafeFile
}
defer func() {
if temporary != "" {
_ = unix.Unlinkat(directory, temporary, 0)
}
}()
if err := requireSingleRegularAt(directory, target); err != nil {
return ErrUnsafeFile
}
if err := unix.Renameat(directory, temporary, directory, target); err != nil {
return ErrUnsafeFile
}
temporary = ""
if err := unix.Fsync(directory); err != nil {
return ErrUnsafeFile
}
return nil
}
func openCanonicalParentDirectory(path string) (int, string, error) {
components := strings.Split(strings.TrimPrefix(path, string(os.PathSeparator)), string(os.PathSeparator))
if len(components) < 2 || components[0] == "" || components[len(components)-1] == "" {
return -1, "", ErrUnsafeFile
}
directory, err := unix.Open(string(os.PathSeparator), unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY, 0)
if err != nil {
return -1, "", err
}
for _, component := range components[:len(components)-1] {
next, err := unix.Openat(directory, component, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0)
if err != nil {
unix.Close(directory)
return -1, "", err
}
unix.Close(directory)
directory = next
}
return directory, components[len(components)-1], nil
}
func requireSingleRegularAt(directory int, name string) error {
var stat unix.Stat_t
if err := unix.Fstatat(directory, name, &stat, unix.AT_SYMLINK_NOFOLLOW); err != nil || stat.Mode&unix.S_IFMT != unix.S_IFREG || stat.Nlink != 1 {
return ErrUnsafeFile
}
return nil
}
func writePrivateTemporaryAt(directory int, contents []byte) (string, error) {
for attempt := 0; attempt < 16; attempt++ {
name, err := randomTemporaryName()
if err != nil {
return "", err
}
descriptor, err := unix.Openat(directory, name, unix.O_WRONLY|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0o600)
if errors.Is(err, unix.EEXIST) {
continue
}
if err != nil {
return "", err
}
file := os.NewFile(uintptr(descriptor), "tht-safeio-replacement")
if file == nil {
unix.Close(descriptor)
return "", ErrUnsafeFile
}
if err := file.Chmod(0o600); 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 "", ErrUnsafeFile
}