//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 removeCanonicalPrivateRegular(path string) error { directory, target, err := openCanonicalParentDirectory(path) if err != nil { return ErrUnsafeFile } defer unix.Close(directory) if err := requirePrivateRegularAt(directory, target); err != nil { return ErrUnsafeFile } if err := unix.Unlinkat(directory, target, 0); err != nil { return ErrUnsafeFile } 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 requirePrivateRegularAt(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 || stat.Mode&0o7777 != 0o600 { 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 }