diff --git a/tools/tht/internal/authconfig/store.go b/tools/tht/internal/authconfig/store.go index 7264c2b8..c770ecef 100644 --- a/tools/tht/internal/authconfig/store.go +++ b/tools/tht/internal/authconfig/store.go @@ -6,7 +6,6 @@ import ( "io" "os" "path/filepath" - "runtime" "github.com/aritmolab/thothii/tools/tht/internal/safeio" "github.com/gofrs/flock" @@ -112,26 +111,15 @@ func decodeStrictYAML(contents []byte, destination any) error { } func requirePrivateDirectory(directory string) error { - if err := safeio.ValidateCanonicalPath(directory); err != nil { - return err - } - info, err := os.Lstat(directory) - if err != nil || !info.IsDir() || info.Mode()&os.ModeSymlink != 0 { - return safeio.ErrUnsafeFile - } - resolved, err := filepath.EvalSymlinks(directory) - if err != nil || resolved != directory { - return safeio.ErrUnsafeFile - } - if runtime.GOOS != "windows" && info.Mode().Perm() != 0o700 { - return safeio.ErrUnsafeFile - } - return nil + return safeio.ValidatePrivateDirectory(directory) } func readPrivateFile(path string) ([]byte, error) { + if err := safeio.ValidatePrivateRegular(path); err != nil { + return nil, err + } before, err := os.Lstat(path) - if err != nil || !before.Mode().IsRegular() || before.Mode()&os.ModeSymlink != 0 || (runtime.GOOS != "windows" && before.Mode().Perm() != 0o600) { + if err != nil { return nil, safeio.ErrUnsafeFile } contents, err := safeio.ReadCanonicalRegular(path, maxYAMLBytes) @@ -139,7 +127,10 @@ func readPrivateFile(path string) ([]byte, error) { return nil, err } after, err := os.Lstat(path) - if err != nil || !after.Mode().IsRegular() || after.Mode()&os.ModeSymlink != 0 || (runtime.GOOS != "windows" && after.Mode().Perm() != 0o600) || !os.SameFile(before, after) { + if err != nil || !os.SameFile(before, after) { + return nil, safeio.ErrUnsafeFile + } + if err := safeio.ValidatePrivateRegular(path); err != nil { return nil, safeio.ErrUnsafeFile } return contents, nil @@ -149,14 +140,18 @@ func acquireLock(directory string) (*flock.Flock, error) { path := filepath.Join(directory, lockFileName) if _, err := os.Lstat(path); errors.Is(err, os.ErrNotExist) { if err := safeio.WriteCanonicalNewFile(path, nil, 0o600); err != nil { - info, statErr := os.Lstat(path) - if statErr != nil || !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || (runtime.GOOS != "windows" && info.Mode().Perm() != 0o600) { + // A competing mutation may have created the lock after our Lstat. Accept only + // that exact safe/private lock; every other creation failure remains unsafe. + if validateErr := safeio.ValidatePrivateRegular(path); validateErr != nil { return nil, safeio.ErrUnsafeFile } } } else if err != nil { return nil, safeio.ErrUnsafeFile } + if err := safeio.ValidatePrivateRegular(path); err != nil { + return nil, err + } lock := flock.New(path, flock.SetPermissions(0o600)) if err := lock.Lock(); err != nil { return nil, errInvalidAuthenticationConfig diff --git a/tools/tht/internal/authconfig/store_windows_test.go b/tools/tht/internal/authconfig/store_windows_test.go new file mode 100644 index 00000000..c2bc6513 --- /dev/null +++ b/tools/tht/internal/authconfig/store_windows_test.go @@ -0,0 +1,92 @@ +//go:build windows + +package authconfig + +import ( + "errors" + "path/filepath" + "runtime" + "testing" + + "github.com/aritmolab/thothii/tools/tht/internal/safeio" + "golang.org/x/sys/windows" +) + +func TestLoadRejectsPermissiveWindowsDirectoryAuthAndUsersDACLs(t *testing.T) { + for name, makePermissive := range map[string]func(t *testing.T, directory string){ + "directory": func(t *testing.T, directory string) { + setPermissiveAuthDACL(t, directory) + }, + "auth file": func(t *testing.T, directory string) { + setPermissiveAuthDACL(t, filepath.Join(directory, "auth.yaml")) + }, + "users file": func(t *testing.T, directory string) { + setPermissiveAuthDACL(t, filepath.Join(directory, "users.yaml")) + }, + } { + t.Run(name, func(t *testing.T) { + directory := writePrivateWindowsAuthFiles(t) + makePermissive(t, directory) + if _, _, err := Load(directory); !errors.Is(err, safeio.ErrUnsafeFile) { + t.Fatalf("Load() error = %v, want ErrUnsafeFile", err) + } + }) + } +} + +func TestMutateUsersCreatesAndRejectsPermissiveWindowsLockDACL(t *testing.T) { + directory := writePrivateWindowsAuthFiles(t) + if err := MutateUsers(directory, func(*Registry) error { return nil }); err != nil { + t.Fatal(err) + } + lockPath := filepath.Join(directory, lockFileName) + if err := safeio.ValidatePrivateRegular(lockPath); err != nil { + t.Fatalf("lock privacy error = %v", err) + } + setPermissiveAuthDACL(t, lockPath) + if err := MutateUsers(directory, func(*Registry) error { return nil }); !errors.Is(err, safeio.ErrUnsafeFile) { + t.Fatalf("MutateUsers() error = %v, want ErrUnsafeFile", err) + } +} + +func writePrivateWindowsAuthFiles(t *testing.T) string { + t.Helper() + directory := writeAuthFiles(t, defaultAuthYAML, registryYAML(adminUserYAML("admin", "Admin", true, "admin"))) + if err := safeio.ProtectPrivateDirectory(directory); err != nil { + t.Fatal(err) + } + for _, name := range []string{"auth.yaml", "users.yaml"} { + if err := safeio.ProtectPrivateRegular(filepath.Join(directory, name)); err != nil { + t.Fatal(err) + } + } + return directory +} + +func setPermissiveAuthDACL(t *testing.T, path string) { + t.Helper() + world, err := windows.StringToSid("S-1-1-0") + if err != nil { + t.Fatal(err) + } + var pinner runtime.Pinner + pinner.Pin(world) + defer pinner.Unpin() + acl, err := windows.ACLFromEntries([]windows.EXPLICIT_ACCESS{{ + AccessPermissions: windows.GENERIC_READ | windows.GENERIC_WRITE, + AccessMode: windows.GRANT_ACCESS, + Trustee: windows.TRUSTEE{ + TrusteeForm: windows.TRUSTEE_IS_SID, + TrusteeType: windows.TRUSTEE_IS_GROUP, + TrusteeValue: windows.TrusteeValueFromSID(world), + }, + }}, nil) + if err != nil { + t.Fatal(err) + } + if err := windows.SetNamedSecurityInfo(path, windows.SE_FILE_OBJECT, + windows.DACL_SECURITY_INFORMATION|windows.PROTECTED_DACL_SECURITY_INFORMATION, + nil, nil, acl, nil); err != nil { + t.Fatal(err) + } +} diff --git a/tools/tht/internal/safeio/files.go b/tools/tht/internal/safeio/files.go index 8a0a0fa1..0ea916f3 100644 --- a/tools/tht/internal/safeio/files.go +++ b/tools/tht/internal/safeio/files.go @@ -77,6 +77,10 @@ func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error return ErrUnsafeFile } defer file.Close() + if err := ProtectPrivateRegular(path); err != nil { + _ = os.Remove(path) + return ErrUnsafeFile + } if _, err := file.Write(contents); err != nil { _ = os.Remove(path) return ErrUnsafeFile diff --git a/tools/tht/internal/safeio/files_windows.go b/tools/tht/internal/safeio/files_windows.go index 020a86c7..0040971d 100644 --- a/tools/tht/internal/safeio/files_windows.go +++ b/tools/tht/internal/safeio/files_windows.go @@ -57,6 +57,10 @@ func ReadCanonicalRegular(path string, maximum int64) ([]byte, error) { } func openWindowsComponent(path string, directory bool) (windows.Handle, error) { + return openWindowsComponentWithAccess(path, directory, windows.GENERIC_READ) +} + +func openWindowsComponentWithAccess(path string, directory bool, access uint32) (windows.Handle, error) { flags := uint32(windows.FILE_FLAG_OPEN_REPARSE_POINT) if directory { flags |= windows.FILE_FLAG_BACKUP_SEMANTICS @@ -65,7 +69,7 @@ func openWindowsComponent(path string, directory bool) (windows.Handle, error) { } handle, err := windows.CreateFile( windows.StringToUTF16Ptr(path), - windows.GENERIC_READ, + access, windowsRetainedHandleShareMode, nil, windows.OPEN_EXISTING, diff --git a/tools/tht/internal/safeio/private_unix.go b/tools/tht/internal/safeio/private_unix.go new file mode 100644 index 00000000..8630a63c --- /dev/null +++ b/tools/tht/internal/safeio/private_unix.go @@ -0,0 +1,61 @@ +//go:build !windows + +package safeio + +import ( + "os" + "path/filepath" +) + +// ProtectPrivateDirectory sets the private directory mode used for local authentication state. +func ProtectPrivateDirectory(path string) error { + if err := ValidateCanonicalPath(path); err != nil { + return err + } + if err := os.Chmod(path, 0o700); err != nil { + return ErrUnsafeFile + } + return ValidatePrivateDirectory(path) +} + +// ValidatePrivateDirectory requires a canonical, non-symlinked directory with no group or world access. +func ValidatePrivateDirectory(path string) error { + if err := requireCanonicalDirectory(path); err != nil { + return err + } + info, err := os.Lstat(path) + if err != nil || !info.IsDir() || info.Mode()&os.ModeSymlink != 0 || !isExactPrivateMode(info.Mode(), 0o700) { + return ErrUnsafeFile + } + return nil +} + +// ProtectPrivateRegular sets the private file mode used for local authentication files. +func ProtectPrivateRegular(path string) error { + if err := ValidateCanonicalPath(path); err != nil { + return err + } + if err := os.Chmod(path, 0o600); err != nil { + return ErrUnsafeFile + } + return ValidatePrivateRegular(path) +} + +// ValidatePrivateRegular requires a canonical, single-link private regular file. +func ValidatePrivateRegular(path string) error { + if err := ValidateCanonicalPath(path); err != nil { + return err + } + if err := requireCanonicalDirectory(filepath.Dir(path)); err != nil { + return err + } + info, err := os.Lstat(path) + if err != nil || !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || !hasSingleLink(info) || !isExactPrivateMode(info.Mode(), 0o600) { + return ErrUnsafeFile + } + return nil +} + +func isExactPrivateMode(mode os.FileMode, permissions os.FileMode) bool { + return mode.Perm() == permissions && mode&(os.ModeSetuid|os.ModeSetgid|os.ModeSticky) == 0 +} diff --git a/tools/tht/internal/safeio/private_windows.go b/tools/tht/internal/safeio/private_windows.go new file mode 100644 index 00000000..49ca9283 --- /dev/null +++ b/tools/tht/internal/safeio/private_windows.go @@ -0,0 +1,192 @@ +//go:build windows + +package safeio + +import ( + "path/filepath" + "runtime" + "strings" + "unsafe" + + "golang.org/x/sys/windows" +) + +// ProtectPrivateDirectory sets a protected DACL containing only the current owner. +func ProtectPrivateDirectory(path string) error { + parents, target, err := openCanonicalWindowsParent(path) + if err != nil { + return ErrUnsafeFile + } + defer parents.Close() + handle, err := openWindowsComponentWithAccess(filepath.Join(parents.directory, target), true, windows.GENERIC_READ|windows.WRITE_DAC|windows.WRITE_OWNER) + if err != nil { + return ErrUnsafeFile + } + defer windows.CloseHandle(handle) + if err := setOwnerOnlyDACL(handle); err != nil { + return ErrUnsafeFile + } + return validateOwnerOnlyDACL(handle) +} + +// ValidatePrivateDirectory requires a canonical directory protected for its current owner only. +func ValidatePrivateDirectory(path string) error { + parents, target, err := openCanonicalWindowsParent(path) + if err != nil { + return ErrUnsafeFile + } + defer parents.Close() + handle, err := openWindowsComponent(filepath.Join(parents.directory, target), true) + if err != nil { + return ErrUnsafeFile + } + defer windows.CloseHandle(handle) + if err := validateOwnerOnlyDACL(handle); err != nil { + return ErrUnsafeFile + } + return nil +} + +// ProtectPrivateRegular sets a protected DACL containing only the current owner. +func ProtectPrivateRegular(path string) error { + parents, target, err := openCanonicalWindowsParent(path) + if err != nil { + return ErrUnsafeFile + } + defer parents.Close() + handle, err := openWindowsComponentWithAccess(filepath.Join(parents.directory, target), false, windows.GENERIC_READ|windows.WRITE_DAC|windows.WRITE_OWNER) + if err != nil { + return ErrUnsafeFile + } + defer windows.CloseHandle(handle) + if err := setOwnerOnlyDACL(handle); err != nil { + return ErrUnsafeFile + } + return validateOwnerOnlyDACL(handle) +} + +// ValidatePrivateRegular requires a canonical, single-link file protected for its current owner only. +func ValidatePrivateRegular(path string) error { + parents, target, err := openCanonicalWindowsParent(path) + if err != nil { + return ErrUnsafeFile + } + defer parents.Close() + handle, err := openWindowsComponent(filepath.Join(parents.directory, target), false) + if err != nil { + return ErrUnsafeFile + } + defer windows.CloseHandle(handle) + if err := validateOwnerOnlyDACL(handle); err != nil { + return ErrUnsafeFile + } + return nil +} + +type windowsParentHandles struct { + directory string + handles []windows.Handle +} + +func (parents *windowsParentHandles) Close() { + for index := len(parents.handles) - 1; index >= 0; index-- { + _ = windows.CloseHandle(parents.handles[index]) + } +} + +// openCanonicalWindowsParent retains every directory handle from the volume root through the +// target parent without FILE_SHARE_DELETE. The resulting parent cannot be renamed or replaced by +// a reparse point while an operation uses its absolute child paths. +func openCanonicalWindowsParent(path string) (*windowsParentHandles, string, error) { + if err := ValidateCanonicalPath(path); err != nil { + return nil, "", err + } + volume := filepath.VolumeName(path) + root := volume + `\` + components := strings.Split(strings.TrimPrefix(path, root), `\`) + if volume == "" || len(components) == 0 || components[0] == "" { + return nil, "", ErrUnsafeFile + } + parents := &windowsParentHandles{directory: root} + rootHandle, err := openWindowsComponent(root, true) + if err != nil { + return nil, "", err + } + parents.handles = append(parents.handles, rootHandle) + for _, component := range components[:len(components)-1] { + parents.directory = filepath.Join(parents.directory, component) + handle, err := openWindowsComponent(parents.directory, true) + if err != nil { + parents.Close() + return nil, "", err + } + parents.handles = append(parents.handles, handle) + } + return parents, components[len(components)-1], nil +} + +func setOwnerOnlyDACL(handle windows.Handle) error { + sid, err := currentOwnerSID() + if err != nil { + return err + } + var pinner runtime.Pinner + pinner.Pin(sid) + defer pinner.Unpin() + 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), + }, + }}, nil) + if err != nil { + return err + } + 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) +} + +func validateOwnerOnlyDACL(handle windows.Handle) error { + ownerSID, err := currentOwnerSID() + if err != nil { + return err + } + 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 +} + +func currentOwnerSID() (*windows.SID, error) { + user, err := windows.GetCurrentProcessToken().GetTokenUser() + if err != nil || user == nil || user.User.Sid == nil { + return nil, ErrUnsafeFile + } + return user.User.Sid, nil +} diff --git a/tools/tht/internal/safeio/private_windows_test.go b/tools/tht/internal/safeio/private_windows_test.go new file mode 100644 index 00000000..2be8812c --- /dev/null +++ b/tools/tht/internal/safeio/private_windows_test.go @@ -0,0 +1,169 @@ +//go:build windows + +package safeio + +import ( + "errors" + "os" + "path/filepath" + "runtime" + "testing" + + "golang.org/x/sys/windows" +) + +func TestPrivateWindowsDACLRejectsPermissiveDirectoryAndRegularFile(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) + } + if err := ValidatePrivateDirectory(directory); err != nil { + t.Fatalf("ValidatePrivateDirectory() protected directory error = %v", 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) + } + if err := ValidatePrivateRegular(path); err != nil { + t.Fatalf("ValidatePrivateRegular() protected file error = %v", err) + } + + for name, path := range map[string]string{"directory": directory, "regular file": path} { + t.Run(name, func(t *testing.T) { + if err := setPermissiveDACL(path); err != nil { + t.Fatal(err) + } + var err error + if name == "directory" { + err = ValidatePrivateDirectory(path) + } else { + err = ValidatePrivateRegular(path) + } + if !errors.Is(err, ErrUnsafeFile) { + t.Fatalf("private validation error = %v, want ErrUnsafeFile", err) + } + }) + } +} + +func TestReplaceCanonicalRegularCreatesPrivateTemporaryAndReplacement(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("old"), 0o600); err != nil { + t.Fatal(err) + } + if err := ProtectPrivateRegular(path); err != nil { + t.Fatal(err) + } + + temporary, err := writePrivateTemporary(directory, []byte("temporary")) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.Remove(temporary) }) + if err := ValidatePrivateRegular(temporary); err != nil { + t.Fatalf("temporary DACL error = %v", err) + } + + if err := ReplaceCanonicalRegular(path, []byte("replacement"), 0o600); err != nil { + t.Fatal(err) + } + if err := ValidatePrivateRegular(path); err != nil { + t.Fatalf("replacement DACL error = %v", err) + } +} + +func TestReplaceCanonicalRegularRejectsReparseParent(t *testing.T) { + root := t.TempDir() + realDirectory := filepath.Join(root, "real") + if err := os.Mkdir(realDirectory, 0o700); err != nil { + t.Fatal(err) + } + if err := ProtectPrivateDirectory(realDirectory); err != nil { + t.Fatal(err) + } + target := filepath.Join(realDirectory, "users.yaml") + if err := os.WriteFile(target, []byte("old"), 0o600); err != nil { + t.Fatal(err) + } + if err := ProtectPrivateRegular(target); err != nil { + t.Fatal(err) + } + + parentLink := filepath.Join(root, "reparse-parent") + if err := os.Symlink(realDirectory, parentLink); err != nil { + t.Skipf("Windows host does not permit test symlink creation: %v", err) + } + if err := ReplaceCanonicalRegular(filepath.Join(parentLink, "users.yaml"), []byte("new"), 0o600); !errors.Is(err, ErrUnsafeFile) { + t.Fatalf("reparse-parent replacement error = %v, want ErrUnsafeFile", err) + } +} + +func TestOpenCanonicalWindowsParentBlocksParentRename(t *testing.T) { + root := t.TempDir() + realDirectory := filepath.Join(root, "real") + if err := os.Mkdir(realDirectory, 0o700); err != nil { + t.Fatal(err) + } + if err := ProtectPrivateDirectory(realDirectory); err != nil { + t.Fatal(err) + } + target := filepath.Join(realDirectory, "users.yaml") + if err := os.WriteFile(target, []byte("old"), 0o600); err != nil { + t.Fatal(err) + } + if err := ProtectPrivateRegular(target); err != nil { + t.Fatal(err) + } + parents, _, err := openCanonicalWindowsParent(target) + if err != nil { + t.Fatal(err) + } + renamed := realDirectory + "-renamed" + if err := os.Rename(realDirectory, renamed); err == nil { + parents.Close() + t.Fatal("parent rename succeeded while canonical replacement handles were retained") + } + parents.Close() + if err := os.Rename(realDirectory, renamed); err != nil { + t.Fatalf("parent rename after closing canonical replacement handles: %v", err) + } +} + +func setPermissiveDACL(path string) error { + world, err := windows.StringToSid("S-1-1-0") + if err != nil { + return err + } + var pinner runtime.Pinner + pinner.Pin(world) + defer pinner.Unpin() + acl, err := windows.ACLFromEntries([]windows.EXPLICIT_ACCESS{{ + AccessPermissions: windows.GENERIC_READ | windows.GENERIC_WRITE, + AccessMode: windows.GRANT_ACCESS, + Trustee: windows.TRUSTEE{ + TrusteeForm: windows.TRUSTEE_IS_SID, + TrusteeType: windows.TRUSTEE_IS_GROUP, + TrusteeValue: windows.TrusteeValueFromSID(world), + }, + }}, nil) + if err != nil { + return err + } + return windows.SetNamedSecurityInfo(path, windows.SE_FILE_OBJECT, + windows.DACL_SECURITY_INFORMATION|windows.PROTECTED_DACL_SECURITY_INFORMATION, + nil, nil, acl, nil) +} diff --git a/tools/tht/internal/safeio/replace_windows.go b/tools/tht/internal/safeio/replace_windows.go index 257c8060..5b1302fd 100644 --- a/tools/tht/internal/safeio/replace_windows.go +++ b/tools/tht/internal/safeio/replace_windows.go @@ -14,38 +14,41 @@ import ( const windowsReplaceMoveFlags = windows.MOVEFILE_REPLACE_EXISTING | windows.MOVEFILE_WRITE_THROUGH func replaceCanonicalRegular(path string, contents []byte) error { - directory := filepath.Dir(path) - if err := requireCanonicalDirectory(directory); err != nil || !safeExistingRegular(path) { + parents, target, err := openCanonicalWindowsParent(path) + if err != nil { return ErrUnsafeFile } - temporary, err := writePrivateTemporary(directory, contents) + defer parents.Close() + if len(parents.handles) == 0 || validateOwnerOnlyDACL(parents.handles[len(parents.handles)-1]) != nil || !safeExistingRegular(filepath.Join(parents.directory, target)) { + return ErrUnsafeFile + } + temporary, err := writePrivateTemporary(parents.directory, contents) if err != nil { return ErrUnsafeFile } defer func() { _ = os.Remove(temporary) }() - if !safeExistingRegular(path) { + if !safeExistingRegular(filepath.Join(parents.directory, target)) { return ErrUnsafeFile } from, err := windows.UTF16PtrFromString(temporary) if err != nil { return ErrUnsafeFile } - to, err := windows.UTF16PtrFromString(path) + to, err := windows.UTF16PtrFromString(filepath.Join(parents.directory, target)) if err != nil { return ErrUnsafeFile } if err := windows.MoveFileEx(from, to, windowsReplaceMoveFlags); err != nil { return ErrUnsafeFile } + if err := ValidatePrivateRegular(filepath.Join(parents.directory, target)); err != nil { + return ErrUnsafeFile + } return nil } func safeExistingRegular(path string) bool { - handle, err := openWindowsComponent(path, false) - if err != nil { - return false - } - return windows.CloseHandle(handle) == nil + return ValidatePrivateRegular(path) == nil } func writePrivateTemporary(directory string, contents []byte) (string, error) { @@ -62,6 +65,11 @@ func writePrivateTemporary(directory string, contents []byte) (string, error) { 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)