refactoring

This commit is contained in:
Boring
2022-10-11 07:47:31 +08:00
parent d1787b1dbe
commit e64ee25bac
15 changed files with 364 additions and 897 deletions
+55 -55
View File
@@ -154,7 +154,7 @@ DWORD NTAPI MmpUserThreadStart(LPVOID lpThreadParameter) {
//
// Allocate and replace ThreadLocalStoragePointer for new thread
//
EnterCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
EnterCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
record = PMMP_TLSP_RECORD(RtlAllocateHeap(RtlProcessHeap(), 0, sizeof(MMP_TLSP_RECORD)));
if (record) {
@@ -172,7 +172,7 @@ DWORD NTAPI MmpUserThreadStart(LPVOID lpThreadParameter) {
NtCurrentTeb()->ThreadLocalStoragePointer = record->TlspMmpBlock;
InsertTailList(&MmpGlobalDataPtr->MmpTls.MmpThreadLocalStoragePointer, &record->InMmpThreadLocalStoragePointer);
InsertTailList(&MmpGlobalDataPtr->MmpTls->MmpThreadLocalStoragePointer, &record->InMmpThreadLocalStoragePointer);
success = true;
}
else {
@@ -180,17 +180,17 @@ DWORD NTAPI MmpUserThreadStart(LPVOID lpThreadParameter) {
}
}
LeaveCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
LeaveCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
//
// Handle MemoryModule Tls data
//
if (success) {
RtlAcquireSRWLockShared(&MmpGlobalDataPtr->MmpTls.MmpTlsListLock);
RtlAcquireSRWLockShared(&MmpGlobalDataPtr->MmpTls->MmpTlsListLock);
auto ThreadLocalStoragePointer = (PVOID*)NtCurrentTeb()->ThreadLocalStoragePointer;
PLIST_ENTRY entry = MmpGlobalDataPtr->MmpTls.MmpTlsList.Flink;
while (entry != &MmpGlobalDataPtr->MmpTls.MmpTlsList) {
PLIST_ENTRY entry = MmpGlobalDataPtr->MmpTls->MmpTlsList.Flink;
while (entry != &MmpGlobalDataPtr->MmpTls->MmpTlsList) {
PTLS_ENTRY tls = CONTAINING_RECORD(entry, TLS_ENTRY, TlsEntryLinks);
auto len = tls->TlsDirectory.EndAddressOfRawData - tls->TlsDirectory.StartAddressOfRawData;
@@ -212,16 +212,16 @@ DWORD NTAPI MmpUserThreadStart(LPVOID lpThreadParameter) {
entry = entry->Flink;
}
RtlReleaseSRWLockShared(&MmpGlobalDataPtr->MmpTls.MmpTlsListLock);
RtlReleaseSRWLockShared(&MmpGlobalDataPtr->MmpTls->MmpTlsListLock);
}
if (!success) {
return ERROR_NOT_ENOUGH_MEMORY;
}
EnterCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
++MmpGlobalDataPtr->MmpTls.MmpActiveThreadCount;
LeaveCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
EnterCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
++MmpGlobalDataPtr->MmpTls->MmpActiveThreadCount;
LeaveCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
__skip_tls:
return Context.ThreadStartRoutine(Context.ThreadParameter);
@@ -257,7 +257,7 @@ NTSTATUS NTAPI HookNtCreateThread(
Context.Rdx = ULONG64(_Context);
#endif
status = MmpGlobalDataPtr->MmpTls.Hooks.OriginNtCreateThread(
status = MmpGlobalDataPtr->MmpTls->Hooks.OriginNtCreateThread(
ThreadHandle,
DesiredAccess,
ObjectAttributes,
@@ -294,7 +294,7 @@ NTSTATUS NTAPI HookNtCreateThreadEx(
Context->ThreadStartRoutine = PTHREAD_START_ROUTINE(StartRoutine);
Context->ThreadParameter = Argument;
NTSTATUS status = MmpGlobalDataPtr->MmpTls.Hooks.OriginNtCreateThreadEx(
NTSTATUS status = MmpGlobalDataPtr->MmpTls->Hooks.OriginNtCreateThreadEx(
ThreadHandle,
DesiredAccess,
ObjectAttributes,
@@ -322,10 +322,10 @@ VOID NTAPI HookLdrShutdownThread(VOID) {
//
// Find our tlsp record
//
EnterCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
EnterCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
entry = MmpGlobalDataPtr->MmpTls.MmpThreadLocalStoragePointer.Flink;
while (entry != &MmpGlobalDataPtr->MmpTls.MmpThreadLocalStoragePointer) {
entry = MmpGlobalDataPtr->MmpTls->MmpThreadLocalStoragePointer.Flink;
while (entry != &MmpGlobalDataPtr->MmpTls->MmpThreadLocalStoragePointer) {
auto p = CONTAINING_RECORD(entry, MMP_TLSP_RECORD, InMmpThreadLocalStoragePointer);
if (p->UniqueThread == NtCurrentThreadId()) {
@@ -344,19 +344,19 @@ VOID NTAPI HookLdrShutdownThread(VOID) {
entry = entry->Flink;
}
--MmpGlobalDataPtr->MmpTls.MmpActiveThreadCount;
--MmpGlobalDataPtr->MmpTls->MmpActiveThreadCount;
LeaveCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
LeaveCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
//
// Free MemoryModule Tls data
//
RtlAcquireSRWLockExclusive(&MmpGlobalDataPtr->MmpTls.MmpTlsListLock);
RtlAcquireSRWLockExclusive(&MmpGlobalDataPtr->MmpTls->MmpTlsListLock);
if (record) {
auto TlspMmpBlock = (PVOID*)record->TlspMmpBlock;
entry = MmpGlobalDataPtr->MmpTls.MmpTlsList.Flink;
while (entry != &MmpGlobalDataPtr->MmpTls.MmpTlsList) {
entry = MmpGlobalDataPtr->MmpTls->MmpTlsList.Flink;
while (entry != &MmpGlobalDataPtr->MmpTls->MmpTlsList) {
auto p = CONTAINING_RECORD(entry, TLS_ENTRY, TlsEntryLinks);
RtlFreeHeap(RtlProcessHeap(), 0, TlspMmpBlock[p->TlsDirectory.Characteristics]);
@@ -367,17 +367,17 @@ VOID NTAPI HookLdrShutdownThread(VOID) {
RtlFreeHeap(RtlProcessHeap(), 0, TlspMmpBlock);
}
else {
if (MmpGlobalDataPtr->MmpTls.MmpTlsList.Flink != &MmpGlobalDataPtr->MmpTls.MmpTlsList) {
if (MmpGlobalDataPtr->MmpTls->MmpTlsList.Flink != &MmpGlobalDataPtr->MmpTls->MmpTlsList) {
assert(false);
}
}
RtlReleaseSRWLockExclusive(&MmpGlobalDataPtr->MmpTls.MmpTlsListLock);
RtlReleaseSRWLockExclusive(&MmpGlobalDataPtr->MmpTls->MmpTlsListLock);
//
// Call the original function
//
MmpGlobalDataPtr->MmpTls.Hooks.OriginLdrShutdownThread();
MmpGlobalDataPtr->MmpTls->Hooks.OriginLdrShutdownThread();
}
BOOL NTAPI PreHookNtSetInformationProcess() {
@@ -428,7 +428,7 @@ BOOL NTAPI PreHookNtSetInformationProcess() {
);
if (NT_SUCCESS(status)) {
EnterCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
EnterCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
for (DWORD i = 0; i < CurrentThreadCount; ++i) {
auto const& LdrTls = ProcessTlsInformation->ThreadData[i];
auto const& MmpTls = tmpTlsInformation->ThreadData[i];
@@ -438,9 +438,9 @@ BOOL NTAPI PreHookNtSetInformationProcess() {
record->TlspLdrBlock = LdrTls.TlsVector;
record->TlspMmpBlock = MmpTls.TlsVector;
record->UniqueThread = LdrTls.ThreadId;
InsertTailList(&MmpGlobalDataPtr->MmpTls.MmpThreadLocalStoragePointer, &record->InMmpThreadLocalStoragePointer);
InsertTailList(&MmpGlobalDataPtr->MmpTls->MmpThreadLocalStoragePointer, &record->InMmpThreadLocalStoragePointer);
}
LeaveCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
LeaveCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
}
}
@@ -458,7 +458,7 @@ NTSTATUS NTAPI HookNtSetInformationProcess(
_In_ ULONG ProcessInformationLength) {
if (ProcessInformationClass != ProcessTlsInformation) {
return MmpGlobalDataPtr->MmpTls.Hooks.OriginNtSetInformationProcess(
return MmpGlobalDataPtr->MmpTls->Hooks.OriginNtSetInformationProcess(
ProcessHandle,
ProcessInformationClass,
ProcessInformation,
@@ -532,7 +532,7 @@ NTSTATUS NTAPI HookNtSetInformationProcess(
}
}
status = MmpGlobalDataPtr->MmpTls.Hooks.OriginNtSetInformationProcess(
status = MmpGlobalDataPtr->MmpTls->Hooks.OriginNtSetInformationProcess(
hProcess,
ProcessInformationClass,
Tls,
@@ -542,14 +542,14 @@ NTSTATUS NTAPI HookNtSetInformationProcess(
//
// Modify our mapping
//
EnterCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
EnterCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
for (ULONG i = 0; i < Tls->ThreadDataCount; ++i) {
bool found = false;
PLIST_ENTRY entry = MmpGlobalDataPtr->MmpTls.MmpThreadLocalStoragePointer.Flink;
PLIST_ENTRY entry = MmpGlobalDataPtr->MmpTls->MmpThreadLocalStoragePointer.Flink;
// Find thread-spec tlsp
while (entry != &MmpGlobalDataPtr->MmpTls.MmpThreadLocalStoragePointer) {
while (entry != &MmpGlobalDataPtr->MmpTls->MmpThreadLocalStoragePointer) {
PMMP_TLSP_RECORD j = CONTAINING_RECORD(entry, MMP_TLSP_RECORD, InMmpThreadLocalStoragePointer);
@@ -593,7 +593,7 @@ NTSTATUS NTAPI HookNtSetInformationProcess(
ProcessTlsInformation->ThreadData[i].ThreadId = Tls->ThreadData[i].ThreadId;
}
}
LeaveCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
LeaveCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
} while (false);
@@ -605,7 +605,7 @@ NTSTATUS NTAPI MmpAcquireTlsIndex(_Out_ PULONG TlsIndex) {
*TlsIndex = -1;
ULONG Index = RtlFindClearBitsAndSet(&MmpGlobalDataPtr->MmpTls.MmpTlsBitmap, 1, 0);
ULONG Index = RtlFindClearBitsAndSet(&MmpGlobalDataPtr->MmpTls->MmpTlsBitmap, 1, 0);
if (Index != -1) {
*TlsIndex = Index;
return STATUS_SUCCESS;
@@ -666,9 +666,9 @@ NTSTATUS NTAPI MmpAllocateTlsEntry(
Entry->TlsDirectory.Characteristics =
*PULONG(Entry->TlsDirectory.AddressOfIndex) = TlsIndex;
RtlAcquireSRWLockExclusive(&MmpGlobalDataPtr->MmpTls.MmpTlsListLock);
InsertTailList(&MmpGlobalDataPtr->MmpTls.MmpTlsList, &Entry->TlsEntryLinks);
RtlReleaseSRWLockExclusive(&MmpGlobalDataPtr->MmpTls.MmpTlsListLock);
RtlAcquireSRWLockExclusive(&MmpGlobalDataPtr->MmpTls->MmpTlsListLock);
InsertTailList(&MmpGlobalDataPtr->MmpTls->MmpTlsList, &Entry->TlsEntryLinks);
RtlReleaseSRWLockExclusive(&MmpGlobalDataPtr->MmpTls->MmpTlsListLock);
*lpTlsEntry = Entry;
*lpTlsIndex = TlsIndex;
@@ -677,20 +677,20 @@ NTSTATUS NTAPI MmpAllocateTlsEntry(
NTSTATUS NTAPI MmpReleaseTlsEntry(_In_ PLDR_DATA_TABLE_ENTRY lpModuleEntry) {
RtlAcquireSRWLockExclusive(&MmpGlobalDataPtr->MmpTls.MmpTlsListLock);
RtlAcquireSRWLockExclusive(&MmpGlobalDataPtr->MmpTls->MmpTlsListLock);
for (auto entry = MmpGlobalDataPtr->MmpTls.MmpTlsList.Flink; entry != &MmpGlobalDataPtr->MmpTls.MmpTlsList; entry = entry->Flink) {
for (auto entry = MmpGlobalDataPtr->MmpTls->MmpTlsList.Flink; entry != &MmpGlobalDataPtr->MmpTls->MmpTlsList; entry = entry->Flink) {
auto p = CONTAINING_RECORD(entry, TLS_ENTRY, TlsEntryLinks);
if (p->ModuleEntry == lpModuleEntry) {
RemoveEntryList(&p->TlsEntryLinks);
RtlClearBit(&MmpGlobalDataPtr->MmpTls.MmpTlsBitmap, p->TlsDirectory.Characteristics);
RtlClearBit(&MmpGlobalDataPtr->MmpTls->MmpTlsBitmap, p->TlsDirectory.Characteristics);
RtlFreeHeap(RtlProcessHeap(), 0, p);
break;
}
}
RtlReleaseSRWLockExclusive(&MmpGlobalDataPtr->MmpTls.MmpTlsListLock);
RtlReleaseSRWLockExclusive(&MmpGlobalDataPtr->MmpTls->MmpTlsListLock);
return STATUS_SUCCESS;
}
@@ -723,7 +723,7 @@ NTSTATUS NTAPI MmpHandleTlsData(_In_ PLDR_DATA_TABLE_ENTRY lpModuleEntry) {
return STATUS_INSUFFICIENT_RESOURCES;
}
auto ThreadCount = MmpGlobalDataPtr->MmpTls.MmpActiveThreadCount;
auto ThreadCount = MmpGlobalDataPtr->MmpTls->MmpActiveThreadCount;
auto success = true;
auto Length = sizeof(PROCESS_TLS_INFORMATION) + (ThreadCount - 1) * sizeof(THREAD_TLS_INFORMATION);
auto ProcessTlsInformation = PPROCESS_TLS_INFORMATION(RtlAllocateHeap(RtlProcessHeap(), HEAP_ZERO_MEMORY, Length));
@@ -791,26 +791,26 @@ BOOL NTAPI MmpTlsInitialize() {
//
// Capture thread count
//
MmpGlobalDataPtr->MmpTls.MmpActiveThreadCount = MmpGetThreadCount();
MmpGlobalDataPtr->MmpTls->MmpActiveThreadCount = MmpGetThreadCount();
//
// Initialize tlsp
//
InitializeCriticalSection(&MmpGlobalDataPtr->MmpTls.MmpTlspLock);
InitializeListHead(&MmpGlobalDataPtr->MmpTls.MmpThreadLocalStoragePointer);
InitializeCriticalSection(&MmpGlobalDataPtr->MmpTls->MmpTlspLock);
InitializeListHead(&MmpGlobalDataPtr->MmpTls->MmpThreadLocalStoragePointer);
//
// Initialize tls list
//
InitializeListHead(&MmpGlobalDataPtr->MmpTls.MmpTlsList);
RtlInitializeSRWLock(&MmpGlobalDataPtr->MmpTls.MmpTlsListLock);
InitializeListHead(&MmpGlobalDataPtr->MmpTls->MmpTlsList);
RtlInitializeSRWLock(&MmpGlobalDataPtr->MmpTls->MmpTlsListLock);
PULONG buffer = PULONG(RtlAllocateHeap(RtlProcessHeap(), HEAP_ZERO_MEMORY, MMP_TLSP_INDEX_BUFFER_SIZE));
if (!buffer) RtlRaiseStatus(STATUS_NO_MEMORY);
RtlFillMemory(buffer, MMP_START_TLS_INDEX / 8, -1);
RtlInitializeBitMap(&MmpGlobalDataPtr->MmpTls.MmpTlsBitmap, buffer, MMP_MAXIMUM_TLS_INDEX);
RtlInitializeBitMap(&MmpGlobalDataPtr->MmpTls->MmpTlsBitmap, buffer, MMP_MAXIMUM_TLS_INDEX);
if (NtCurrentTeb()->ThreadLocalStoragePointer) {
if (!PreHookNtSetInformationProcess()) {
@@ -822,17 +822,17 @@ BOOL NTAPI MmpTlsInitialize() {
// Hook functions
//
MmpGlobalDataPtr->MmpTls.Hooks.OriginNtCreateThread = NtCreateThread;
MmpGlobalDataPtr->MmpTls.Hooks.OriginNtCreateThreadEx = NtCreateThreadEx;
MmpGlobalDataPtr->MmpTls.Hooks.OriginLdrShutdownThread = LdrShutdownThread;
MmpGlobalDataPtr->MmpTls.Hooks.OriginNtSetInformationProcess = NtSetInformationProcess;
MmpGlobalDataPtr->MmpTls->Hooks.OriginNtCreateThread = NtCreateThread;
MmpGlobalDataPtr->MmpTls->Hooks.OriginNtCreateThreadEx = NtCreateThreadEx;
MmpGlobalDataPtr->MmpTls->Hooks.OriginLdrShutdownThread = LdrShutdownThread;
MmpGlobalDataPtr->MmpTls->Hooks.OriginNtSetInformationProcess = NtSetInformationProcess;
DetourTransactionBegin();
DetourUpdateThread(NtCurrentThread());
DetourAttach((PVOID*)&MmpGlobalDataPtr->MmpTls.Hooks.OriginNtCreateThread, HookNtCreateThread);
DetourAttach((PVOID*)&MmpGlobalDataPtr->MmpTls.Hooks.OriginNtCreateThreadEx, HookNtCreateThreadEx);
DetourAttach((PVOID*)&MmpGlobalDataPtr->MmpTls.Hooks.OriginLdrShutdownThread, HookLdrShutdownThread);
DetourAttach((PVOID*)&MmpGlobalDataPtr->MmpTls.Hooks.OriginNtSetInformationProcess, HookNtSetInformationProcess);
DetourAttach((PVOID*)&MmpGlobalDataPtr->MmpTls->Hooks.OriginNtCreateThread, HookNtCreateThread);
DetourAttach((PVOID*)&MmpGlobalDataPtr->MmpTls->Hooks.OriginNtCreateThreadEx, HookNtCreateThreadEx);
DetourAttach((PVOID*)&MmpGlobalDataPtr->MmpTls->Hooks.OriginLdrShutdownThread, HookLdrShutdownThread);
DetourAttach((PVOID*)&MmpGlobalDataPtr->MmpTls->Hooks.OriginNtSetInformationProcess, HookNtSetInformationProcess);
DetourTransactionCommit();
return TRUE;