diff --git a/MemoryModule/ImportTable.cpp b/MemoryModule/ImportTable.cpp new file mode 100644 index 0000000..81dece0 --- /dev/null +++ b/MemoryModule/ImportTable.cpp @@ -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); +} diff --git a/MemoryModule/ImportTable.h b/MemoryModule/ImportTable.h new file mode 100644 index 0000000..5ea803d --- /dev/null +++ b/MemoryModule/ImportTable.h @@ -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); diff --git a/MemoryModule/Initialize.cpp b/MemoryModule/Initialize.cpp index ee845af..86f4d0d 100644 --- a/MemoryModule/Initialize.cpp +++ b/MemoryModule/Initialize.cpp @@ -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; diff --git a/MemoryModule/Loader.cpp b/MemoryModule/Loader.cpp index 2a18bcd..a3f2a13 100644 --- a/MemoryModule/Loader.cpp +++ b/MemoryModule/Loader.cpp @@ -1,4 +1,5 @@ #include "stdafx.h" +#include 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; diff --git a/MemoryModule/MemoryModule.cpp b/MemoryModule/MemoryModule.cpp index 3333428..166d942 100644 --- a/MemoryModule/MemoryModule.cpp +++ b/MemoryModule/MemoryModule.cpp @@ -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; diff --git a/MemoryModule/MemoryModule.h b/MemoryModule/MemoryModule.h index 5d3b093..7a8d80c 100644 --- a/MemoryModule/MemoryModule.h +++ b/MemoryModule/MemoryModule.h @@ -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 diff --git a/MemoryModule/MemoryModule.vcxproj b/MemoryModule/MemoryModule.vcxproj index 22bee5d..b152c6d 100644 --- a/MemoryModule/MemoryModule.vcxproj +++ b/MemoryModule/MemoryModule.vcxproj @@ -45,6 +45,7 @@ + @@ -106,6 +107,7 @@ + diff --git a/MemoryModule/MemoryModule.vcxproj.filters b/MemoryModule/MemoryModule.vcxproj.filters index 4364a80..9413f3c 100644 --- a/MemoryModule/MemoryModule.vcxproj.filters +++ b/MemoryModule/MemoryModule.vcxproj.filters @@ -105,6 +105,9 @@ Source Files + + Source Files + @@ -254,6 +257,9 @@ Header Files + + Header Files + diff --git a/MemoryModule/MmpGlobalData.h b/MemoryModule/MmpGlobalData.h index 6bc58f8..35752ec 100644 --- a/MemoryModule/MmpGlobalData.h +++ b/MemoryModule/MmpGlobalData.h @@ -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; diff --git a/MemoryModule/stdafx.h b/MemoryModule/stdafx.h index a15eb40..90d4463 100644 --- a/MemoryModule/stdafx.h +++ b/MemoryModule/stdafx.h @@ -21,6 +21,9 @@ //memory module base support #include "MemoryModule.h" +//import table support +#include "ImportTable.h" + //LDR_DATA_TABLE_ENTRY #include "LdrEntry.h" diff --git a/test/test.cpp b/test/test.cpp index 3e18578..98d5b49 100644 --- a/test/test.cpp +++ b/test/test.cpp @@ -4,6 +4,31 @@ //PMMP_GLOBAL_DATA MmpGlobalDataPtr = *(PMMP_GLOBAL_DATA*)GetProcAddress(GetModuleHandleA("MemoryModule.dll"), "MmpGlobalDataPtr"); +static void DisplayStatus() { + printf( + "\ +MemoryModulePP [Version %d.%d%s]\n\n\t\ +MmpFeatures = %08X\n\n\t\ +LdrpModuleBaseAddressIndex = %p\n\t\ +NtdllLdrEntry = %p\n\t\ +RtlRbInsertNodeEx = %p\n\t\ +RtlRbRemoveNode = %p\n\n\t\ +LdrpInvertedFunctionTable = %p\n\n\t\ +LdrpHashTable = %p\n\n\ +", + MmpGlobalDataPtr->MajorVersion, + MEMORY_MODULE_GET_MINOR_VERSION(MmpGlobalDataPtr->MinorVersion), + MEMORY_MODULE_IS_PREVIEW(MmpGlobalDataPtr->MinorVersion) ? " Preview" : "", + MmpGlobalDataPtr->MmpFeatures, + MmpGlobalDataPtr->MmpBaseAddressIndex->LdrpModuleBaseAddressIndex, + MmpGlobalDataPtr->MmpBaseAddressIndex->NtdllLdrEntry, + MmpGlobalDataPtr->MmpBaseAddressIndex->_RtlRbInsertNodeEx, + MmpGlobalDataPtr->MmpBaseAddressIndex->_RtlRbRemoveNode, + MmpGlobalDataPtr->MmpInvertedFunctionTable->LdrpInvertedFunctionTable, + MmpGlobalDataPtr->MmpLdrEntry->LdrpHashTable + ); +} + static PVOID ReadDllFile(LPCSTR FileName) { LPVOID buffer; size_t size; @@ -21,21 +46,6 @@ static PVOID ReadDllFile(LPCSTR FileName) { return buffer; } -static void DisplayStatus() { - printf( - "MemoryModulePP [Version %d.%d]\n\n\tMmpFeatures = %08X\n\n\tLdrpModuleBaseAddressIndex = %p\n\tNtdllLdrEntry = %p\n\tRtlRbInsertNodeEx = %p\n\tRtlRbRemoveNode = %p\n\n\tLdrpInvertedFunctionTable = %p\n\n\tLdrpHashTable = %p\n\n", - MmpGlobalDataPtr->MajorVersion, - MmpGlobalDataPtr->MinorVersion, - MmpGlobalDataPtr->MmpFeatures, - MmpGlobalDataPtr->MmpBaseAddressIndex->LdrpModuleBaseAddressIndex, - MmpGlobalDataPtr->MmpBaseAddressIndex->NtdllLdrEntry, - MmpGlobalDataPtr->MmpBaseAddressIndex->_RtlRbInsertNodeEx, - MmpGlobalDataPtr->MmpBaseAddressIndex->_RtlRbRemoveNode, - MmpGlobalDataPtr->MmpInvertedFunctionTable->LdrpInvertedFunctionTable, - MmpGlobalDataPtr->MmpLdrEntry->LdrpHashTable - ); -} - PVOID ReadDllFile2(LPCSTR FileName) { CHAR path[MAX_PATH + 4]; DWORD len = GetModuleFileNameA(nullptr, path, sizeof(path)); @@ -52,91 +62,62 @@ PVOID ReadDllFile2(LPCSTR FileName) { return nullptr; } -int test() { - LPVOID buffer = ReadDllFile2("a.dll"); +#define LIBRARY_PATH "\\\\DESKTOP-1145141919810\\Debug\\" - HMEMORYMODULE m1 = nullptr, m2 = m1; - HMODULE hModule = nullptr; - FARPROC pfn = nullptr; - DWORD MemoryModuleFeatures = 0; +HMODULE WINAPI MyLoadLibrary(LPCSTR lpModuleName) { + HMODULE hModule; + PVOID buffer; - typedef int(*_exception)(int code); - _exception exception = nullptr; - HRSRC hRsrc; - DWORD SizeofRes; - HGLOBAL gRes; - char str[10]; - - LdrQuerySystemMemoryModuleFeatures(&MemoryModuleFeatures); - if (MemoryModuleFeatures != MEMORY_FEATURE_ALL) { - printf("not support all features on this version of windows.\n"); + if (0 == _stricmp(lpModuleName, "CTestClassLibrary1.dll")) { + buffer = ReadDllFile(LIBRARY_PATH"CTestClassLibrary1.dll"); } - - if (!NT_SUCCESS(LdrLoadDllMemoryExW(&m1, nullptr, 0, buffer, 0, L"kernel64", nullptr))) goto end; - LoadLibraryW(L"wininet.dll"); - if (!NT_SUCCESS(LdrLoadDllMemoryExW(&m2, nullptr, 0, buffer, 0, L"kernel128", nullptr))) goto end; - - //forward export - hModule = (HMODULE)m1; - pfn = (decltype(pfn))(GetProcAddress(hModule, "Socket")); //ws2_32.WSASocketW - pfn = (decltype(pfn))(GetProcAddress(hModule, "VerifyTruse")); //wintrust.WinVerifyTrust - hModule = (HMODULE)m2; - pfn = (decltype(pfn))(GetProcAddress(hModule, "Socket")); - pfn = (decltype(pfn))(GetProcAddress(hModule, "VerifyTruse")); - - //exception - hModule = (HMODULE)m1; - exception = (_exception)GetProcAddress(hModule, "exception"); - if (exception) { - for (int i = 0; i < 5; ++i)exception(i); + else if (0 == _stricmp(lpModuleName, "CTestClassLibrary2.dll")) { + buffer = ReadDllFile(LIBRARY_PATH"CTestClassLibrary2.dll"); } - - //tls - pfn = GetProcAddress(hModule, "thread"); - if (pfn && pfn()) { - printf("thread test failed.\n"); - } - - //resource - if (!LoadStringA(hModule, 101, str, 10)) { - printf("load string failed.\n"); + else if (0 == _stricmp(lpModuleName, "CTestClassLibrary1Dep.dll")) { + buffer = ReadDllFile(LIBRARY_PATH"CTestClassLibrary1Dep.dll"); } else { - printf("%s\n", str); - } - if (!(hRsrc = FindResourceA(hModule, MAKEINTRESOURCEA(102), "BINARY"))) { - printf("find binary resource failed.\n"); - } - else { - if ((SizeofRes = SizeofResource(hModule, hRsrc)) != 0x10) { - printf("invalid res size.\n"); - } - else { - if (!(gRes = LoadResource(hModule, hRsrc))) { - printf("load res failed.\n"); - } - else { - if (!LockResource(gRes))printf("lock res failed.\n"); - else { - printf("resource test success.\n"); - } - } - } + return nullptr; } -end: + hModule = LoadLibraryMemoryExA(buffer, 0, lpModuleName, nullptr, 0); delete[]buffer; - if (m1)LdrUnloadDllMemory(m1); - FreeLibrary(LoadLibraryW(L"wininet.dll")); - FreeLibrary(GetModuleHandleW(L"wininet.dll")); - if (m2)LdrUnloadDllMemory(m2); + return hModule; +} - return 0; +VOID TestImportTableResolver() { + + // + // Register the import table resolver. + // + HANDLE hResolver = MmRegisterImportTableResolver(MyLoadLibrary, FreeLibraryMemory); + + // + // |-> CTestClassLibrary1.dll -> CTestClassLibrary1Dep.dll + // CTestClient.dll -| + // |-> CTestClassLibrary2.dll + // + + PVOID Client = ReadDllFile2("CTestClient.dll"); + HMODULE hm = LoadLibraryMemoryEx(Client, 0, TEXT("CTestClient.dll"), nullptr, 0); + delete[]Client; + + if (hm) { + auto pfn = GetProcAddress(hm, "TestProc"); + if (pfn) { + pfn(); + } + + FreeLibraryMemory(hm); + } + + MmRemoveImportTableResolver(hResolver); } int main() { DisplayStatus(); - test(); + TestImportTableResolver(); return 0; }