From feee4ee648536850d06ed4db637a95bf7726c597 Mon Sep 17 00:00:00 2001 From: mptyl Date: Tue, 18 Aug 2026 14:45:24 +0200 Subject: [PATCH] refactor(windows): unify private claim primitives --- tools/tht/internal/safeio/claim_windows.go | 139 +++++---------------- 1 file changed, 33 insertions(+), 106 deletions(-) diff --git a/tools/tht/internal/safeio/claim_windows.go b/tools/tht/internal/safeio/claim_windows.go index 86ec8986..1173b6c5 100644 --- a/tools/tht/internal/safeio/claim_windows.go +++ b/tools/tht/internal/safeio/claim_windows.go @@ -3,9 +3,6 @@ package safeio import ( - "errors" - "io" - "os" "path/filepath" "golang.org/x/sys/windows" @@ -14,7 +11,6 @@ import ( type windowsPrivateRegular struct { parents *windowsParentHandles handle windows.Handle - path string info windows.ByHandleFileInformation } @@ -47,91 +43,59 @@ func openWindowsPrivateRegular(path string, links uint32) (*windowsPrivateRegula return &windowsPrivateRegular{ parents: parents, handle: handle, - path: filepath.Join(parents.directory, target), info: value.info, }, nil } -func claimCanonicalPrivateRegular(source, claim string) (bool, error) { - sourceFile, err := openWindowsPrivateRegular(source, 1) - if err != nil { - if windowsPrivateClaimPairExists(source, claim) { - return false, nil - } - if isWindowsNotFound(err) && windowsPrivateClaimAbsentOrOrphan(claim) { - return false, nil - } - return false, ErrUnsafeFile - } - defer sourceFile.Close() - if err := windows.CreateHardLink(windows.StringToUTF16Ptr(claim), windows.StringToUTF16Ptr(sourceFile.path), 0); err != nil { - if errors.Is(err, windows.ERROR_FILE_EXISTS) || errors.Is(err, windows.ERROR_ALREADY_EXISTS) { - if existing, existingErr := openWindowsPrivateRegularWithAllowedLinks(claim, 1, 2); existingErr == nil { - existing.Close() - return false, nil - } - } - return false, ErrUnsafeFile - } - if err := windows.GetFileInformationByHandle(sourceFile.handle, &sourceFile.info); err != nil || sourceFile.info.NumberOfLinks != 2 { - return false, ErrUnsafeFile - } - claimFile, err := openWindowsPrivateRegular(claim, 2) +func claimCanonicalPrivateRegular(source, claim string) (claimed bool, resultErr error) { + directory, sourceName, claimName, err := openWindowsPrivateClaimDirectory(source, claim) if err != nil { return false, ErrUnsafeFile } - defer claimFile.Close() - if !sameWindowsPrivateFile(sourceFile.info, claimFile.info) { - return false, ErrUnsafeFile - } - return true, nil + defer func() { + if closeErr := directory.Close(); closeErr != nil && resultErr == nil { + claimed = false + resultErr = ErrUnsafeFile + } + }() + return directory.ClaimRegular(sourceName, claimName) } -func readCanonicalPrivateClaim(source, claim string, maximum int64) ([]byte, bool, error) { - sourceFile, err := openWindowsPrivateRegular(source, 2) +func readCanonicalPrivateClaim(source, claim string, maximum int64) (contents []byte, found bool, resultErr error) { + directory, sourceName, claimName, err := openWindowsPrivateClaimDirectory(source, claim) if err != nil { - if isWindowsNotFound(err) && windowsPrivateClaimAbsentOrOrphan(claim) { - return nil, false, nil + return nil, false, ErrUnsafeFile + } + defer func() { + if closeErr := directory.Close(); closeErr != nil && resultErr == nil { + contents = nil + found = false + resultErr = ErrUnsafeFile } - return nil, false, ErrUnsafeFile - } - defer sourceFile.Close() - claimFile, err := openWindowsPrivateRegular(claim, 2) - if err != nil { - if isWindowsNotFound(err) { - return nil, false, nil - } - return nil, false, ErrUnsafeFile - } - defer claimFile.Close() - if !sameWindowsPrivateFile(sourceFile.info, claimFile.info) { - return nil, false, ErrUnsafeFile - } - file := os.NewFile(uintptr(sourceFile.handle), "tht-safeio-oidc-claim") - if file == nil { - return nil, false, ErrUnsafeFile - } - contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1)) - file.Close() - sourceFile.handle = 0 - if readErr != nil || int64(len(contents)) > maximum { - return nil, false, ErrUnsafeFile - } - if err := windows.GetFileInformationByHandle(claimFile.handle, &claimFile.info); err != nil || claimFile.info.NumberOfLinks != 2 || !sameWindowsPrivateFile(sourceFile.info, claimFile.info) { - return nil, false, ErrUnsafeFile - } - return contents, true, nil + }() + return directory.ReadClaim(sourceName, claimName, maximum) } -func removeCanonicalPrivateClaim(source, claim string) (removed bool, resultErr error) { +func openWindowsPrivateClaimDirectory(source, claim string) (PrivateDirectoryHandle, string, string, error) { parentPath := filepath.Dir(source) sourceName := filepath.Base(source) claimName := filepath.Base(claim) if parentPath != filepath.Dir(claim) || !validPrivateLeafName(sourceName) || !validPrivateLeafName(claimName) { - return false, ErrUnsafeFile + return nil, "", "", ErrUnsafeFile } directory, found, err := OpenPrivateDirectory(parentPath, false) if err != nil || !found { + if directory != nil { + _ = directory.Close() + } + return nil, "", "", ErrUnsafeFile + } + return directory, sourceName, claimName, nil +} + +func removeCanonicalPrivateClaim(source, claim string) (removed bool, resultErr error) { + directory, sourceName, claimName, err := openWindowsPrivateClaimDirectory(source, claim) + if err != nil { return false, ErrUnsafeFile } defer func() { @@ -143,43 +107,6 @@ func removeCanonicalPrivateClaim(source, claim string) (removed bool, resultErr return directory.RemoveClaim(sourceName, claimName) } -func openWindowsPrivateRegularWithAllowedLinks(path string, allowed ...uint32) (*windowsPrivateRegular, error) { - for _, links := range allowed { - value, err := openWindowsPrivateRegular(path, links) - if err == nil { - return value, nil - } - } - return nil, ErrUnsafeFile -} - -func windowsPrivateClaimAbsentOrOrphan(claim string) bool { - claimFile, err := openWindowsPrivateRegular(claim, 1) - if err != nil { - return isWindowsNotFound(err) - } - claimFile.Close() - return true -} - -func windowsPrivateClaimPairExists(source, claim string) bool { - sourceFile, sourceErr := openWindowsPrivateRegular(source, 2) - if sourceErr != nil { - return false - } - defer sourceFile.Close() - claimFile, claimErr := openWindowsPrivateRegular(claim, 2) - if claimErr != nil { - return false - } - defer claimFile.Close() - return sameWindowsPrivateFile(sourceFile.info, claimFile.info) -} - -func isWindowsNotFound(err error) bool { - return errors.Is(err, windows.ERROR_FILE_NOT_FOUND) || errors.Is(err, windows.ERROR_PATH_NOT_FOUND) -} - func sameWindowsPrivateFile(left, right windows.ByHandleFileInformation) bool { return left.VolumeSerialNumber == right.VolumeSerialNumber && left.FileIndexHigh == right.FileIndexHigh && left.FileIndexLow == right.FileIndexLow && left.FileSizeHigh == right.FileSizeHigh && left.FileSizeLow == right.FileSizeLow && left.LastWriteTime == right.LastWriteTime }