Files
ThothII/tools/tht/internal/authconfig/projection_lock_linux.go

178 lines
4.6 KiB
Go

//go:build linux
package authconfig
import (
"context"
"errors"
"os"
"sync"
"time"
"golang.org/x/sys/unix"
)
type transactionLock struct {
directoryFD int
fileFD int
closed bool
}
type projectionLockHooks struct {
beforeCreate func()
afterOpen func()
}
var projectionLockHookState struct {
sync.Mutex
hooks projectionLockHooks
}
func setProjectionLockHooksForTest(hooks projectionLockHooks) func() {
projectionLockHookState.Lock()
previous := projectionLockHookState.hooks
projectionLockHookState.hooks = hooks
projectionLockHookState.Unlock()
return func() {
projectionLockHookState.Lock()
projectionLockHookState.hooks = previous
projectionLockHookState.Unlock()
}
}
func projectionLockHooksForTest() projectionLockHooks {
projectionLockHookState.Lock()
defer projectionLockHookState.Unlock()
return projectionLockHookState.hooks
}
func acquireTransactionLock(ctx context.Context, directory string) (*transactionLock, error) {
if ctx == nil || ctx.Err() != nil || requirePrivateDirectory(directory) != nil {
return nil, errProjectionIntegrity
}
directoryFD, err := unix.Open(directory, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0)
if err != nil {
return nil, errProjectionIntegrity
}
lock := &transactionLock{directoryFD: directoryFD, fileFD: -1}
fail := func() (*transactionLock, error) {
_ = lock.Close()
return nil, errProjectionIntegrity
}
if !validTransactionDirectory(directoryFD) {
return fail()
}
fileFD, err := openTransactionLockFile(directoryFD)
if err != nil {
return fail()
}
lock.fileFD = fileFD
if hook := projectionLockHooksForTest().afterOpen; hook != nil {
hook()
}
var opened unix.Stat_t
if unix.Fstat(fileFD, &opened) != nil || !validTransactionLockFile(opened) || !sameTransactionLockName(directoryFD, opened) {
return fail()
}
if err := flockTransactionFile(ctx, fileFD); err != nil {
_ = lock.Close()
return nil, err
}
if !sameTransactionLockName(directoryFD, opened) {
return fail()
}
return lock, nil
}
func openTransactionLockFile(directoryFD int) (int, error) {
for attempt := 0; attempt < 2; attempt++ {
fd, err := unix.Openat(directoryFD, transactionLockFileName, unix.O_RDWR|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0)
if err == nil {
return fd, nil
}
if !errors.Is(err, unix.ENOENT) {
return -1, errProjectionIntegrity
}
if hook := projectionLockHooksForTest().beforeCreate; hook != nil {
hook()
}
fd, err = unix.Openat(directoryFD, transactionLockFileName, unix.O_RDWR|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0o600)
if err == nil {
if unix.Fchmod(fd, 0o600) != nil || unix.Fsync(fd) != nil || unix.Fsync(directoryFD) != nil {
unix.Close(fd)
return -1, errProjectionIntegrity
}
return fd, nil
}
if !errors.Is(err, unix.EEXIST) {
return -1, errProjectionIntegrity
}
}
return -1, errProjectionIntegrity
}
func validTransactionDirectory(fd int) bool {
var stat unix.Stat_t
return unix.Fstat(fd, &stat) == nil && stat.Mode&unix.S_IFMT == unix.S_IFDIR && stat.Uid == uint32(os.Geteuid()) && stat.Mode&0o077 == 0
}
func validTransactionLockFile(stat unix.Stat_t) bool {
return stat.Mode&unix.S_IFMT == unix.S_IFREG && stat.Nlink == 1 && stat.Uid == uint32(os.Geteuid()) && stat.Mode&0o7777 == 0o600
}
func sameTransactionLockName(directoryFD int, opened unix.Stat_t) bool {
var named unix.Stat_t
return unix.Fstatat(directoryFD, transactionLockFileName, &named, unix.AT_SYMLINK_NOFOLLOW) == nil && validTransactionLockFile(named) && named.Dev == opened.Dev && named.Ino == opened.Ino
}
func flockTransactionFile(ctx context.Context, fd int) error {
for {
if err := ctx.Err(); err != nil {
return err
}
err := unix.Flock(fd, unix.LOCK_EX|unix.LOCK_NB)
if err == nil {
return nil
}
if !errors.Is(err, unix.EWOULDBLOCK) {
return errProjectionIntegrity
}
if hook := coordinatorHooks().onOuterLockContention; hook != nil {
hook()
}
timer := time.NewTimer(10 * time.Millisecond)
select {
case <-ctx.Done():
if !timer.Stop() {
<-timer.C
}
return ctx.Err()
case <-timer.C:
}
}
}
func (lock *transactionLock) Close() error {
if lock == nil || lock.closed {
return nil
}
lock.closed = true
var err error
if lock.fileFD >= 0 {
if unlockErr := unix.Flock(lock.fileFD, unix.LOCK_UN); unlockErr != nil {
err = errors.Join(err, errProjectionCleanup)
}
if closeErr := unix.Close(lock.fileFD); closeErr != nil {
err = errors.Join(err, errProjectionCleanup)
}
lock.fileFD = -1
}
if lock.directoryFD >= 0 {
if closeErr := unix.Close(lock.directoryFD); closeErr != nil {
err = errors.Join(err, errProjectionCleanup)
}
lock.directoryFD = -1
}
return err
}