diff --git a/tools/tht/internal/authconfig/store_windows_test.go b/tools/tht/internal/authconfig/store_windows_test.go index c2bc6513..f5cc65fa 100644 --- a/tools/tht/internal/authconfig/store_windows_test.go +++ b/tools/tht/internal/authconfig/store_windows_test.go @@ -4,8 +4,10 @@ package authconfig import ( "errors" + "fmt" "path/filepath" "runtime" + "sync" "testing" "github.com/aritmolab/thothii/tools/tht/internal/safeio" @@ -49,9 +51,52 @@ func TestMutateUsersCreatesAndRejectsPermissiveWindowsLockDACL(t *testing.T) { } } -func writePrivateWindowsAuthFiles(t *testing.T) string { +func TestMutateUsersConcurrentWindowsLockCreationNeverObservesDefaultDACL(t *testing.T) { + users := adminUserYAML("admin", "Admin", true, "admin") + for index := 0; index < 32; index++ { + users += userYAML(index) + } + directory := writePrivateWindowsAuthFiles(t, registryYAML(users)) + + var group sync.WaitGroup + var ready sync.WaitGroup + start := make(chan struct{}) + errors := make(chan error, 32) + for index := 0; index < 32; index++ { + index := index + group.Add(1) + ready.Add(1) + go func() { + defer group.Done() + ready.Done() + <-start + errors <- MutateUsers(directory, func(registry *Registry) error { + registry.Users[index+1].DisplayName = fmt.Sprintf("Windows Updated %d", index) + return nil + }) + }() + } + ready.Wait() + close(start) + group.Wait() + close(errors) + for err := range errors { + if err != nil { + t.Fatalf("MutateUsers() concurrent lock creation error = %v", err) + } + } + if err := safeio.ValidatePrivateRegular(filepath.Join(directory, lockFileName)); err != nil { + t.Fatalf("concurrently created lock DACL error = %v", err) + } +} + +func writePrivateWindowsAuthFiles(t *testing.T, users ...string) string { t.Helper() - directory := writeAuthFiles(t, defaultAuthYAML, registryYAML(adminUserYAML("admin", "Admin", true, "admin"))) + registry := registryYAML(adminUserYAML("admin", "Admin", true, "admin")) + if len(users) > 0 { + registry = users[0] + } + directory := writeAuthFiles(t, defaultAuthYAML, registry) if err := safeio.ProtectPrivateDirectory(directory); err != nil { t.Fatal(err) } diff --git a/tools/tht/internal/safeio/files.go b/tools/tht/internal/safeio/files.go index 0ea916f3..cfc54b91 100644 --- a/tools/tht/internal/safeio/files.go +++ b/tools/tht/internal/safeio/files.go @@ -72,20 +72,21 @@ func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error } else if !errors.Is(err, os.ErrNotExist) { return ErrUnsafeFile } - file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, mode) + file, err := createCanonicalNewPrivateFile(path, mode) if err != nil { return ErrUnsafeFile } - defer file.Close() - if err := ProtectPrivateRegular(path); err != nil { - _ = os.Remove(path) - return ErrUnsafeFile - } if _, err := file.Write(contents); err != nil { + _ = file.Close() _ = os.Remove(path) return ErrUnsafeFile } if err := file.Sync(); err != nil { + _ = file.Close() + _ = os.Remove(path) + return ErrUnsafeFile + } + if err := file.Close(); err != nil { _ = os.Remove(path) return ErrUnsafeFile } diff --git a/tools/tht/internal/safeio/private_unix.go b/tools/tht/internal/safeio/private_unix.go index 8630a63c..4aed559d 100644 --- a/tools/tht/internal/safeio/private_unix.go +++ b/tools/tht/internal/safeio/private_unix.go @@ -41,6 +41,19 @@ func ProtectPrivateRegular(path string) error { return ValidatePrivateRegular(path) } +func createCanonicalNewPrivateFile(path string, mode os.FileMode) (*os.File, error) { + file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, mode) + if err != nil { + return nil, err + } + if err := ProtectPrivateRegular(path); err != nil { + _ = file.Close() + _ = os.Remove(path) + return nil, ErrUnsafeFile + } + return file, nil +} + // ValidatePrivateRegular requires a canonical, single-link private regular file. func ValidatePrivateRegular(path string) error { if err := ValidateCanonicalPath(path); err != nil { diff --git a/tools/tht/internal/safeio/private_windows.go b/tools/tht/internal/safeio/private_windows.go index 49ca9283..69ed4e89 100644 --- a/tools/tht/internal/safeio/private_windows.go +++ b/tools/tht/internal/safeio/private_windows.go @@ -3,6 +3,7 @@ package safeio import ( + "os" "path/filepath" "runtime" "strings" @@ -65,6 +66,45 @@ func ProtectPrivateRegular(path string) error { return validateOwnerOnlyDACL(handle) } +// createCanonicalNewPrivateFile installs the owner-only protected DACL in the CreateFile call, so +// another mutation can never observe a newly-created lock with an inherited/default DACL. +func createCanonicalNewPrivateFile(path string, mode os.FileMode) (*os.File, error) { + security, err := newOwnerOnlySecurityDescriptor() + if err != nil { + return nil, ErrUnsafeFile + } + defer security.Close() + attributes := &windows.SecurityAttributes{ + Length: uint32(unsafe.Sizeof(windows.SecurityAttributes{})), + SecurityDescriptor: security.descriptor, + } + handle, err := windows.CreateFile( + windows.StringToUTF16Ptr(path), + windows.GENERIC_WRITE, + windowsRetainedHandleShareMode, + attributes, + windows.CREATE_NEW, + windows.FILE_ATTRIBUTE_NORMAL, + 0, + ) + runtime.KeepAlive(security) + if err != nil { + return nil, err + } + if err := validateOwnerOnlyDACL(handle); err != nil { + _ = windows.CloseHandle(handle) + _ = os.Remove(path) + return nil, ErrUnsafeFile + } + file := os.NewFile(uintptr(handle), "tht-safeio-private") + if file == nil { + _ = windows.CloseHandle(handle) + _ = os.Remove(path) + return nil, ErrUnsafeFile + } + return file, nil +} + // ValidatePrivateRegular requires a canonical, single-link file protected for its current owner only. func ValidatePrivateRegular(path string) error { parents, target, err := openCanonicalWindowsParent(path) @@ -125,29 +165,84 @@ func openCanonicalWindowsParent(path string) (*windowsParentHandles, string, err return parents, components[len(components)-1], nil } -func setOwnerOnlyDACL(handle windows.Handle) error { - sid, err := currentOwnerSID() - if err != nil { - return err +type ownerOnlyDACL struct { + sid *windows.SID + acl *windows.ACL + pinner runtime.Pinner +} + +func newOwnerOnlyDACL() (*ownerOnlyDACL, error) { + tokenUser, err := windows.GetCurrentProcessToken().GetTokenUser() + if err != nil || tokenUser == nil || tokenUser.User.Sid == nil { + return nil, ErrUnsafeFile } - var pinner runtime.Pinner - pinner.Pin(sid) - defer pinner.Unpin() + sid, err := tokenUser.User.Sid.Copy() + if err != nil { + return nil, ErrUnsafeFile + } + owner := &ownerOnlyDACL{sid: sid} + owner.pinner.Pin(owner.sid) acl, err := windows.ACLFromEntries([]windows.EXPLICIT_ACCESS{{ AccessPermissions: windows.GENERIC_ALL, AccessMode: windows.GRANT_ACCESS, Trustee: windows.TRUSTEE{ TrusteeForm: windows.TRUSTEE_IS_SID, TrusteeType: windows.TRUSTEE_IS_USER, - TrusteeValue: windows.TrusteeValueFromSID(sid), + TrusteeValue: windows.TrusteeValueFromSID(owner.sid), }, }}, nil) + if err != nil { + owner.pinner.Unpin() + return nil, err + } + owner.acl = acl + return owner, nil +} + +func (owner *ownerOnlyDACL) Close() { + owner.pinner.Unpin() +} + +type ownerOnlySecurityDescriptor struct { + *ownerOnlyDACL + descriptor *windows.SECURITY_DESCRIPTOR +} + +func newOwnerOnlySecurityDescriptor() (*ownerOnlySecurityDescriptor, error) { + owner, err := newOwnerOnlyDACL() + if err != nil { + return nil, err + } + descriptor, err := windows.NewSecurityDescriptor() + if err == nil { + err = descriptor.SetOwner(owner.sid, false) + } + if err == nil { + err = descriptor.SetDACL(owner.acl, true, false) + } + if err == nil { + err = descriptor.SetControl(windows.SE_DACL_PROTECTED, windows.SE_DACL_PROTECTED) + } + if err != nil || !descriptor.IsValid() { + owner.Close() + return nil, ErrUnsafeFile + } + return &ownerOnlySecurityDescriptor{ownerOnlyDACL: owner, descriptor: descriptor}, nil +} + +func (descriptor *ownerOnlySecurityDescriptor) Close() { + descriptor.ownerOnlyDACL.Close() +} + +func setOwnerOnlyDACL(handle windows.Handle) error { + owner, err := newOwnerOnlyDACL() if err != nil { return err } + defer owner.Close() return windows.SetSecurityInfo(handle, windows.SE_FILE_OBJECT, windows.OWNER_SECURITY_INFORMATION|windows.DACL_SECURITY_INFORMATION|windows.PROTECTED_DACL_SECURITY_INFORMATION, - sid, nil, acl, nil) + owner.sid, nil, owner.acl, nil) } func validateOwnerOnlyDACL(handle windows.Handle) error { @@ -155,32 +250,40 @@ func validateOwnerOnlyDACL(handle windows.Handle) error { if err != nil { return err } + return withWindowsSecurityDescriptor(handle, func(descriptor *windows.SECURITY_DESCRIPTOR) error { + owner, _, err := descriptor.Owner() + if err != nil || owner == nil || !windows.EqualSid(owner, ownerSID) { + return ErrUnsafeFile + } + control, _, err := descriptor.Control() + if err != nil || control&windows.SE_DACL_PROTECTED == 0 { + return ErrUnsafeFile + } + dacl, defaulted, err := descriptor.DACL() + if err != nil || defaulted || dacl == nil || dacl.AceCount != 1 { + return ErrUnsafeFile + } + var ace *windows.ACCESS_ALLOWED_ACE + if err := windows.GetAce(dacl, 0, &ace); err != nil || ace == nil || ace.Header.AceType != windows.ACCESS_ALLOWED_ACE_TYPE || ace.Header.AceFlags != 0 || ace.Mask != windows.GENERIC_ALL { + return ErrUnsafeFile + } + aceSID := (*windows.SID)(unsafe.Pointer(&ace.SidStart)) + if !windows.EqualSid(aceSID, ownerSID) { + return ErrUnsafeFile + } + return nil + }) +} + +// withWindowsSecurityDescriptor confines inspection to x/sys's Go-owned descriptor copy. Its +// GetSecurityInfo wrapper releases the native LocalAlloc result with LocalFree before returning. +func withWindowsSecurityDescriptor(handle windows.Handle, inspect func(*windows.SECURITY_DESCRIPTOR) error) error { descriptor, err := windows.GetSecurityInfo(handle, windows.SE_FILE_OBJECT, windows.OWNER_SECURITY_INFORMATION|windows.DACL_SECURITY_INFORMATION) if err != nil || descriptor == nil { return ErrUnsafeFile } - owner, _, err := descriptor.Owner() - if err != nil || owner == nil || !windows.EqualSid(owner, ownerSID) { - return ErrUnsafeFile - } - control, _, err := descriptor.Control() - if err != nil || control&windows.SE_DACL_PROTECTED == 0 { - return ErrUnsafeFile - } - dacl, defaulted, err := descriptor.DACL() - if err != nil || defaulted || dacl == nil || dacl.AceCount != 1 { - return ErrUnsafeFile - } - var ace *windows.ACCESS_ALLOWED_ACE - if err := windows.GetAce(dacl, 0, &ace); err != nil || ace == nil || ace.Header.AceType != windows.ACCESS_ALLOWED_ACE_TYPE || ace.Header.AceFlags != 0 || ace.Mask != windows.GENERIC_ALL { - return ErrUnsafeFile - } - aceSID := (*windows.SID)(unsafe.Pointer(&ace.SidStart)) - if !windows.EqualSid(aceSID, ownerSID) { - return ErrUnsafeFile - } - return nil + return inspect(descriptor) } func currentOwnerSID() (*windows.SID, error) { @@ -188,5 +291,9 @@ func currentOwnerSID() (*windows.SID, error) { if err != nil || user == nil || user.User.Sid == nil { return nil, ErrUnsafeFile } - return user.User.Sid, nil + sid, err := user.User.Sid.Copy() + if err != nil { + return nil, ErrUnsafeFile + } + return sid, nil } diff --git a/tools/tht/internal/safeio/private_windows_test.go b/tools/tht/internal/safeio/private_windows_test.go index 2be8812c..26985f3f 100644 --- a/tools/tht/internal/safeio/private_windows_test.go +++ b/tools/tht/internal/safeio/private_windows_test.go @@ -53,6 +53,63 @@ func TestPrivateWindowsDACLRejectsPermissiveDirectoryAndRegularFile(t *testing.T } } +func TestCreateCanonicalNewPrivateFileInstallsOwnerOnlyDACLAtCreation(t *testing.T) { + directory := filepath.Join(t.TempDir(), "auth") + if err := os.Mkdir(directory, 0o700); err != nil { + t.Fatal(err) + } + if err := ProtectPrivateDirectory(directory); err != nil { + t.Fatal(err) + } + + path := filepath.Join(directory, ".auth.lock") + file, err := createCanonicalNewPrivateFile(path, 0o600) + if err != nil { + t.Fatal(err) + } + defer file.Close() + if err := ValidatePrivateRegular(path); err != nil { + t.Fatalf("new lock DACL error = %v", err) + } +} + +func TestWithWindowsSecurityDescriptorKeepsOwnedDescriptorValidDuringInspection(t *testing.T) { + directory := filepath.Join(t.TempDir(), "auth") + if err := os.Mkdir(directory, 0o700); err != nil { + t.Fatal(err) + } + if err := ProtectPrivateDirectory(directory); err != nil { + t.Fatal(err) + } + path := filepath.Join(directory, "users.yaml") + if err := os.WriteFile(path, []byte("private"), 0o600); err != nil { + t.Fatal(err) + } + if err := ProtectPrivateRegular(path); err != nil { + t.Fatal(err) + } + + handle, err := openWindowsComponent(path, false) + if err != nil { + t.Fatal(err) + } + defer windows.CloseHandle(handle) + if err := withWindowsSecurityDescriptor(handle, func(descriptor *windows.SECURITY_DESCRIPTOR) error { + runtime.GC() + owner, _, err := descriptor.Owner() + if err != nil || owner == nil { + t.Fatalf("descriptor owner error = %v", err) + } + dacl, _, err := descriptor.DACL() + if err != nil || dacl == nil || dacl.AceCount != 1 { + t.Fatalf("descriptor DACL error = %v", err) + } + return nil + }); err != nil { + t.Fatalf("withWindowsSecurityDescriptor() error = %v", err) + } +} + func TestReplaceCanonicalRegularCreatesPrivateTemporaryAndReplacement(t *testing.T) { directory := filepath.Join(t.TempDir(), "auth") if err := os.Mkdir(directory, 0o700); err != nil { diff --git a/tools/tht/internal/safeio/replace_windows.go b/tools/tht/internal/safeio/replace_windows.go index 5b1302fd..98da45c5 100644 --- a/tools/tht/internal/safeio/replace_windows.go +++ b/tools/tht/internal/safeio/replace_windows.go @@ -58,24 +58,17 @@ func writePrivateTemporary(directory string, contents []byte) (string, error) { return "", err } path := filepath.Join(directory, name) - file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o600) + file, err := createCanonicalNewPrivateFile(path, 0o600) if errors.Is(err, os.ErrExist) { continue } if err != nil { return "", err } - if err := ProtectPrivateRegular(path); err != nil { - _ = file.Close() - _ = os.Remove(path) - return "", err - } - if err := file.Chmod(0o600); err == nil { - var written int - written, err = file.Write(contents) - if err == nil && written != len(contents) { - err = io.ErrShortWrite - } + var written int + written, err = file.Write(contents) + if err == nil && written != len(contents) { + err = io.ErrShortWrite } if err == nil { err = file.Sync()