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

463 lines
16 KiB
Go

//go:build !windows
package safeio
import (
"errors"
"io"
"os"
"sort"
"time"
"golang.org/x/sys/unix"
)
type unixPrivateDirectory struct {
descriptor int
}
func openPrivateDirectory(path string, ensure bool) (PrivateDirectoryHandle, bool, error) {
parents, err := openCanonicalUnixParent(path)
if err != nil {
return nil, false, ErrUnsafeFile
}
defer parents.Close()
return openPrivateUnixDirectoryAt(parents.parent, parents.target, ensure)
}
func openPrivateUnixDirectoryAt(parent int, name string, ensure bool) (PrivateDirectoryHandle, bool, error) {
if parent < 0 || !validPrivateLeafName(name) {
return nil, false, ErrUnsafeFile
}
for attempt := 0; attempt < 2; attempt++ {
descriptor, err := unix.Openat(parent, name, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0)
if err == nil {
value := &unixPrivateDirectory{descriptor: descriptor}
if value.Validate() != nil {
_ = value.Close()
return nil, false, ErrUnsafeFile
}
return value, true, nil
}
if !errors.Is(err, unix.ENOENT) {
return nil, false, ErrUnsafeFile
}
if !ensure {
if unix.Faccessat(parent, ".", unix.W_OK|unix.X_OK, unix.AT_EACCESS) != nil {
return nil, false, ErrUnsafeFile
}
return nil, false, nil
}
if err := unix.Mkdirat(parent, name, 0o700); err != nil && !errors.Is(err, unix.EEXIST) {
return nil, false, ErrUnsafeFile
}
}
return nil, false, ErrUnsafeFile
}
func (directory *unixPrivateDirectory) Close() error {
if directory == nil || directory.descriptor < 0 {
return nil
}
err := unix.Close(directory.descriptor)
directory.descriptor = -1
if err != nil {
return ErrUnsafeFile
}
return nil
}
func (directory *unixPrivateDirectory) Validate() error {
if directory == nil || directory.descriptor < 0 {
return ErrUnsafeFile
}
var stat unix.Stat_t
if err := unix.Fstat(directory.descriptor, &stat); err != nil || !privateUnixDirectoryStat(&stat) {
return ErrUnsafeFile
}
return nil
}
func (directory *unixPrivateDirectory) OpenChild(name string, ensure bool) (PrivateDirectoryHandle, bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) {
return nil, false, ErrUnsafeFile
}
child, found, err := openPrivateUnixDirectoryAt(directory.descriptor, name, ensure)
if err != nil || directory.Validate() != nil {
if child != nil {
_ = child.Close()
}
return nil, false, ErrUnsafeFile
}
return child, found, nil
}
func privateUnixRootRegular(stat *unix.Stat_t, links uint64) bool {
return stat != nil && stat.Mode&unix.S_IFMT == unix.S_IFREG && uint64(stat.Nlink) == links &&
stat.Uid == uint32(os.Geteuid()) && stat.Mode&0o7777 == 0o600
}
func sameUnixRootRegular(left, right unix.Stat_t) bool {
return left.Dev == right.Dev && left.Ino == right.Ino && left.Size == right.Size &&
left.Mtim == right.Mtim && left.Ctim == right.Ctim && left.Mode == right.Mode && left.Nlink == right.Nlink
}
func requirePrivateUnixRootRegularAt(directory int, name string, links uint64) (unix.Stat_t, error) {
var stat unix.Stat_t
if err := unix.Fstatat(directory, name, &stat, unix.AT_SYMLINK_NOFOLLOW); err != nil || !privateUnixRootRegular(&stat, links) {
return unix.Stat_t{}, ErrUnsafeFile
}
return stat, nil
}
func privateUnixRootRegularAtAllowedLinks(directory int, name string, allowed ...uint64) error {
var stat unix.Stat_t
if err := unix.Fstatat(directory, name, &stat, unix.AT_SYMLINK_NOFOLLOW); err != nil {
return ErrUnsafeFile
}
for _, links := range allowed {
if privateUnixRootRegular(&stat, links) {
return nil
}
}
return ErrUnsafeFile
}
func requireSamePrivateUnixRootPairAt(directory int, source, claim string) error {
left, err := requirePrivateUnixRootRegularAt(directory, source, 2)
if err != nil {
return ErrUnsafeFile
}
right, err := requirePrivateUnixRootRegularAt(directory, claim, 2)
if err != nil || !sameUnixPrivateFile(left, right) {
return ErrUnsafeFile
}
return nil
}
func isPrivateUnixRootClaimAbsentOrOrphanAt(directory int, source, claim string) bool {
var sourceStat, claimStat unix.Stat_t
sourceErr := unix.Fstatat(directory, source, &sourceStat, unix.AT_SYMLINK_NOFOLLOW)
claimErr := unix.Fstatat(directory, claim, &claimStat, unix.AT_SYMLINK_NOFOLLOW)
if errors.Is(sourceErr, unix.ENOENT) && errors.Is(claimErr, unix.ENOENT) {
return true
}
return errors.Is(sourceErr, unix.ENOENT) && claimErr == nil && privateUnixRootRegular(&claimStat, 1)
}
func writeAllPrivateRoot(file *os.File, contents []byte) error {
for written := 0; written < len(contents); {
count, err := file.Write(contents[written:])
written += count
if err != nil {
return err
}
if count == 0 {
return io.ErrShortWrite
}
}
return file.Sync()
}
func (directory *unixPrivateDirectory) CreateRegular(name string, contents []byte) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) || len(contents) == 0 {
return false, ErrUnsafeFile
}
descriptor, err := unix.Openat(directory.descriptor, name,
unix.O_WRONLY|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0o600)
if errors.Is(err, unix.EEXIST) {
if _, existingErr := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); existingErr != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return false, nil
}
if err != nil {
return false, ErrUnsafeFile
}
file := os.NewFile(uintptr(descriptor), "tht-safeio-private-root-create")
if file == nil {
_ = unix.Close(descriptor)
_ = unix.Unlinkat(directory.descriptor, name, 0)
return false, ErrUnsafeFile
}
failed := true
defer func() {
if failed {
_ = unix.Unlinkat(directory.descriptor, name, 0)
}
}()
if unix.Fchmod(descriptor, 0o600) != nil {
_ = file.Close()
return false, ErrUnsafeFile
}
var stat unix.Stat_t
if unix.Fstat(descriptor, &stat) != nil || !privateUnixRootRegular(&stat, 1) || writeAllPrivateRoot(file, contents) != nil || file.Close() != nil {
return false, ErrUnsafeFile
}
if directory.Validate() != nil || unix.Fsync(directory.descriptor) != nil {
return false, ErrUnsafeFile
}
failed = false
return true, nil
}
func (directory *unixPrivateDirectory) CreateRegularFile(name string) (*os.File, bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) {
return nil, false, ErrUnsafeFile
}
descriptor, err := unix.Openat(directory.descriptor, name,
unix.O_RDWR|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0o600)
if errors.Is(err, unix.EEXIST) {
if _, existingErr := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); existingErr != nil || directory.Validate() != nil {
return nil, false, ErrUnsafeFile
}
return nil, false, nil
}
if err != nil {
return nil, false, ErrUnsafeFile
}
file := os.NewFile(uintptr(descriptor), "tht-safeio-private-root-stream")
if file == nil {
_ = unix.Close(descriptor)
_ = unix.Unlinkat(directory.descriptor, name, 0)
return nil, false, ErrUnsafeFile
}
failed := true
defer func() {
if failed {
_ = file.Close()
_ = unix.Unlinkat(directory.descriptor, name, 0)
}
}()
if unix.Fchmod(descriptor, 0o600) != nil {
return nil, false, ErrUnsafeFile
}
var stat unix.Stat_t
if unix.Fstat(descriptor, &stat) != nil || !privateUnixRootRegular(&stat, 1) {
return nil, false, ErrUnsafeFile
}
if directory.Validate() != nil || unix.Fsync(directory.descriptor) != nil {
return nil, false, ErrUnsafeFile
}
failed = false
return file, true, nil
}
func (directory *unixPrivateDirectory) ReadRegular(name string, maximum int64) ([]byte, bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) || maximum < 0 || maximum == int64(^uint64(0)>>1) {
return nil, false, ErrUnsafeFile
}
before, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1)
if err != nil {
var stat unix.Stat_t
if statErr := unix.Fstatat(directory.descriptor, name, &stat, unix.AT_SYMLINK_NOFOLLOW); errors.Is(statErr, unix.ENOENT) {
return nil, false, nil
}
return nil, false, ErrUnsafeFile
}
descriptor, err := unix.Openat(directory.descriptor, name, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0)
if err != nil {
return nil, false, ErrUnsafeFile
}
file := os.NewFile(uintptr(descriptor), "tht-safeio-private-root-read")
if file == nil {
_ = unix.Close(descriptor)
return nil, false, ErrUnsafeFile
}
contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1))
var opened unix.Stat_t
statErr := unix.Fstat(descriptor, &opened)
closeErr := file.Close()
var current unix.Stat_t
currentErr := unix.Fstatat(directory.descriptor, name, &current, unix.AT_SYMLINK_NOFOLLOW)
if readErr != nil || statErr != nil || closeErr != nil || currentErr != nil || int64(len(contents)) > maximum ||
!privateUnixRootRegular(&opened, 1) || !privateUnixRootRegular(&current, 1) ||
!sameUnixRootRegular(before, opened) || !sameUnixRootRegular(opened, current) || directory.Validate() != nil {
return nil, false, ErrUnsafeFile
}
return contents, true, nil
}
func (directory *unixPrivateDirectory) ReplaceRegular(name string, contents []byte) error {
if directory.Validate() != nil || !validPrivateLeafName(name) || len(contents) == 0 {
return ErrUnsafeFile
}
if _, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); err != nil {
return ErrUnsafeFile
}
temporary, err := writePrivateTemporaryAt(directory.descriptor, contents)
if err != nil {
return ErrUnsafeFile
}
defer func() { _ = unix.Unlinkat(directory.descriptor, temporary, 0) }()
if _, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); err != nil || directory.Validate() != nil {
return ErrUnsafeFile
}
if err := unix.Renameat(directory.descriptor, temporary, directory.descriptor, name); err != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil {
return ErrUnsafeFile
}
return nil
}
func (directory *unixPrivateDirectory) RemoveRegular(name string) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) {
return false, ErrUnsafeFile
}
if _, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); err != nil {
var stat unix.Stat_t
if statErr := unix.Fstatat(directory.descriptor, name, &stat, unix.AT_SYMLINK_NOFOLLOW); errors.Is(statErr, unix.ENOENT) {
return false, nil
}
return false, ErrUnsafeFile
}
if err := unix.Unlinkat(directory.descriptor, name, 0); err != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return true, nil
}
func (directory *unixPrivateDirectory) ListPage(
maximumEntries int,
afterName string,
validName func(string) bool,
validLinks func(string, uint64) bool,
) (PrivateDirectoryPage, error) {
if directory.Validate() != nil || maximumEntries < 1 || maximumEntries > 4096 || validName == nil || validLinks == nil || (afterName != "" && !validName(afterName)) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
var before unix.Stat_t
if unix.Fstat(directory.descriptor, &before) != nil || !privateUnixDirectoryStat(&before) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
duplicate, err := unix.Dup(directory.descriptor)
if err != nil {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
file := os.NewFile(uintptr(duplicate), "tht-safeio-private-root-list")
if file == nil {
_ = unix.Close(duplicate)
return PrivateDirectoryPage{}, ErrUnsafeFile
}
defer file.Close()
seen := make(map[string]struct{}, maximumEntries+1)
selected := make([]PrivateDirectoryEntry, 0, maximumEntries+1)
scanned := 0
for {
entries, readErr := file.ReadDir(1)
if readErr != nil && !errors.Is(readErr, io.EOF) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if len(entries) == 0 {
break
}
if len(entries) != 1 {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
scanned++
if scanned > maximumPrivateDirectoryPageScanEntries {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
name := entries[0].Name()
if !validPrivateLeafName(name) || !validName(name) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if _, duplicate := seen[name]; duplicate {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
seen[name] = struct{}{}
var stat unix.Stat_t
if unix.Fstatat(directory.descriptor, name, &stat, unix.AT_SYMLINK_NOFOLLOW) != nil ||
!privateUnixRootRegular(&stat, uint64(stat.Nlink)) || !validLinks(name, uint64(stat.Nlink)) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if name > afterName {
selected = appendBoundedPrivateDirectoryEntry(selected, PrivateDirectoryEntry{
Name: name, ModifiedUnixMs: time.Unix(stat.Mtim.Sec, stat.Mtim.Nsec).UnixMilli(),
}, maximumEntries+1)
}
if errors.Is(readErr, io.EOF) {
break
}
}
var after unix.Stat_t
if unix.Fstat(directory.descriptor, &after) != nil || !privateUnixDirectoryStat(&after) ||
before.Dev != after.Dev || before.Ino != after.Ino || before.Mode != after.Mode || before.Mtim != after.Mtim || before.Ctim != after.Ctim || directory.Validate() != nil {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
sort.Slice(selected, func(left, right int) bool { return selected[left].Name < selected[right].Name })
more := len(selected) > maximumEntries
if more {
selected = selected[:maximumEntries]
}
return PrivateDirectoryPage{Entries: selected, More: more}, nil
}
func (directory *unixPrivateDirectory) ClaimRegular(source, claim string) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) {
return false, ErrUnsafeFile
}
if err := privateUnixRootRegularAtAllowedLinks(directory.descriptor, source, 1); err != nil {
if requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim) == nil || isPrivateUnixRootClaimAbsentOrOrphanAt(directory.descriptor, source, claim) {
return false, nil
}
return false, ErrUnsafeFile
}
if err := unix.Linkat(directory.descriptor, source, directory.descriptor, claim, 0); err != nil {
if errors.Is(err, unix.EEXIST) && privateUnixRootRegularAtAllowedLinks(directory.descriptor, claim, 1, 2) == nil {
return false, nil
}
return false, ErrUnsafeFile
}
if err := requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim); err != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return true, nil
}
func (directory *unixPrivateDirectory) ReadClaim(source, claim string, maximum int64) ([]byte, bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) || maximum < 0 || maximum == int64(^uint64(0)>>1) {
return nil, false, ErrUnsafeFile
}
if err := requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim); err != nil {
if isPrivateUnixRootClaimAbsentOrOrphanAt(directory.descriptor, source, claim) {
return nil, false, nil
}
return nil, false, ErrUnsafeFile
}
descriptor, err := unix.Openat(directory.descriptor, source, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0)
if err != nil {
return nil, false, ErrUnsafeFile
}
file := os.NewFile(uintptr(descriptor), "tht-safeio-private-root-claim")
if file == nil {
_ = unix.Close(descriptor)
return nil, false, ErrUnsafeFile
}
contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1))
var after unix.Stat_t
statErr := unix.Fstat(descriptor, &after)
closeErr := file.Close()
if readErr != nil || statErr != nil || closeErr != nil || int64(len(contents)) > maximum || !privateUnixRootRegular(&after, 2) ||
requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim) != nil || directory.Validate() != nil {
return nil, false, ErrUnsafeFile
}
return contents, true, nil
}
func (directory *unixPrivateDirectory) RemoveClaim(source, claim string) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) {
return false, ErrUnsafeFile
}
if err := requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim); err != nil {
if isPrivateUnixRootClaimAbsentOrOrphanAt(directory.descriptor, source, claim) {
return false, nil
}
return false, ErrUnsafeFile
}
if err := unix.Unlinkat(directory.descriptor, source, 0); err != nil || privateUnixRootRegularAtAllowedLinks(directory.descriptor, claim, 1) != nil ||
unix.Unlinkat(directory.descriptor, claim, 0) != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return true, nil
}