Merge branch 'ImportTableResolver' into MmpTlsFiber

This commit is contained in:
Boring
2023-12-14 19:44:48 +08:00
11 changed files with 338 additions and 190 deletions
+178
View File
@@ -0,0 +1,178 @@
#include "stdafx.h"
typedef struct _MMP_IAT_HANDLE {
HMODULE hModule;
PMM_IAT_RESOLVER lpResolver;
}MMP_IAT_HANDLE, * PMMP_IAT_HANDLE;
HMODULE MmpLoadLibraryA(
_In_ LPCSTR lpModuleName,
_Out_ PMM_IAT_RESOLVER* lpModuleResolver) {
HMODULE hModule = nullptr;
PMM_IAT_RESOLVER resolver = nullptr;
EnterCriticalSection(&MmpGlobalDataPtr->MmpIat->MmpIatResolverListLock);
PLIST_ENTRY lpResolver = MmpGlobalDataPtr->MmpIat->MmpIatResolverList.Flink;
while (lpResolver != &MmpGlobalDataPtr->MmpIat->MmpIatResolverList) {
PMM_IAT_RESOLVER entry = CONTAINING_RECORD(lpResolver, MM_IAT_RESOLVER, MM_IAT_RESOLVER::InMmpIatResolverList);
hModule = entry->LoadLibraryProv(lpModuleName);
if (hModule) {
resolver = entry;
++entry->ReferenceCount;
break;
}
lpResolver = lpResolver->Flink;
}
LeaveCriticalSection(&MmpGlobalDataPtr->MmpIat->MmpIatResolverListLock);
*lpModuleResolver = resolver;
return hModule;
}
VOID MemoryFreeImportTable(_In_ PMEMORYMODULE hMemoryModule) {
EnterCriticalSection(&MmpGlobalDataPtr->MmpIat->MmpIatResolverListLock);
PMMP_IAT_HANDLE list = (PMMP_IAT_HANDLE)hMemoryModule->hModulesList;
for (DWORD i = 0; i < hMemoryModule->dwModulesCount; ++i) {
auto entry = list[i];
entry.lpResolver->FreeLibraryProv(entry.hModule);
--entry.lpResolver->ReferenceCount;
}
LeaveCriticalSection(&MmpGlobalDataPtr->MmpIat->MmpIatResolverListLock);
RtlFreeHeap(NtCurrentPeb()->ProcessHeap, 0, hMemoryModule->hModulesList);
hMemoryModule->hModulesList = nullptr;
hMemoryModule->dwModulesCount = 0;
}
NTSTATUS MemoryResolveImportTable(
_In_ LPBYTE base,
_In_ PIMAGE_NT_HEADERS lpNtHeaders,
_In_ PMEMORYMODULE hMemoryModule) {
NTSTATUS status = STATUS_SUCCESS;
PIMAGE_IMPORT_DESCRIPTOR importDesc = nullptr;
DWORD count = 0;
do {
__try {
PIMAGE_DATA_DIRECTORY dir = &lpNtHeaders->OptionalHeader.DataDirectory[IMAGE_DIRECTORY_ENTRY_IMPORT];
PIMAGE_IMPORT_DESCRIPTOR iat = nullptr;
if (dir && dir->Size) {
iat = importDesc = PIMAGE_IMPORT_DESCRIPTOR(lpNtHeaders->OptionalHeader.ImageBase + dir->VirtualAddress);
}
if (iat) {
while (iat->Name) {
++count;
++iat;
}
}
if (importDesc && count) {
PMMP_IAT_HANDLE handles = (PMMP_IAT_HANDLE)RtlAllocateHeap(NtCurrentPeb()->ProcessHeap, HEAP_ZERO_MEMORY, sizeof(MMP_IAT_HANDLE) * count);
hMemoryModule->hModulesList = handles;
if (!hMemoryModule->hModulesList) {
status = STATUS_NO_MEMORY;
break;
}
for (DWORD i = 0; i < count; ++i, ++importDesc) {
uintptr_t* thunkRef;
FARPROC* funcRef;
PMM_IAT_RESOLVER resolver;
HMODULE handle = MmpLoadLibraryA((LPCSTR)(base + importDesc->Name), &resolver);
if (!handle) {
status = STATUS_DLL_NOT_FOUND;
break;
}
handles[hMemoryModule->dwModulesCount].hModule = handle;
handles[hMemoryModule->dwModulesCount++].lpResolver = resolver;
thunkRef = (uintptr_t*)(base + (importDesc->OriginalFirstThunk ? importDesc->OriginalFirstThunk : importDesc->FirstThunk));
funcRef = (FARPROC*)(base + importDesc->FirstThunk);
while (*thunkRef) {
*funcRef = GetProcAddress(
handle,
IMAGE_SNAP_BY_ORDINAL(*thunkRef) ? (LPCSTR)IMAGE_ORDINAL(*thunkRef) : (LPCSTR)PIMAGE_IMPORT_BY_NAME(base + (*thunkRef))->Name
);
if (!*funcRef) {
status = STATUS_ENTRYPOINT_NOT_FOUND;
break;
}
++thunkRef;
++funcRef;
}
if (!NT_SUCCESS(status))break;
}
}
}
__except (EXCEPTION_EXECUTE_HANDLER) {
status = GetExceptionCode();
}
} while (false);
if (!NT_SUCCESS(status)) {
MemoryFreeImportTable(hMemoryModule);
}
return status;
}
HANDLE WINAPI MmRegisterImportTableResolver(
_In_ MM_IAT_RESOLVER_ENTRY LoadLibraryProv,
_In_ MM_IAT_FREE_ENTRY FreeLibraryProv) {
HANDLE heap = RtlProcessHeap();
PMM_IAT_RESOLVER resolver = (PMM_IAT_RESOLVER)RtlAllocateHeap(heap, 0, sizeof(MM_IAT_RESOLVER));
if (resolver) {
EnterCriticalSection(&MmpGlobalDataPtr->MmpIat->MmpIatResolverListLock);
resolver->ReferenceCount = 1;
resolver->LoadLibraryProv = LoadLibraryProv;
resolver->FreeLibraryProv = FreeLibraryProv;
InsertTailList(&MmpGlobalDataPtr->MmpIat->MmpIatResolverList, &resolver->InMmpIatResolverList);
LeaveCriticalSection(&MmpGlobalDataPtr->MmpIat->MmpIatResolverListLock);
}
return resolver;
}
_Success_(return)
BOOL WINAPI MmRemoveImportTableResolver(_In_ HANDLE hMmIatResolver) {
HANDLE heap = RtlProcessHeap();
if (hMmIatResolver == &MmpGlobalDataPtr->MmpIat->MmpIatResolverHead) {
return FALSE;
}
PMM_IAT_RESOLVER resolver = CONTAINING_RECORD(hMmIatResolver, MM_IAT_RESOLVER, MM_IAT_RESOLVER::InMmpIatResolverList);
EnterCriticalSection(&MmpGlobalDataPtr->MmpIat->MmpIatResolverListLock);
if (resolver->ReferenceCount > 1) {
LeaveCriticalSection(&MmpGlobalDataPtr->MmpIat->MmpIatResolverListLock);
return FALSE;
}
RemoveHeadList(&resolver->InMmpIatResolverList);
LeaveCriticalSection(&MmpGlobalDataPtr->MmpIat->MmpIatResolverListLock);
return RtlFreeHeap(heap, 0, hMmIatResolver);
}
+31
View File
@@ -0,0 +1,31 @@
#pragma once
typedef HMODULE(WINAPI* MM_IAT_RESOLVER_ENTRY)(LPCSTR lpModuleName);
typedef BOOL(WINAPI* MM_IAT_FREE_ENTRY)(HMODULE hModule);
typedef struct _MM_IAT_RESOLVER {
LIST_ENTRY InMmpIatResolverList;
MM_IAT_RESOLVER_ENTRY LoadLibraryProv;
MM_IAT_FREE_ENTRY FreeLibraryProv;
DWORD ReferenceCount;
}MM_IAT_RESOLVER, * PMM_IAT_RESOLVER;
VOID MemoryFreeImportTable(_In_ PMEMORYMODULE hMemoryModule);
NTSTATUS MemoryResolveImportTable(
_In_ LPBYTE base,
_In_ PIMAGE_NT_HEADERS lpNtHeaders,
_In_ PMEMORYMODULE hMemoryModule
);
HANDLE WINAPI MmRegisterImportTableResolver(
_In_ MM_IAT_RESOLVER_ENTRY LoadLibraryProv,
_In_ MM_IAT_FREE_ENTRY FreeLibraryProv
);
_Success_(return)
BOOL WINAPI MmRemoveImportTableResolver(_In_ HANDLE hMmIatResolver);
+17 -2
View File
@@ -5,6 +5,10 @@
PMMP_GLOBAL_DATA MmpGlobalDataPtr;
#if MEMORY_MODULE_IS_PREVIEW(MEMORY_MODULE_MINOR_VERSION)
#pragma message("WARNING: You are using a preview version of MemoryModulePP.")
#endif
PRTL_RB_TREE FindLdrpModuleBaseAddressIndex() {
PRTL_RB_TREE LdrpModuleBaseAddressIndex = nullptr;
PLDR_DATA_TABLE_ENTRY_WIN10 nt10 = decltype(nt10)(MmpGlobalDataPtr->MmpBaseAddressIndex->NtdllLdrEntry);
@@ -427,8 +431,10 @@ NTSTATUS InitializeLockHeld() {
status = MmpAllocateGlobalData();
if (!NT_SUCCESS(status)) {
if (status == STATUS_ALREADY_INITIALIZED) {
if ((MmpGlobalDataPtr->MajorVersion < MEMORY_MODULE_MAJOR_VERSION) ||
(MmpGlobalDataPtr->MajorVersion == MEMORY_MODULE_MAJOR_VERSION && MmpGlobalDataPtr->MinorVersion < MEMORY_MODULE_MINOR_VERSION)) {
if ((MmpGlobalDataPtr->MajorVersion != MEMORY_MODULE_MAJOR_VERSION) ||
MEMORY_MODULE_IS_PREVIEW(MmpGlobalDataPtr->MinorVersion) != MEMORY_MODULE_IS_PREVIEW(MEMORY_MODULE_MINOR_VERSION) ||
(MEMORY_MODULE_IS_PREVIEW(MEMORY_MODULE_MINOR_VERSION) ? MmpGlobalDataPtr->MinorVersion != MEMORY_MODULE_MINOR_VERSION :
MmpGlobalDataPtr->MinorVersion < MEMORY_MODULE_MINOR_VERSION)) {
status = STATUS_NOT_SUPPORTED;
}
else {
@@ -458,6 +464,7 @@ NTSTATUS InitializeLockHeld() {
MmpGlobalDataPtr->MmpTls = (PMMP_TLS_DATA)((LPBYTE)MmpGlobalDataPtr->MmpLdrEntry + sizeof(MMP_LDR_ENTRY_DATA));
MmpGlobalDataPtr->MmpDotNet = (PMMP_DOT_NET_DATA)((LPBYTE)MmpGlobalDataPtr->MmpTls + sizeof(MMP_TLS_DATA));
MmpGlobalDataPtr->MmpFunctions = (PMMP_FUNCTIONS)((LPBYTE)MmpGlobalDataPtr->MmpDotNet + sizeof(MMP_DOT_NET_DATA));
MmpGlobalDataPtr->MmpIat = (PMMP_IAT_DATA)((LPBYTE)MmpGlobalDataPtr->MmpFunctions + sizeof(MMP_FUNCTIONS));
PLDR_DATA_TABLE_ENTRY pNtdllEntry = RtlFindLdrTableEntryByBaseName(L"ntdll.dll");
MmpGlobalDataPtr->MmpBaseAddressIndex->NtdllLdrEntry = pNtdllEntry;
@@ -480,6 +487,14 @@ NTSTATUS InitializeLockHeld() {
MmpGlobalDataPtr->MmpFunctions->_MmpHandleTlsData = MmpHandleTlsData;
MmpGlobalDataPtr->MmpFunctions->_MmpReleaseTlsEntry = MmpReleaseTlsEntry;
InitializeCriticalSection(&MmpGlobalDataPtr->MmpIat->MmpIatResolverListLock);
InitializeListHead(&MmpGlobalDataPtr->MmpIat->MmpIatResolverList);
InitializeListHead(&MmpGlobalDataPtr->MmpIat->MmpIatResolverHead.InMmpIatResolverList);
MmpGlobalDataPtr->MmpIat->MmpIatResolverHead.LoadLibraryProv = LoadLibraryA;
MmpGlobalDataPtr->MmpIat->MmpIatResolverHead.FreeLibraryProv = FreeLibrary;
MmpGlobalDataPtr->MmpIat->MmpIatResolverHead.ReferenceCount = 1;
InsertTailList(&MmpGlobalDataPtr->MmpIat->MmpIatResolverList, &MmpGlobalDataPtr->MmpIat->MmpIatResolverHead.InMmpIatResolverList);
MmpTlsInitialize();
MmpGlobalDataPtr->MmpDotNet->Initialized = MmpGlobalDataPtr->MmpDotNet->PreHooked = FALSE;
+13 -3
View File
@@ -1,4 +1,5 @@
#include "stdafx.h"
#include <cmath>
NTSTATUS NTAPI LdrMapDllMemory(
_In_ HMEMORYMODULE ViewBase,
@@ -64,6 +65,7 @@ NTSTATUS NTAPI LdrLoadDllMemoryExW(
if (dwFlags & LOAD_FLAGS_USE_DLL_NAME && (!DllName || !DllFullName))return STATUS_INVALID_PARAMETER_3;
if (DllName) {
int length = (int)wcslen(DllName);
PLIST_ENTRY ListHead = &NtCurrentPeb()->Ldr->InLoadOrderModuleList, ListEntry = ListHead->Flink;
PIMAGE_NT_HEADERS h1 = RtlImageNtHeader(BufferAddress), h2 = nullptr;
if (!h1)return STATUS_INVALID_IMAGE_FORMAT;
@@ -74,11 +76,19 @@ NTSTATUS NTAPI LdrLoadDllMemoryExW(
/* Check if it's being unloaded */
if (!CurEntry->InMemoryOrderLinks.Flink) continue;
auto dist = (CurEntry->BaseDllName.Length / sizeof(wchar_t)) - length;
bool equal = false;
if (dist == 0 || dist == 4) {
equal = !_wcsnicmp(DllName, CurEntry->BaseDllName.Buffer, length);
}
else {
continue;
}
/* Check if name matches */
if (!_wcsnicmp(DllName, CurEntry->BaseDllName.Buffer, (CurEntry->BaseDllName.Length / sizeof(wchar_t)) - 4) ||
!_wcsnicmp(DllName, CurEntry->BaseDllName.Buffer, CurEntry->BaseDllName.Length / sizeof(wchar_t))) {
if (equal) {
/* Let's compare their headers */
if (!(h2 = RtlImageNtHeader(CurEntry->DllBase)))continue;
if (!(module = MapMemoryModuleHandle((HMEMORYMODULE)CurEntry->DllBase)))continue;
+1 -89
View File
@@ -85,86 +85,6 @@ NTSTATUS MmpInitializeStructure(DWORD ImageFileSize, LPCVOID ImageFileBuffer, PI
return STATUS_SUCCESS;
}
NTSTATUS MemoryResolveImportTable(
_In_ LPBYTE base,
_In_ PIMAGE_NT_HEADERS lpNtHeaders,
_In_ PMEMORYMODULE hMemoryModule) {
NTSTATUS status = STATUS_SUCCESS;
PIMAGE_IMPORT_DESCRIPTOR importDesc = nullptr;
DWORD count = 0;
do {
__try {
PIMAGE_DATA_DIRECTORY dir = GET_HEADER_DICTIONARY(lpNtHeaders, IMAGE_DIRECTORY_ENTRY_IMPORT);
PIMAGE_IMPORT_DESCRIPTOR iat = nullptr;
if (dir && dir->Size) {
iat = importDesc = PIMAGE_IMPORT_DESCRIPTOR(lpNtHeaders->OptionalHeader.ImageBase + dir->VirtualAddress);
}
if (iat) {
while (iat->Name) {
++count;
++iat;
}
}
if (importDesc && count) {
hMemoryModule->hModulesList = (HMODULE*)RtlAllocateHeap(NtCurrentPeb()->ProcessHeap, HEAP_ZERO_MEMORY, sizeof(HMODULE) * count);
if (!hMemoryModule->hModulesList) {
status = STATUS_NO_MEMORY;
break;
}
for (DWORD i = 0; i < count; ++i, ++importDesc) {
uintptr_t* thunkRef;
FARPROC* funcRef;
HMODULE handle = LoadLibraryA((LPCSTR)(base + importDesc->Name));
if (!handle) {
status = STATUS_DLL_NOT_FOUND;
break;
}
hMemoryModule->hModulesList[hMemoryModule->dwModulesCount++] = handle;
thunkRef = (uintptr_t*)(base + (importDesc->OriginalFirstThunk ? importDesc->OriginalFirstThunk : importDesc->FirstThunk));
funcRef = (FARPROC*)(base + importDesc->FirstThunk);
while (*thunkRef) {
*funcRef = GetProcAddress(
handle,
IMAGE_SNAP_BY_ORDINAL(*thunkRef) ? (LPCSTR)IMAGE_ORDINAL(*thunkRef) : (LPCSTR)PIMAGE_IMPORT_BY_NAME(base + (*thunkRef))->Name
);
if (!*funcRef) {
status = STATUS_ENTRYPOINT_NOT_FOUND;
break;
}
++thunkRef;
++funcRef;
}
if (!NT_SUCCESS(status))break;
}
}
}
__except (EXCEPTION_EXECUTE_HANDLER) {
status = GetExceptionCode();
}
} while (false);
if (!NT_SUCCESS(status)) {
for (DWORD i = 0; i < hMemoryModule->dwModulesCount; ++i)
FreeLibrary(hMemoryModule->hModulesList[i]);
RtlFreeHeap(NtCurrentPeb()->ProcessHeap, 0, hMemoryModule->hModulesList);
hMemoryModule->hModulesList = nullptr;
hMemoryModule->dwModulesCount = 0;
}
return status;
}
NTSTATUS MemorySetSectionProtection(
_In_ LPBYTE base,
_In_ PIMAGE_NT_HEADERS lpNtHeaders) {
@@ -404,15 +324,7 @@ BOOL MemoryFreeLibrary(HMEMORYMODULE mod) {
if (!module) return FALSE;
if (module->loadFromLdrLoadDllMemory && !module->underUnload)return FALSE;
if (module->hModulesList) {
for (DWORD i = 0; i < module->dwModulesCount; ++i) {
if (module->hModulesList[i]) {
FreeLibrary(module->hModulesList[i]);
}
}
RtlFreeHeap(NtCurrentPeb()->ProcessHeap, 0, module->hModulesList);
}
if (module->hModulesList)MemoryFreeImportTable(module);
if (module->codeBase) VirtualFree(mod, 0, MEM_RELEASE);
return TRUE;
+1 -7
View File
@@ -47,7 +47,7 @@ typedef struct _MEMORYMODULE {
LPBYTE codeBase; //codeBase == ImageBase
PVOID lpReserved;
HMODULE* hModulesList; //Import module handles
PVOID hModulesList; //Import module handles
DWORD dwModulesCount; //number of module handles
DWORD dwReferenceCount;
@@ -71,12 +71,6 @@ extern "C" {
_In_ DWORD size
);
NTSTATUS MemoryResolveImportTable(
_In_ LPBYTE base,
_In_ PIMAGE_NT_HEADERS lpNtHeaders,
_In_ PMEMORYMODULE hMemoryModule
);
NTSTATUS MemorySetSectionProtection(
_In_ LPBYTE base,
_In_ PIMAGE_NT_HEADERS lpNtHeaders
+2
View File
@@ -45,6 +45,7 @@
<ClCompile Include="..\3rdparty\Detours\disolx86.cpp" />
<ClCompile Include="..\3rdparty\Detours\image.cpp" />
<ClCompile Include="..\3rdparty\Detours\modules.cpp" />
<ClCompile Include="ImportTable.cpp" />
<ClCompile Include="Initialize.cpp" />
<ClCompile Include="LoadDllMemoryApi.cpp" />
<ClCompile Include="MemoryModule.cpp" />
@@ -106,6 +107,7 @@
<ClInclude Include="..\3rdparty\phnt\include\phnt_windows.h" />
<ClInclude Include="..\3rdparty\phnt\include\subprocesstag.h" />
<ClInclude Include="..\3rdparty\phnt\include\winsta.h" />
<ClInclude Include="ImportTable.h" />
<ClInclude Include="LoadDllMemoryApi.h" />
<ClInclude Include="LoaderPrivate.h" />
<ClInclude Include="MemoryModule.h" />
@@ -105,6 +105,9 @@
<ClCompile Include="MmpTlsFiber.cpp">
<Filter>Source Files</Filter>
</ClCompile>
<ClCompile Include="ImportTable.cpp">
<Filter>Source Files</Filter>
</ClCompile>
</ItemGroup>
<ItemGroup>
<ClInclude Include="MemoryModule.h">
@@ -254,6 +257,9 @@
<ClInclude Include="MmpTlsp.h">
<Filter>Header Files</Filter>
</ClInclude>
<ClInclude Include="ImportTable.h">
<Filter>Header Files</Filter>
</ClInclude>
</ItemGroup>
<ItemGroup>
<None Include="..\README.md">
+19 -3
View File
@@ -72,6 +72,15 @@ typedef struct _MMP_FUNCTIONS {
decltype(&MmpReleaseTlsEntry) _MmpReleaseTlsEntry;
}MMP_FUNCTIONS, * PMMP_FUNCTIONS;
//ImportTable.cpp
typedef struct _MMP_IAT_DATA {
LIST_ENTRY MmpIatResolverList;
CRITICAL_SECTION MmpIatResolverListLock;
MM_IAT_RESOLVER MmpIatResolverHead;
}MMP_IAT_DATA, * PMMP_IAT_DATA;
typedef enum class _WINDOWS_VERSION :BYTE {
null,
xp,
@@ -86,8 +95,12 @@ typedef enum class _WINDOWS_VERSION :BYTE {
invalid
}WINDOWS_VERSION;
#define MEMORY_MODULE_MAJOR_VERSION 1
#define MEMORY_MODULE_MINOR_VERSION 3
#define MEMORY_MODULE_MAKE_PREVIEW(MinorVersion) (0x8000|(MinorVersion))
#define MEMORY_MODULE_IS_PREVIEW(MinorVersion) (!!(0x8000&(MinorVersion)))
#define MEMORY_MODULE_GET_MINOR_VERSION(MinorVersion) (~0x8000&(MinorVersion))
#define MEMORY_MODULE_MAJOR_VERSION 2
#define MEMORY_MODULE_MINOR_VERSION MEMORY_MODULE_MAKE_PREVIEW(0)
typedef struct _MMP_GLOBAL_DATA {
@@ -122,6 +135,8 @@ typedef struct _MMP_GLOBAL_DATA {
PMMP_FUNCTIONS MmpFunctions;
PMMP_IAT_DATA MmpIat;
}MMP_GLOBAL_DATA, * PMMP_GLOBAL_DATA;
#define MMP_GLOBAL_DATA_SIZE (\
@@ -131,7 +146,8 @@ typedef struct _MMP_GLOBAL_DATA {
sizeof(MMP_LDR_ENTRY_DATA) + \
sizeof(MMP_TLS_DATA) + \
sizeof(MMP_DOT_NET_DATA) + \
sizeof(PMMP_FUNCTIONS)\
sizeof(MMP_FUNCTIONS) + \
sizeof(PMMP_IAT_DATA)\
)
extern PMMP_GLOBAL_DATA MmpGlobalDataPtr;
+3
View File
@@ -21,6 +21,9 @@
//memory module base support
#include "MemoryModule.h"
//import table support
#include "ImportTable.h"
//LDR_DATA_TABLE_ENTRY
#include "LdrEntry.h"