From 4df17abea7ecbef2bba8632a0d8af30dcf0a6ca8 Mon Sep 17 00:00:00 2001 From: Billy Laws Date: Tue, 12 Aug 2025 00:31:53 +0100 Subject: [PATCH] Windows: Avoid potential race between thread init and termination With the prior locking, we could try to terminate a thread while it was initializing and perhaps holding global locks leading to deadlock. --- Source/Windows/ARM64EC/Module.cpp | 24 ++++++++++++------------ Source/Windows/WOW64/Module.cpp | 26 +++++++++++++------------- 2 files changed, 25 insertions(+), 25 deletions(-) diff --git a/Source/Windows/ARM64EC/Module.cpp b/Source/Windows/ARM64EC/Module.cpp index ff502a0d1..efa9c147c 100644 --- a/Source/Windows/ARM64EC/Module.cpp +++ b/Source/Windows/ARM64EC/Module.cpp @@ -911,6 +911,7 @@ void BTCpu64NotifyMemoryDirty(void* Address, SIZE_T Size) { void BTCpu64NotifyReadFile(HANDLE Handle, void* Address, SIZE_T Size, BOOL After, NTSTATUS Status) {} NTSTATUS ThreadInit() { + std::scoped_lock Lock(ThreadCreationMutex); FEX::Windows::InitCRTThread(); static constexpr size_t EmulatorStackSize = 0x40000; const uint64_t EmulatorStack = reinterpret_cast(::VirtualAlloc(nullptr, EmulatorStackSize, MEM_COMMIT | MEM_RESERVE, PAGE_READWRITE)); @@ -945,7 +946,6 @@ NTSTATUS ThreadInit() { Exception::LoadStateFromECContext(Thread, CPUArea.ContextAmd64().AMD64_Context); { - std::scoped_lock Lock(ThreadCreationMutex); auto ThreadTID = GetCurrentThreadId(); Threads.emplace(ThreadTID, Thread); if (StatAllocHandler) { @@ -959,29 +959,29 @@ NTSTATUS ThreadInit() { } NTSTATUS ThreadTerm(HANDLE Thread, LONG ExitCode) { - const auto [Err, CPUArea] = GetThreadCPUArea(Thread); - if (Err) { - return Err; - } - auto* OldThreadState = CPUArea.ThreadState(); - CPUArea.ThreadState() = nullptr; - THREAD_BASIC_INFORMATION Info; - if (NTSTATUS Err = NtQueryInformationThread(Thread, ThreadBasicInformation, &Info, sizeof(Info), nullptr); Err) { + if (auto Err = NtQueryInformationThread(Thread, ThreadBasicInformation, &Info, sizeof(Info), nullptr); Err) { return Err; } const auto ThreadTID = reinterpret_cast(Info.ClientId.UniqueThread); + bool Self = ThreadTID == GetCurrentThreadId(); + + const auto [Err, CPUArea] = GetThreadCPUArea(Thread); + if (Err) { + return Err; + } + { std::scoped_lock Lock(ThreadCreationMutex); Threads.erase(ThreadTID); if (StatAllocHandler) { - StatAllocHandler->DeallocateSlot(OldThreadState->ThreadStats); + StatAllocHandler->DeallocateSlot(CPUArea.ThreadState()->ThreadStats); } } - FEX::Windows::CallRetStack::DestroyThread(OldThreadState); - CTX->DestroyThread(OldThreadState); + FEX::Windows::CallRetStack::DestroyThread(CPUArea.ThreadState()); + CTX->DestroyThread(CPUArea.ThreadState()); ::VirtualFree(reinterpret_cast(CPUArea.EmulatorStackLimit()), 0, MEM_RELEASE); if (ThreadTID == GetCurrentThreadId()) { FEX::Windows::DeinitCRTThread(); diff --git a/Source/Windows/WOW64/Module.cpp b/Source/Windows/WOW64/Module.cpp index 04464326a..b8429aa58 100644 --- a/Source/Windows/WOW64/Module.cpp +++ b/Source/Windows/WOW64/Module.cpp @@ -559,6 +559,7 @@ void BTCpuProcessInit() { void BTCpuProcessTerm(HANDLE Handle, BOOL After, ULONG Status) {} void BTCpuThreadInit() { + std::scoped_lock Lock(ThreadCreationMutex); FEX::Windows::InitCRTThread(); auto* Thread = CTX->CreateThread(0, 0); @@ -567,7 +568,6 @@ void BTCpuThreadInit() { GetTLS().ThreadState() = Thread; GetTLS().ControlWord().fetch_or(ControlBits::WOW_CPU_AREA_DIRTY, std::memory_order::relaxed); - std::scoped_lock Lock(ThreadCreationMutex); auto ThreadTID = GetCurrentThreadId(); Threads.emplace(ThreadTID, Thread); if (StatAllocHandler) { @@ -576,30 +576,30 @@ void BTCpuThreadInit() { } void BTCpuThreadTerm(HANDLE Thread, LONG ExitCode) { - const auto [Err, TLS] = GetThreadTLS(Thread); - if (Err) { - return; - } - - auto* ThreadState = TLS.ThreadState(); - THREAD_BASIC_INFORMATION Info; - if (NTSTATUS Err = NtQueryInformationThread(Thread, ThreadBasicInformation, &Info, sizeof(Info), nullptr); Err) { + if (auto Err = NtQueryInformationThread(Thread, ThreadBasicInformation, &Info, sizeof(Info), nullptr); Err) { return; } const auto ThreadTID = reinterpret_cast(Info.ClientId.UniqueThread); + bool Self = ThreadTID == GetCurrentThreadId(); + + auto [Err, TLS] = GetThreadTLS(Thread); + if (Err) { + return; + } + { std::scoped_lock Lock(ThreadCreationMutex); Threads.erase(ThreadTID); if (StatAllocHandler) { - StatAllocHandler->DeallocateSlot(ThreadState->ThreadStats); + StatAllocHandler->DeallocateSlot(TLS.ThreadState()->ThreadStats); } } - FEX::Windows::CallRetStack::DestroyThread(ThreadState); - CTX->DestroyThread(ThreadState); - if (ThreadTID == GetCurrentThreadId()) { + FEX::Windows::CallRetStack::DestroyThread(TLS.ThreadState()); + CTX->DestroyThread(TLS.ThreadState()); + if (Self) { FEX::Windows::DeinitCRTThread(); } }