// SPDX-License-Identifier: MIT #include #include #include #include #include #include #include "InvalidationTracker.h" #include #include namespace FEX::Windows { InvalidationTracker::InvalidationTracker(FEXCore::Context::Context& CTX, const std::unordered_map& Threads) : CTX {CTX} , Threads {Threads} { FEX_CONFIG_OPT(SMCChecks, SMCCHECKS); SMCDetectionDisabled = (SMCChecks == FEXCore::Config::CONFIG_SMC_NONE); MEMORY_BASIC_INFORMATION Info; uint64_t Address = 0; while (VirtualQuery(reinterpret_cast(Address), &Info, sizeof(Info))) { uint64_t BaseAddress = reinterpret_cast(Info.BaseAddress); if (Info.State == MEM_COMMIT) { HandleMemoryProtectionNotification(BaseAddress, Info.RegionSize, Info.Protect); } Address = BaseAddress + Info.RegionSize; } } static bool ProtHasExec(ULONG Prot) { return (Prot & (PAGE_EXECUTE | PAGE_EXECUTE_READ | PAGE_EXECUTE_READWRITE | PAGE_EXECUTE_WRITECOPY)) != 0; } static bool ProtIsReadable(ULONG Prot) { return (Prot & (PAGE_READONLY | PAGE_READWRITE | PAGE_WRITECOPY | PAGE_EXECUTE | PAGE_EXECUTE_READ | PAGE_EXECUTE_READWRITE | PAGE_EXECUTE_WRITECOPY)) != 0; } static bool ProtIsWritable(ULONG Prot) { return (Prot & (PAGE_READWRITE | PAGE_WRITECOPY | PAGE_EXECUTE_READWRITE | PAGE_EXECUTE_WRITECOPY)) != 0; } void InvalidationTracker::HandleMemoryProtectionNotification(uint64_t Address, uint64_t Size, ULONG Prot) { const auto AlignedBase = Address & FEXCore::Utils::FEX_PAGE_MASK; const auto AlignedSize = (Address - AlignedBase + Size + FEXCore::Utils::FEX_PAGE_SIZE - 1) & FEXCore::Utils::FEX_PAGE_MASK; const bool NeedsInvalidate = [&]() { std::unique_lock Lock(IntervalsLock); FEXCore::IntervalList::Interval ProtInterval {AlignedBase, AlignedBase + AlignedSize}; const bool HasExec = ProtHasExec(Prot); const bool EffectiveExec = HasExec || (DEPDisabled && ProtIsReadable(Prot)); const bool EffectiveRWX = EffectiveExec && ProtIsWritable(Prot); if (EffectiveExec) { XIntervals.Insert(ProtInterval); if (EffectiveRWX) { LogMan::Msg::DFmt("Add SMC interval: {:X} - {:X}", AlignedBase, AlignedBase + AlignedSize); RWXIntervals.Insert(ProtInterval); } if (DEPDisabled && !HasExec) { DEPPromotedIntervals.Insert(ProtInterval); } return true; } else if (XIntervals.Intersect(ProtInterval)) { XIntervals.Remove(ProtInterval); RWXIntervals.Remove(ProtInterval); if (DEPDisabled) { DEPPromotedIntervals.Remove(ProtInterval); } return true; } return false; }(); if (NeedsInvalidate) { // IntervalsLock cannot be held during invalidation InvalidateIntervalInternal(AlignedBase, AlignedSize); } } void InvalidationTracker::HandleProcessExecuteFlagsChange(ULONG Flags) { const bool DisableDEP = (Flags & MEM_EXECUTE_OPTION_ENABLE) != 0; std::scoped_lock CodeLock(CTX.GetCodeInvalidationMutex()); std::unique_lock Lock(IntervalsLock); if (DisableDEP == DEPDisabled) { return; } DEPDisabled = DisableDEP; if (DisableDEP) { DEPPromotedIntervals.Clear(); MEMORY_BASIC_INFORMATION Info; uint64_t Address = 0; while (VirtualQuery(reinterpret_cast(Address), &Info, sizeof(Info))) { uint64_t BaseAddress = reinterpret_cast(Info.BaseAddress); if (Info.State == MEM_COMMIT && ProtIsReadable(Info.Protect) && !ProtHasExec(Info.Protect)) { const auto AlignedBase = BaseAddress & FEXCore::Utils::FEX_PAGE_MASK; const auto AlignedSize = (BaseAddress - AlignedBase + Info.RegionSize + FEXCore::Utils::FEX_PAGE_SIZE - 1) & FEXCore::Utils::FEX_PAGE_MASK; FEXCore::IntervalList::Interval ProtInterval {AlignedBase, AlignedBase + AlignedSize}; XIntervals.Insert(ProtInterval); if (ProtIsWritable(Info.Protect)) { RWXIntervals.Insert(ProtInterval); } DEPPromotedIntervals.Insert(ProtInterval); } Address = BaseAddress + Info.RegionSize; } } else { for (const auto& Interval : DEPPromotedIntervals) { XIntervals.Remove(Interval); RWXIntervals.Remove(Interval); } DEPPromotedIntervals.Clear(); } // Invalidate all cached code: previously-compiled blocks may contain NoExec stubs for addresses // that are now executable (or reference regions whose executability just changed). InvalidateIntervalInternalLocked(0, std::numeric_limits::max()); } void InvalidationTracker::HandleImageMap(std::string_view Name, uint64_t Address) { auto* Nt = RtlImageNtHeader(reinterpret_cast(Address)); auto* SectionsBegin = IMAGE_FIRST_SECTION(Nt); auto* SectionsEnd = SectionsBegin + Nt->FileHeader.NumberOfSections; uint64_t LastExecutableSectionEnd = 0; for (auto* Section = SectionsBegin; Section != SectionsEnd; Section++) { if (Section->Characteristics & IMAGE_SCN_MEM_EXECUTE) { std::unique_lock Lock(IntervalsLock); uint64_t SectionBase = Address + Section->VirtualAddress; uint64_t SectionEnd = SectionBase + Section->Misc.VirtualSize; XIntervals.Insert({SectionBase, SectionEnd}); LastExecutableSectionEnd = std::max(LastExecutableSectionEnd, SectionEnd); if (Section->Characteristics & IMAGE_SCN_MEM_WRITE) { LogMan::Msg::DFmt("Add image SMC interval: {:X} - {:X}", SectionBase, SectionBase + Section->Misc.VirtualSize); RWXIntervals.Insert({SectionBase, SectionBase + Section->Misc.VirtualSize}); } } } FEX_CONFIG_OPT(MonoHacks, MONOHACKS); if (MonoHacks && (Name == "mono-2.0-bdwgc.dll" || Name == "mono.dll")) { FEX_CONFIG_OPT(MaxInst, MAXINST); FEX_CONFIG_OPT(Multiblock, MULTIBLOCK); if (Multiblock && MaxInst() >= 500) { // Require these settings to ensure we can safely hook all SMC sites in a single block CTX.MarkMonoDetected(); MonoBackpatcherDetectionPending = true; MonoBase = Address; MonoEnd = LastExecutableSectionEnd; } else { LogMan::Msg::IFmt("Not applying mono hacks, Multiblock with MaxInst >= 500 required"); } } } InvalidationTracker::InvalidateContainingSectionResult InvalidationTracker::InvalidateContainingSection(uint64_t Address, bool Free) { MEMORY_BASIC_INFORMATION Info; if (NtQueryVirtualMemory(NtCurrentProcess(), reinterpret_cast(Address), MemoryBasicInformation, &Info, sizeof(Info), nullptr)) { return {Address, 0}; } const auto SectionBase = reinterpret_cast(Info.AllocationBase); auto SectionSize = reinterpret_cast(Info.BaseAddress) + Info.RegionSize - SectionBase; while (!NtQueryVirtualMemory(NtCurrentProcess(), reinterpret_cast(SectionBase + SectionSize), MemoryBasicInformation, &Info, sizeof(Info), nullptr) && reinterpret_cast(Info.AllocationBase) == SectionBase) { SectionSize += Info.RegionSize; } InvalidateIntervalInternal(SectionBase, SectionSize); if (Free) { std::unique_lock Lock(IntervalsLock); XIntervals.Remove({SectionBase, SectionBase + SectionSize}); RWXIntervals.Remove({SectionBase, SectionBase + SectionSize}); } return {SectionBase, SectionSize}; } void InvalidationTracker::InvalidateAlignedInterval(uint64_t Address, uint64_t Size, bool Free) { if (!Address) { // Match the Windows behaviour when passed a NULL base address. Size = std::numeric_limits::max(); } const auto AlignedBase = Address & FEXCore::Utils::FEX_PAGE_MASK; const auto AlignedSize = std::max(Size, (Address - AlignedBase + Size + FEXCore::Utils::FEX_PAGE_SIZE - 1) & FEXCore::Utils::FEX_PAGE_MASK); InvalidateIntervalInternal(AlignedBase, AlignedSize); if (Free) { std::unique_lock Lock(IntervalsLock); XIntervals.Remove({AlignedBase, AlignedBase + AlignedSize}); RWXIntervals.Remove({AlignedBase, AlignedBase + AlignedSize}); } } void InvalidationTracker::ReprotectRWXIntervals(uint64_t Address, uint64_t Size) { ProtectRWXIntervalsInternal(Address, Size, false); } bool InvalidationTracker::HandleRWXAccessViolation(FEXCore::Core::InternalThreadState* Thread, uint64_t HostPc, uint64_t FaultAddress) { const auto [NeedsInvalidate, UntrapProt] = [&](uint64_t Address) -> std::pair { std::shared_lock Lock(IntervalsLock); if (!RWXIntervals.Query(Address).Enclosed) { return {false, 0}; } return {true, GetUntrapProt(Address)}; }(FaultAddress); if (NeedsInvalidate) { // IntervalsLock cannot be held during invalidation { std::scoped_lock Lock(CTX.GetCodeInvalidationMutex()); InvalidateIntervalInternalLocked(FaultAddress & FEXCore::Utils::FEX_PAGE_MASK, FEXCore::Utils::FEX_PAGE_SIZE); // Invalidate, then unprotect the faulting page with the compilation lock held to ensure that any racing invalidations are not dropped. ULONG TmpProt; void* TmpAddress = reinterpret_cast(FaultAddress); SIZE_T TmpSize = 1; NtProtectVirtualMemory(NtCurrentProcess(), &TmpAddress, &TmpSize, UntrapProt, &TmpProt); } DetectMonoBackpatcherBlock(Thread, HostPc); return true; } return false; } bool InvalidationTracker::BeginUntrackedWriteLocked(uint64_t Address, uint64_t Size) { return ProtectRWXIntervalsInternal(Address, Size, true); } FEXCore::HLE::ExecutableRangeInfo InvalidationTracker::QueryExecutableRange(uint64_t Address) { std::shared_lock Lock(IntervalsLock); const auto XResult = XIntervals.Query(Address); if (!XResult.Enclosed) { return {}; } const auto RWXResult = RWXIntervals.Query(Address); if (RWXResult.Enclosed) { return {RWXResult.Interval.Offset, RWXResult.Interval.End - RWXResult.Interval.Offset, true}; } else if (RWXResult.Size && RWXResult.Size < XResult.Size) { return {XResult.Interval.Offset, RWXResult.Interval.Offset - XResult.Interval.Offset, false}; } return {XResult.Interval.Offset, XResult.Interval.End - XResult.Interval.Offset, false}; } void InvalidationTracker::DetectMonoBackpatcherBlock(FEXCore::Core::InternalThreadState* Thread, uint64_t HostPc) { if (!MonoBackpatcherDetectionPending) { return; } if (!CTX.IsAddressInCodeBuffer(Thread, HostPc)) { return; } uint64_t RIP = CTX.RestoreRIPFromHostPC(Thread, HostPc); if (!RIP || RIP < MonoBase || RIP >= MonoEnd) { return; } static constexpr uint8_t XChgOp = 0x87; if (*reinterpret_cast(RIP) != XChgOp && *reinterpret_cast(RIP + 1) != XChgOp) { return; } uint64_t BlockEntry = CTX.GetGuestBlockEntry(Thread); LogMan::Msg::DFmt("Detected mono backpatcher at: {:X}", BlockEntry); DisableSMCDetection(); { std::scoped_lock CodeLock(CTX.GetCodeInvalidationMutex()); CTX.MarkMonoBackpatcherBlock(BlockEntry); } InvalidateAlignedInterval(BlockEntry, FEXCore::Utils::FEX_PAGE_SIZE, false); } void InvalidationTracker::DisableSMCDetection() { std::unique_lock Lock(IntervalsLock); SMCDetectionDisabled = true; uint64_t Address = 0; // Reprotect all RWX intervals as writable FEXCore::IntervalList::QueryResult Query; do { Query = RWXIntervals.Query(Address); if (Query.Enclosed) { void* TmpAddress = reinterpret_cast(Address); SIZE_T TmpSize = static_cast(Query.Size); ULONG TmpProt; NtProtectVirtualMemory(NtCurrentProcess(), &TmpAddress, &TmpSize, GetUntrapProt(Address), &TmpProt); } Address += Query.Size; } while (Query.Size); } ULONG InvalidationTracker::GetTrapProt(uint64_t Address) const { if (DEPDisabled && DEPPromotedIntervals.Query(Address).Enclosed) { return PAGE_READONLY; } return PAGE_EXECUTE_READ; } ULONG InvalidationTracker::GetUntrapProt(uint64_t Address) const { if (DEPDisabled && DEPPromotedIntervals.Query(Address).Enclosed) { return PAGE_READWRITE; } return PAGE_EXECUTE_READWRITE; } void InvalidationTracker::InvalidateIntervalInternal(uint64_t Address, uint64_t Size) { std::scoped_lock CodeLock(CTX.GetCodeInvalidationMutex()); InvalidateIntervalInternalLocked(Address, Size); } void InvalidationTracker::InvalidateIntervalInternalLocked(uint64_t Address, uint64_t Size) { // NOTE: This assumes CodeInvalidationMutex is locked by the caller CTX.InvalidateCodeBuffersCodeRange(Address, Size); for (auto Thread : Threads) { CTX.InvalidateThreadCachedCodeRange(Thread.second, Address, Size); } } bool InvalidationTracker::ProtectRWXIntervalsInternal(uint64_t Address, uint64_t Size, bool ForWriteLocked) { const auto End = Address + Size; std::shared_lock Lock(IntervalsLock); if (SMCDetectionDisabled) { return false; } bool HitRWXInterval = false; do { const auto Query = RWXIntervals.Query(Address); if (Query.Enclosed) { if (!HitRWXInterval) { if (ForWriteLocked) { // If we are protecting as writable, then the entire range must be invalidated before any protections are // applied and the invalidation mutex must be locked throughout. // Do this lazily only when an RWX region is actually hit. // NOTE: This assumes CodeInvalidationMutex is locked by the caller InvalidateIntervalInternalLocked(Address, Size); } HitRWXInterval = true; } void* TmpAddress = reinterpret_cast(Address); SIZE_T TmpSize = static_cast(std::min(End, Address + Query.Size) - Address); ULONG TmpProt; NtProtectVirtualMemory(NtCurrentProcess(), &TmpAddress, &TmpSize, ForWriteLocked ? GetUntrapProt(Address) : GetTrapProt(Address), &TmpProt); } else if (!Query.Size) { // No more regions past `Address` in the interval list break; } Address += Query.Size; } while (Address < End); return HitRWXInterval; } } // namespace FEX::Windows