ARM64EC: Install a custom call checker to bypass NTDLL function patches

Some programs will hook the NTDLL exports that FEX depends on, the
regular ARM64EC call checker will detect such patches and invoke the
JIT to run them, which leads to infinite recursion if those same
exports are used during code compilation. Fix this by resolving all
patchable FFSs to their native ARM implementations for all indirect
calls performed by FEX, skipping any x86 patches.
This commit is contained in:
Billy Laws committed 2024-08-09 11:48:18 +00:00
1 parent 85d1b573ef
commit 9d9bd750e2
2 files changed
+52 -2

No files matched your search

+20
View File
@@ -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
+32 -2
View File
@@ -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<typename T>
void WriteModuleRVA(HMODULE Module, LONG RVA, T Data) {
if (!RVA) {
return;
}
void* Address = reinterpret_cast<void*>(reinterpret_cast<uintptr_t>(Module) + RVA);
void* ProtAddress = Address;
SIZE_T ProtSize = sizeof(T);
ULONG Prot;
NtProtectVirtualMemory(NtCurrentProcess(), &ProtAddress, &ProtSize, PAGE_READWRITE, &Prot);
*reinterpret_cast<T*>(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<HMODULE>(&__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<IMAGE_ARM64EC_METADATA*>(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<uint8_t*>(X64ReturnInstr) = 0xc3;
FillNtDllLUTs();
PatchCallChecker();
const auto NtDll = GetModuleHandle("ntdll.dll");
const uintptr_t KiUserExceptionDispatcherFFS = reinterpret_cast<uintptr_t>(GetProcAddress(NtDll, "KiUserExceptionDispatcher"));
Exception::KiUserExceptionDispatcher = NtDllRedirectionLUT[KiUserExceptionDispatcherFFS - NtDllBase] + NtDllBase;