diff --git a/Source/Windows/ARM64EC/Module.S b/Source/Windows/ARM64EC/Module.S index e68d19730..01bef9a04 100644 --- a/Source/Windows/ARM64EC/Module.S +++ b/Source/Windows/ARM64EC/Module.S @@ -103,3 +103,23 @@ wine_syscall: direct_syscall: svc #0x43 ret + + // A replacement for the standard ARM64EC call checker that ignores any FFS patches and always redirects to a function's + // native implementation. As the only library FEX calls into is NTDLL, this is done using a LUT generated at init time. + // Expects the FFS address in x11, exit thunk address in x10 (unused) and it's own address in x9. Return address is in x11. +.global "CheckCall" +"CheckCall": + adrp x9, NtDllBase + ldr x9, [x9, #:lo12:NtDllBase] + subs x16, x11, x9 + b.lo end + adrp x17, NtDllRedirectionLUTSize + ldr x17, [x17, #:lo12:NtDllRedirectionLUTSize] + cmp x16, x17 + b.hi end + adrp x17, NtDllRedirectionLUT + ldr x17, [x17, #:lo12:NtDllRedirectionLUT] + ldr w11, [x17, x16, lsl #2] + add x11, x11, x9 +end: + ret diff --git a/Source/Windows/ARM64EC/Module.cpp b/Source/Windows/ARM64EC/Module.cpp index e9d58e6b1..e3edb6342 100644 --- a/Source/Windows/ARM64EC/Module.cpp +++ b/Source/Windows/ARM64EC/Module.cpp @@ -50,9 +50,12 @@ $end_info$ class ECSyscallHandler; extern "C" { -void* X64ReturnInstr; // See Module.S -extern void* ExitFunctionEC; +extern IMAGE_DOS_HEADER __ImageBase; // Provided by the linker +extern void* ExitFunctionEC; +extern void* CheckCall; + +void* X64ReturnInstr; // See Module.S uintptr_t NtDllBase; // Exports on ARM64EC point to x64 fast forward sequences to allow for redirecting to the JIT if functions are hotpatched. This LUT is from their addresses to the relative addresses of the native code exports. @@ -183,6 +186,32 @@ void FillNtDllLUTs() { } } +template +void WriteModuleRVA(HMODULE Module, LONG RVA, T Data) { + if (!RVA) { + return; + } + + void* Address = reinterpret_cast(reinterpret_cast(Module) + RVA); + void* ProtAddress = Address; + SIZE_T ProtSize = sizeof(T); + ULONG Prot; + NtProtectVirtualMemory(NtCurrentProcess(), &ProtAddress, &ProtSize, PAGE_READWRITE, &Prot); + *reinterpret_cast(Address) = Data; + NtProtectVirtualMemory(NtCurrentProcess(), &ProtAddress, &ProtSize, Prot, nullptr); +} + +void PatchCallChecker() { + // See the comment for CheckCall in Module.S for why this is necessary + const auto Module = reinterpret_cast(&__ImageBase); + ULONG Size; + const auto* LoadConfig = + reinterpret_cast<_IMAGE_LOAD_CONFIG_DIRECTORY64*>(RtlImageDirectoryEntryToData(Module, true, IMAGE_DIRECTORY_ENTRY_LOAD_CONFIG, &Size)); + const auto* CHPEMetadata = reinterpret_cast(LoadConfig->CHPEMetadataPointer); + WriteModuleRVA(Module, CHPEMetadata->__os_arm64x_dispatch_call, &CheckCall); + WriteModuleRVA(Module, CHPEMetadata->__os_arm64x_dispatch_icall, &CheckCall); + WriteModuleRVA(Module, CHPEMetadata->__os_arm64x_dispatch_icall_cfg, &CheckCall); +} } // namespace namespace Exception { @@ -519,6 +548,7 @@ NTSTATUS ProcessInit() { *reinterpret_cast(X64ReturnInstr) = 0xc3; FillNtDllLUTs(); + PatchCallChecker(); const auto NtDll = GetModuleHandle("ntdll.dll"); const uintptr_t KiUserExceptionDispatcherFFS = reinterpret_cast(GetProcAddress(NtDll, "KiUserExceptionDispatcher")); Exception::KiUserExceptionDispatcher = NtDllRedirectionLUT[KiUserExceptionDispatcherFFS - NtDllBase] + NtDllBase;