// SPDX-License-Identifier: MIT #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} {} 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; if (Prot & (PAGE_EXECUTE | PAGE_EXECUTE_READ | PAGE_EXECUTE_READWRITE)) { std::scoped_lock Lock(CTX.GetCodeInvalidationMutex()); for (auto Thread : Threads) { CTX.InvalidateGuestCodeRange(Thread.second, AlignedBase, AlignedSize); } } if (Prot & PAGE_EXECUTE_READWRITE) { LogMan::Msg::DFmt("Add SMC interval: {:X} - {:X}", AlignedBase, AlignedBase + AlignedSize); std::scoped_lock Lock(RWXIntervalsLock); RWXIntervals.Insert({AlignedBase, AlignedBase + AlignedSize}); } else { std::scoped_lock Lock(RWXIntervalsLock); RWXIntervals.Remove({AlignedBase, AlignedBase + AlignedSize}); } } void InvalidationTracker::InvalidateContainingSection(uint64_t Address, bool Free) { MEMORY_BASIC_INFORMATION Info; if (NtQueryVirtualMemory(NtCurrentProcess(), reinterpret_cast(Address), MemoryBasicInformation, &Info, sizeof(Info), nullptr)) { return; } 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; } { std::scoped_lock Lock(CTX.GetCodeInvalidationMutex()); for (auto Thread : Threads) { CTX.InvalidateGuestCodeRange(Thread.second, SectionBase, SectionSize); } } if (Free) { std::scoped_lock Lock(RWXIntervalsLock); RWXIntervals.Remove({SectionBase, SectionBase + SectionSize}); } } void InvalidationTracker::InvalidateAlignedInterval(uint64_t Address, uint64_t Size, bool Free) { 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; { std::scoped_lock Lock(CTX.GetCodeInvalidationMutex()); for (auto Thread : Threads) { CTX.InvalidateGuestCodeRange(Thread.second, AlignedBase, AlignedSize); } } if (Free) { std::scoped_lock Lock(RWXIntervalsLock); RWXIntervals.Remove({AlignedBase, AlignedBase + AlignedSize}); } } void InvalidationTracker::ReprotectRWXIntervals(uint64_t Address, uint64_t Size) { const auto End = Address + Size; std::scoped_lock Lock(RWXIntervalsLock); do { const auto Query = RWXIntervals.Query(Address); if (Query.Enclosed) { void* TmpAddress = reinterpret_cast(Address); SIZE_T TmpSize = static_cast(std::min(End, Address + Query.Size) - Address); ULONG TmpProt; NtProtectVirtualMemory(NtCurrentProcess(), &TmpAddress, &TmpSize, PAGE_EXECUTE_READ, &TmpProt); } else if (!Query.Size) { // No more regions past `Address` in the interval list break; } Address += Query.Size; } while (Address < End); } bool InvalidationTracker::HandleRWXAccessViolation(uint64_t FaultAddress) { const bool NeedsInvalidate = [&](uint64_t Address) { std::unique_lock Lock(RWXIntervalsLock); const bool Enclosed = RWXIntervals.Query(Address).Enclosed; // Invalidate just the single faulting page if (!Enclosed) { return false; } ULONG TmpProt; void* TmpAddress = reinterpret_cast(Address); SIZE_T TmpSize = 1; NtProtectVirtualMemory(NtCurrentProcess(), &TmpAddress, &TmpSize, PAGE_EXECUTE_READWRITE, &TmpProt); return true; }(FaultAddress); if (NeedsInvalidate) { // RWXIntervalsLock cannot be held during invalidation std::scoped_lock Lock(CTX.GetCodeInvalidationMutex()); for (auto Thread : Threads) { CTX.InvalidateGuestCodeRange(Thread.second, FaultAddress & FEXCore::Utils::FEX_PAGE_MASK, FEXCore::Utils::FEX_PAGE_SIZE); } return true; } return false; } } // namespace FEX::Windows