//go:build windows package safeio import ( "errors" "io" "os" "path/filepath" "golang.org/x/sys/windows" ) type windowsPrivateRegular struct { parents *windowsParentHandles handle windows.Handle path string info windows.ByHandleFileInformation } func (value *windowsPrivateRegular) Close() { if value.handle != 0 { _ = windows.CloseHandle(value.handle) value.handle = 0 } if value.parents != nil { value.parents.Close() value.parents = nil } } func openWindowsPrivateRegular(path string, links uint32) (*windowsPrivateRegular, error) { parents, target, err := openCanonicalWindowsParent(path) if err != nil || parents == nil || len(parents.handles) == 0 || validateOwnerOnlyDACL(parents.handles[len(parents.handles)-1]) != nil { if parents != nil { parents.Close() } return nil, ErrUnsafeFile } value, err := openWindowsPrivateRegularAt(parents.handles[len(parents.handles)-1], target, windows.GENERIC_READ, links) if err != nil { parents.Close() return nil, err } handle := value.handle value.handle = 0 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) if err != nil { return false, ErrUnsafeFile } defer claimFile.Close() if !sameWindowsPrivateFile(sourceFile.info, claimFile.info) { return false, ErrUnsafeFile } return true, nil } func readCanonicalPrivateClaim(source, claim string, maximum int64) ([]byte, bool, error) { sourceFile, err := openWindowsPrivateRegular(source, 2) if err != nil { if isWindowsNotFound(err) && windowsPrivateClaimAbsentOrOrphan(claim) { return nil, false, nil } 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 } func removeCanonicalPrivateClaim(source, claim string) (bool, error) { sourceFile, err := openWindowsPrivateRegular(source, 2) if err != nil { if isWindowsNotFound(err) && windowsPrivateClaimAbsentOrOrphan(claim) { return false, nil } return false, ErrUnsafeFile } claimFile, err := openWindowsPrivateRegular(claim, 2) if err != nil { sourceFile.Close() if isWindowsNotFound(err) { return false, nil } return false, ErrUnsafeFile } if !sameWindowsPrivateFile(sourceFile.info, claimFile.info) { sourceFile.Close() claimFile.Close() return false, ErrUnsafeFile } sourceFile.Close() claimFile.Close() if err := windows.DeleteFile(windows.StringToUTF16Ptr(source)); err != nil { return false, ErrUnsafeFile } remaining, err := openWindowsPrivateRegular(claim, 1) if err != nil { return false, ErrUnsafeFile } remaining.Close() if err := windows.DeleteFile(windows.StringToUTF16Ptr(claim)); err != nil { return false, ErrUnsafeFile } return true, nil } 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 }