diff options
Diffstat (limited to 'plugins/rtti/microsoft.cpp')
| -rw-r--r-- | plugins/rtti/microsoft.cpp | 345 |
1 files changed, 118 insertions, 227 deletions
diff --git a/plugins/rtti/microsoft.cpp b/plugins/rtti/microsoft.cpp index 07a33d61..1c600c9d 100644 --- a/plugins/rtti/microsoft.cpp +++ b/plugins/rtti/microsoft.cpp @@ -1,13 +1,25 @@ -#include "rtti.h" +#include "microsoft.h" using namespace BinaryNinja; +using namespace BinaryNinja::RTTI; +using namespace BinaryNinja::RTTI::Microsoft; constexpr int COL_SIG_REV0 = 0; constexpr int COL_SIG_REV1 = 1; -constexpr int RTTI_CONFIDENCE = 100; - constexpr int BCD_HASPCHD = 0x40; +constexpr const char *TYPE_SOURCE_MICROSOFT = "rtti_microsoft"; + + +// This is used internally when processing a `CompleteObjectLocator`. +struct CompleteObjectLocatorInfo +{ + uint64_t classOffset = 0; + std::optional<std::string> baseClassName; + std::optional<uint64_t> baseVft; +}; + + ClassHierarchyDescriptor::ClassHierarchyDescriptor(BinaryView *view, uint64_t address) { BinaryReader reader = BinaryReader(view); @@ -113,7 +125,7 @@ std::optional<CompleteObjectLocator> ReadCompleteObjectorLocator(BinaryView *vie Ref<Type> GetPMDType(BinaryView *view) { - auto typeId = Type::GenerateAutoTypeId("msvc_rtti", QualifiedName("PMD")); + auto typeId = Type::GenerateAutoTypeId(TYPE_SOURCE_MICROSOFT, QualifiedName("PMD")); Ref<Type> typeCache = view->GetTypeById(typeId); if (typeCache == nullptr) @@ -137,7 +149,7 @@ Ref<Type> ClassHierarchyDescriptorType(BinaryView *view, BNPointerBaseType ptrBa Ref<Type> BaseClassDescriptorType(BinaryView *view, BNPointerBaseType ptrBaseTy) { - auto typeId = Type::GenerateAutoTypeId("msvc_rtti", QualifiedName("RTTIBaseClassDescriptor")); + auto typeId = Type::GenerateAutoTypeId(TYPE_SOURCE_MICROSOFT, QualifiedName("RTTIBaseClassDescriptor")); Ref<Type> typeCache = view->GetTypeById(typeId); if (typeCache == nullptr) @@ -192,7 +204,7 @@ Ref<Type> BaseClassArrayType(BinaryView *view, const uint64_t length, BNPointerB Ref<Type> ClassHierarchyDescriptorType(BinaryView *view, BNPointerBaseType ptrBaseTy) { - auto typeId = Type::GenerateAutoTypeId("msvc_rtti", QualifiedName("RTTIClassHierarchyDescriptor")); + auto typeId = Type::GenerateAutoTypeId(TYPE_SOURCE_MICROSOFT, QualifiedName("RTTIClassHierarchyDescriptor")); Ref<Type> typeCache = view->GetTypeById(typeId); if (typeCache == nullptr) @@ -227,7 +239,7 @@ Ref<Type> ClassHierarchyDescriptorType(BinaryView *view, BNPointerBaseType ptrBa Ref<Type> CompleteObjectLocator64Type(BinaryView *view) { - auto typeId = Type::GenerateAutoTypeId("msvc_rtti", QualifiedName("RTTICompleteObjectLocator64")); + auto typeId = Type::GenerateAutoTypeId(TYPE_SOURCE_MICROSOFT, QualifiedName("RTTICompleteObjectLocator64")); Ref<Type> typeCache = view->GetTypeById(typeId); if (typeCache == nullptr) @@ -270,7 +282,7 @@ Ref<Type> CompleteObjectLocator64Type(BinaryView *view) Ref<Type> CompleteObjectLocator32Type(BinaryView *view) { - auto typeId = Type::GenerateAutoTypeId("msvc_rtti", QualifiedName("RTTICompleteObjectLocator32")); + auto typeId = Type::GenerateAutoTypeId(TYPE_SOURCE_MICROSOFT, QualifiedName("RTTICompleteObjectLocator32")); Ref<Type> typeCache = view->GetTypeById(typeId); if (typeCache == nullptr) @@ -317,131 +329,72 @@ Ref<Type> TypeDescriptorType(BinaryView *view, uint64_t length) } -Ref<Metadata> ClassInfo::SerializedMetadata() +std::vector<BaseClassInfo> MicrosoftRTTIProcessor::ProcessClassHierarchyDescriptor(uint64_t address, CompleteObjectLocator &coLocator, const ClassInfo &classInfo) { - std::map<std::string, Ref<Metadata> > classInfoMeta; - classInfoMeta["className"] = new Metadata(className); - if (baseClassName.has_value()) - classInfoMeta["baseClassName"] = new Metadata(baseClassName.value()); - if (classOffset.has_value()) - classInfoMeta["classOffset"] = new Metadata(classOffset.value()); - if (vft.has_value()) - classInfoMeta["vft"] = vft->SerializedMetadata(); - // NOTE: We omit baseVft as it can be resolved manually and just bloats the size. - return new Metadata(classInfoMeta); -} - - -ClassInfo ClassInfo::DeserializedMetadata(const Ref<Metadata> &metadata) -{ - std::map<std::string, Ref<Metadata> > classInfoMeta = metadata->GetKeyValueStore(); - ClassInfo info = {classInfoMeta["className"]->GetString()}; - if (classInfoMeta.find("baseClassName") != classInfoMeta.end()) - info.baseClassName = classInfoMeta["baseClassName"]->GetString(); - if (classInfoMeta.find("classOffset") != classInfoMeta.end()) - info.classOffset = classInfoMeta["classOffset"]->GetUnsignedInteger(); - if (classInfoMeta.find("vft") != classInfoMeta.end()) - info.vft = VirtualFunctionTableInfo::DeserializedMetadata(classInfoMeta["vft"]); - return info; -} + auto startAddr = m_view->GetOriginalImageBase(); + auto resolveAddr = [&](const uint64_t relAddr) { + return coLocator.signature == COL_SIG_REV1 ? startAddr + relAddr : relAddr; + }; + auto ptrBaseTy = coLocator.signature ? RelativeToBinaryStartPointerBaseType : AbsolutePointerBaseType; -Ref<Metadata> VirtualFunctionTableInfo::SerializedMetadata() -{ - std::vector<Ref<Metadata> > funcsMeta; - funcsMeta.reserve(virtualFunctions.size()); - for (auto &vFunc: virtualFunctions) - funcsMeta.emplace_back(vFunc.SerializedMetadata()); - std::map<std::string, Ref<Metadata> > vftMeta; - vftMeta["address"] = new Metadata(address); - vftMeta["functions"] = new Metadata(funcsMeta); - return new Metadata(vftMeta); -} + auto classHierarchyDesc = ClassHierarchyDescriptor(m_view, address); + auto classHierarchyDescName = fmt::format("{}::`RTTI Class Hierarchy Descriptor'", classInfo.className); + m_view->DefineAutoSymbol(new Symbol{DataSymbol, classHierarchyDescName, address}); + m_view->DefineDataVariable(address, + Confidence(ClassHierarchyDescriptorType(m_view, ptrBaseTy), RTTI_CONFIDENCE)); + auto baseClassArrayAddr = resolveAddr(classHierarchyDesc.pBaseClassArray); + auto baseClassArray = BaseClassArray(m_view, baseClassArrayAddr, classHierarchyDesc.numBaseClasses); + auto baseClassArrayName = fmt::format("{}::`RTTI Base Class Array'", classInfo.className); + m_view->DefineAutoSymbol(new Symbol{DataSymbol, baseClassArrayName, baseClassArrayAddr}); + m_view->DefineDataVariable(baseClassArrayAddr, + Confidence(BaseClassArrayType(m_view, baseClassArray.length, ptrBaseTy), + RTTI_CONFIDENCE)); -VirtualFunctionTableInfo VirtualFunctionTableInfo::DeserializedMetadata(const Ref<Metadata> &metadata) -{ - std::map<std::string, Ref<Metadata> > vftMeta = metadata->GetKeyValueStore(); - VirtualFunctionTableInfo vftInfo = {vftMeta["address"]->GetUnsignedInteger()}; - if (vftMeta.find("functions") != vftMeta.end()) + std::vector<BaseClassInfo> baseClasses = {}; + for (auto pBaseClassDescAddr: baseClassArray.descriptors) { - for (auto &entry: vftMeta["functions"]->GetArray()) - vftInfo.virtualFunctions.emplace_back(VirtualFunctionInfo::DeserializedMetadata(entry)); - } - return vftInfo; -} - - -Ref<Metadata> VirtualFunctionInfo::SerializedMetadata() -{ - std::map<std::string, Ref<Metadata> > vFuncMeta; - vFuncMeta["address"] = new Metadata(funcAddr); - return new Metadata(vFuncMeta); -} - - -VirtualFunctionInfo VirtualFunctionInfo::DeserializedMetadata(const Ref<Metadata> &metadata) -{ - std::map<std::string, Ref<Metadata> > vFuncMeta = metadata->GetKeyValueStore(); - VirtualFunctionInfo vFuncInfo = {vFuncMeta["address"]->GetUnsignedInteger()}; - return vFuncInfo; -} + auto baseClassDescAddr = resolveAddr(pBaseClassDescAddr); + auto baseClassDesc = BaseClassDescriptor(m_view, baseClassDescAddr); + auto baseClassTypeDescAddr = resolveAddr(baseClassDesc.pTypeDescriptor); + auto baseClassTypeDesc = TypeDescriptor(m_view, baseClassTypeDescAddr); + auto baseClassName = DemangleNameMS(m_view, allowMangledClassNames, baseClassTypeDesc.name); + if (!baseClassName.has_value()) + { + m_logger->LogWarn("Skipping BaseClassDescriptor with mangled name %llx", baseClassTypeDescAddr); + continue; + } -Ref<Metadata> MicrosoftRTTIProcessor::SerializedMetadata() -{ - std::map<std::string, Ref<Metadata> > classesMeta; - for (auto &[coLocatorAddr, classInfo]: m_classInfo) - { - auto addrStr = std::to_string(coLocatorAddr); - classesMeta[addrStr] = classInfo.SerializedMetadata(); - } + BaseClassInfo baseClassInfo = {baseClassName.value(), (uint64_t)baseClassDesc.where_mdisp}; - std::map<std::string, Ref<Metadata> > msvcMeta; - msvcMeta["classes"] = new Metadata(classesMeta); - return new Metadata(msvcMeta); -} + auto baseClassDescName = fmt::format("{}::`RTTI Base Class Descriptor at ({},{},{},{})", baseClassInfo.className, + baseClassDesc.where_mdisp, baseClassDesc.where_pdisp, + baseClassDesc.where_vdisp, baseClassDesc.attributes); + m_view->DefineAutoSymbol(new Symbol{DataSymbol, baseClassDescName, baseClassDescAddr}); + m_view->DefineDataVariable(baseClassDescAddr, + Confidence(BaseClassDescriptorType(m_view, ptrBaseTy), RTTI_CONFIDENCE)); + auto baseClassTypeDescSymName = fmt::format("class {} `RTTI Type Descriptor'", baseClassInfo.className); + m_view->DefineAutoSymbol(new Symbol{DataSymbol, baseClassTypeDescSymName, baseClassTypeDescAddr}); + m_view->DefineDataVariable(baseClassTypeDescAddr, + Confidence(TypeDescriptorType(m_view, baseClassTypeDesc.name.length()), RTTI_CONFIDENCE)); -void MicrosoftRTTIProcessor::DeserializedMetadata(const Ref<Metadata> &metadata) -{ - std::map<std::string, Ref<Metadata> > msvcMeta = metadata->GetKeyValueStore(); - if (msvcMeta.find("classes") != msvcMeta.end()) - { - for (auto &[coLocatorAddrStr, classInfoMeta]: msvcMeta["classes"]->GetKeyValueStore()) + // If we are not dealing with are own class we should add it as a base class. + if (baseClassDesc.where_mdisp != 0 || baseClassInfo.className != classInfo.className) { - uint64_t coLocatorAddr = std::stoull(coLocatorAddrStr); - m_classInfo[coLocatorAddr] = ClassInfo::DeserializedMetadata(classInfoMeta); + if (baseClassDesc.attributes & BCD_HASPCHD) { + baseClasses.emplace_back(baseClassInfo); + } } } -} - -std::optional<std::string> MicrosoftRTTIProcessor::DemangleName(const std::string &mangledName) -{ - QualifiedName demangledName = {}; - Ref<Type> outType = {}; - if (!DemangleMS(m_view->GetDefaultArchitecture(), mangledName, outType, demangledName, true)) - { - // Try to use LLVM demangler. - if (!DemangleLLVM(mangledName, demangledName, true)) - return allowMangledClassNames ? std::optional(mangledName) : std::nullopt; - auto demangledNameStr = demangledName.GetString(); - size_t beginFind = demangledNameStr.find_first_of(' '); - if (beginFind != std::string::npos) - demangledNameStr.erase(0, beginFind + 1); - size_t endFind = demangledNameStr.find(" `RTTI Type Descriptor Name'"); - if (endFind != std::string::npos) - demangledNameStr.erase(endFind, demangledNameStr.length()); - return demangledNameStr; - } - return demangledName.GetString(); + return baseClasses; } - std::optional<ClassInfo> MicrosoftRTTIProcessor::ProcessRTTI(uint64_t coLocatorAddr) { - // Get complete object locator then check to see if its valid. auto coLocator = ReadCompleteObjectorLocator(m_view, coLocatorAddr); if (!coLocator.has_value()) return std::nullopt; @@ -451,72 +404,10 @@ std::optional<ClassInfo> MicrosoftRTTIProcessor::ProcessRTTI(uint64_t coLocatorA return coLocator->signature == COL_SIG_REV1 ? startAddr + relAddr : relAddr; }; - auto ptrBaseTy = coLocator->signature ? RelativeToBinaryStartPointerBaseType : AbsolutePointerBaseType; - - auto defineClassHierarchyDesc = [&](const uint64_t classHierarchyDescAddr, ClassInfo& classInfo, std::optional<CompleteObjectLocator> coLocator) { - auto classHierarchyDesc = ClassHierarchyDescriptor(m_view, classHierarchyDescAddr); - auto classHierarchyDescName = fmt::format("{}::`RTTI Class Hierarchy Descriptor'", classInfo.className); - m_view->DefineAutoSymbol(new Symbol{DataSymbol, classHierarchyDescName, classHierarchyDescAddr}); - m_view->DefineDataVariable(classHierarchyDescAddr, - Confidence(ClassHierarchyDescriptorType(m_view, ptrBaseTy), RTTI_CONFIDENCE)); - - auto baseClassArrayAddr = resolveAddr(classHierarchyDesc.pBaseClassArray); - auto baseClassArray = BaseClassArray(m_view, baseClassArrayAddr, classHierarchyDesc.numBaseClasses); - auto baseClassArrayName = fmt::format("{}::`RTTI Base Class Array'", classInfo.className); - m_view->DefineAutoSymbol(new Symbol{DataSymbol, baseClassArrayName, baseClassArrayAddr}); - m_view->DefineDataVariable(baseClassArrayAddr, - Confidence(BaseClassArrayType(m_view, baseClassArray.length, ptrBaseTy), - RTTI_CONFIDENCE)); - - std::map<uint64_t, ClassInfo> baseClasses = {}; - for (auto pBaseClassDescAddr: baseClassArray.descriptors) - { - auto baseClassDescAddr = resolveAddr(pBaseClassDescAddr); - auto baseClassDesc = BaseClassDescriptor(m_view, baseClassDescAddr); - - auto baseClassTypeDescAddr = resolveAddr(baseClassDesc.pTypeDescriptor); - auto baseClassTypeDesc = TypeDescriptor(m_view, baseClassTypeDescAddr); - auto baseClassName = DemangleName(baseClassTypeDesc.name); - if (!baseClassName.has_value()) - { - m_logger->LogWarn("Skipping BaseClassDescriptor with mangled name %llx", baseClassTypeDescAddr); - continue; - } - - // TODO: we probably want to maintain this state - auto baseClassInfo = ClassInfo{baseClassName.value()}; - - if (coLocator.has_value()) - { - if (baseClassDesc.where_mdisp == coLocator->offset && !classInfo.baseClassName.has_value() && classInfo.className != baseClassInfo.className) - classInfo.baseClassName = baseClassInfo.className; - } - - auto baseClassDescName = fmt::format("{}::`RTTI Base Class Descriptor at ({},{},{},{})", baseClassInfo.className, - baseClassDesc.where_mdisp, baseClassDesc.where_pdisp, - baseClassDesc.where_vdisp, baseClassDesc.attributes); - m_view->DefineAutoSymbol(new Symbol{DataSymbol, baseClassDescName, baseClassDescAddr}); - m_view->DefineDataVariable(baseClassDescAddr, - Confidence(BaseClassDescriptorType(m_view, ptrBaseTy), RTTI_CONFIDENCE)); - - auto baseClassTypeDescSymName = fmt::format("class {} `RTTI Type Descriptor'", baseClassInfo.className); - m_view->DefineAutoSymbol(new Symbol{DataSymbol, baseClassTypeDescSymName, baseClassTypeDescAddr}); - m_view->DefineDataVariable(baseClassTypeDescAddr, - Confidence(TypeDescriptorType(m_view, baseClassTypeDesc.name.length()), RTTI_CONFIDENCE)); - - if (baseClassDesc.attributes & BCD_HASPCHD) { - auto classHierarchyDescAddr = resolveAddr(baseClassDesc.pClassHierarchyDescriptor); - baseClasses[classHierarchyDescAddr] = baseClassInfo; - } - } - - return baseClasses; - }; - // Get type descriptor then check to see if the class name was demangled. auto typeDescAddr = resolveAddr(coLocator->pTypeDescriptor); auto typeDesc = TypeDescriptor(m_view, typeDescAddr); - auto className = DemangleName(typeDesc.name); + auto className = DemangleNameMS(m_view, allowMangledClassNames, typeDesc.name); if (!className.has_value()) return std::nullopt; @@ -528,39 +419,38 @@ std::optional<ClassInfo> MicrosoftRTTIProcessor::ProcessRTTI(uint64_t coLocatorA m_logger->LogDebug("Skipping CompleteObjectorLocator with anonymous name %llx", coLocatorAddr); return std::nullopt; } - className = fmt::format("ANONYMOUS_{:#x}", coLocatorAddr); + className = fmt::format("anonymous_{:#x}", coLocatorAddr); } - auto classInfo = ClassInfo{className.value()}; - if (coLocator->offset > 0) - classInfo.classOffset = coLocator->offset; - - auto typeDescSymName = fmt::format("class {} `RTTI Type Descriptor'", classInfo.className); - m_view->DefineAutoSymbol(new Symbol{DataSymbol, typeDescSymName, typeDescAddr}); - m_view->DefineDataVariable(typeDescAddr, - Confidence(TypeDescriptorType(m_view, typeDesc.name.length()), RTTI_CONFIDENCE)); + auto classInfo = ClassInfo{RTTIProcessorType::Microsoft, className.value()}; auto classHierarchyDescAddr = resolveAddr(coLocator->pClassHierarchyDescriptor); - auto baseClasses = defineClassHierarchyDesc(classHierarchyDescAddr, classInfo, coLocator); - m_visitedClassHierarchyDescAddrs.insert(classHierarchyDescAddr); - while (baseClasses.size() > 0) + classInfo.baseClasses = ProcessClassHierarchyDescriptor(classHierarchyDescAddr, coLocator.value(), classInfo); + + // Locate the current base class if we are in one. + std::optional<BaseClassInfo> currentBaseClass; + if (coLocator->offset >= 0) { - std::map<uint64_t, ClassInfo> newBaseClasses = {}; - for (auto& [baseClassHierarchyDescAddr, baseClassInfo] : baseClasses) + for (const auto &baseClassInfo: classInfo.baseClasses) { - if (m_visitedClassHierarchyDescAddrs.find(baseClassHierarchyDescAddr) != m_visitedClassHierarchyDescAddrs.end()) - continue; - - newBaseClasses.merge(defineClassHierarchyDesc(baseClassHierarchyDescAddr, baseClassInfo, std::nullopt)); - m_visitedClassHierarchyDescAddrs.insert(baseClassHierarchyDescAddr); + if (baseClassInfo.className != classInfo.className + || baseClassInfo.offset == coLocator->offset) + { + currentBaseClass = baseClassInfo; + break; + } } - - baseClasses = newBaseClasses; } + auto typeDescSymName = fmt::format("class {} `RTTI Type Descriptor'", classInfo.className); + m_view->DefineAutoSymbol(new Symbol{DataSymbol, typeDescSymName, typeDescAddr}); + m_view->DefineDataVariable(typeDescAddr, + Confidence(TypeDescriptorType(m_view, typeDesc.name.length()), RTTI_CONFIDENCE)); + auto coLocatorName = fmt::format("{}::`RTTI Complete Object Locator'", className.value()); - if (classInfo.baseClassName.has_value()) - coLocatorName += fmt::format("{{for `{}'}}", classInfo.baseClassName.value()); + if (currentBaseClass.has_value()) + coLocatorName += fmt::format("{{for `{}'}}", currentBaseClass->className); + m_view->DefineAutoSymbol(new Symbol{DataSymbol, coLocatorName, coLocatorAddr}); if (coLocator->signature == COL_SIG_REV1) m_view->DefineDataVariable(coLocatorAddr, Confidence(CompleteObjectLocator64Type(m_view), RTTI_CONFIDENCE)); @@ -571,7 +461,7 @@ std::optional<ClassInfo> MicrosoftRTTIProcessor::ProcessRTTI(uint64_t coLocatorA } -std::optional<VirtualFunctionTableInfo> MicrosoftRTTIProcessor::ProcessVFT(uint64_t vftAddr, const ClassInfo &classInfo) +std::optional<VirtualFunctionTableInfo> MicrosoftRTTIProcessor::ProcessVFT(uint64_t vftAddr, ClassInfo &classInfo, std::optional<BaseClassInfo> baseClassInfo) { VirtualFunctionTableInfo vftInfo = {vftAddr}; // Gather all virtual functions @@ -614,10 +504,10 @@ std::optional<VirtualFunctionTableInfo> MicrosoftRTTIProcessor::ProcessVFT(uint6 // Create virtual function table type auto vftTypeName = fmt::format("{}::VTable", classInfo.className); - if (classInfo.baseClassName.has_value()) + if (baseClassInfo.has_value()) { - vftTypeName = fmt::format("{}::{}", classInfo.baseClassName.value(), vftTypeName); // TODO: What is the correct form for the name? + vftTypeName = fmt::format("{}::{}", baseClassInfo->className, vftTypeName); } // TODO: Hack the debug type id is used here to allow the PDB type (debug info) to overwrite the RTTI vtable type. auto typeId = Type::GenerateAutoDebugTypeId(vftTypeName); @@ -634,17 +524,17 @@ std::optional<VirtualFunctionTableInfo> MicrosoftRTTIProcessor::ProcessVFT(uint6 auto vftSize = virtualFunctions.size() * addrSize; vftBuilder.SetWidth(vftSize); - if (auto baseVft = classInfo.baseVft) + if (baseClassInfo.has_value() && baseClassInfo->vft.has_value()) { - if (classInfo.baseVft->virtualFunctions.size() <= virtualFunctions.size()) + if (baseClassInfo->vft->virtualFunctions.size() <= virtualFunctions.size()) { // Adjust the current vFunc index to the end of the shared vFuncs. - vFuncIdx = classInfo.baseVft->virtualFunctions.size(); + vFuncIdx = baseClassInfo->vft->virtualFunctions.size(); virtualFunctions.erase(virtualFunctions.begin(), virtualFunctions.begin() + vFuncIdx); // We should set the vtable as a base class so that xrefs are propagated (among other things). // NOTE: this means that `this` params will be assumed pre-adjusted, this is normally fine assuming type propagation // NOTE: never occurs on the vft types. Other-wise we need to change this. - auto baseVftTypeName = fmt::format("{}::VTable", classInfo.baseClassName.value()); + auto baseVftTypeName = fmt::format("{}::VTable", baseClassInfo->className); NamedTypeReferenceBuilder baseVftNTR; baseVftNTR.SetName(baseVftTypeName); // Width is unresolved here so that we can keep non-base vfuncs un-inherited. @@ -676,8 +566,9 @@ std::optional<VirtualFunctionTableInfo> MicrosoftRTTIProcessor::ProcessVFT(uint6 // NOTE: The analyzed function type might not be available here. auto vFuncOffset = vFuncIdx * addrSize; // We have access to a backing function type, use it, otherwise void! + auto vFuncType = vFunc.has_value() ? vFunc.value()->GetType() : Type::VoidType(); vftBuilder.AddMemberAtOffset( - Type::PointerType(addrSize, vFunc.has_value() ? vFunc.value()->GetType() : Type::VoidType(), true), vFuncName, vFuncOffset); + Type::PointerType(addrSize, vFuncType, true), vFuncName, vFuncOffset); vFuncIdx++; } m_view->DefineType(typeId, vftTypeName, @@ -685,28 +576,28 @@ std::optional<VirtualFunctionTableInfo> MicrosoftRTTIProcessor::ProcessVFT(uint6 } auto vftName = fmt::format("{}::`vftable'", classInfo.className); - if (classInfo.baseClassName.has_value()) - vftName += fmt::format("{{for `{}'}}", classInfo.baseClassName.value()); + if (baseClassInfo.has_value()) + vftName += fmt::format("{{for `{}'}}", baseClassInfo->className); m_view->DefineAutoSymbol(new Symbol{DataSymbol, vftName, vftAddr}); m_view->DefineDataVariable(vftAddr, Confidence(Type::NamedType(m_view, vftTypeName), RTTI_CONFIDENCE)); return vftInfo; } -MicrosoftRTTIProcessor::MicrosoftRTTIProcessor(const Ref<BinaryView> &view, bool useMangled, bool checkRData, bool vftSweep, bool allowAnonymous) : m_view(view) +MicrosoftRTTIProcessor::MicrosoftRTTIProcessor(const Ref<BinaryView> &view, bool useMangled, bool checkRData, bool vftSweep, bool allowAnonymous) { + m_view = view; m_logger = new Logger("Microsoft RTTI"); allowMangledClassNames = useMangled; allowAnonymousClassNames = allowAnonymous; checkWritableRData = checkRData; m_classInfo = {}; - m_visitedClassHierarchyDescAddrs = {}; virtualFunctionTableSweep = vftSweep; - auto metadata = view->QueryMetadata(VIEW_METADATA_MSVC); + auto metadata = view->QueryMetadata(VIEW_METADATA_RTTI); if (metadata != nullptr) { // Load in metadata to the processor. - DeserializedMetadata(metadata); + DeserializedMetadata(RTTIProcessorType::Microsoft, metadata); } } @@ -765,12 +656,12 @@ void MicrosoftRTTIProcessor::ProcessRTTI() { if (segment->GetFlags() == (SegmentReadable | SegmentContainsData)) { - m_logger->LogDebug("Attempting to find VirtualFunctionTables in segment %llx", segment->GetStart()); + m_logger->LogDebug("Attempting to find RTTI in segment %llx", segment->GetStart()); scan(segment); } else if (checkWritableRData && rdataSection && rdataSection->GetStart() == segment->GetStart()) { - m_logger->LogDebug("Attempting to find VirtualFunctionTables in writable rdata segment %llx", + m_logger->LogDebug("Attempting to find RTTI in writable rdata segment %llx", segment->GetStart()); scan(segment); } @@ -833,12 +724,12 @@ void MicrosoftRTTIProcessor::ProcessVFT() } } - auto GetCachedVFTInfo = [&](uint64_t vftAddr, const ClassInfo& classInfo) { + auto GetCachedVFTInfo = [&](uint64_t vftAddr, ClassInfo& classInfo, const std::optional<BaseClassInfo>& baseClassInfo = std::nullopt) -> std::optional<VirtualFunctionTableInfo> { // Check in the cache so that we don't process vfts more than once. auto cachedVftInfo = vftFinishedMap.find(vftAddr); if (cachedVftInfo != vftFinishedMap.end()) return cachedVftInfo->second; - auto vftInfo = ProcessVFT(vftAddr, classInfo); + auto vftInfo = ProcessVFT(vftAddr, classInfo, baseClassInfo); vftFinishedMap[vftAddr] = vftInfo; return vftInfo; }; @@ -846,17 +737,17 @@ void MicrosoftRTTIProcessor::ProcessVFT() for (const auto &[coLocatorAddr, vftAddr]: vftMap) { auto classInfo = m_classInfo.find(coLocatorAddr)->second; - if (classInfo.baseClassName.has_value()) + for (auto& baseClassInfo : classInfo.baseClasses) { // Process base vtable and add it to the class info. - for (auto& [baseCoLocAddr, baseClassInfo] : m_classInfo) + for (auto& [baseCoLocAddr, classInfo] : m_classInfo) { - if (baseClassInfo.className == classInfo.baseClassName.value()) + if (classInfo.className == baseClassInfo.className) { uint64_t baseVftAddr = vftMap[baseCoLocAddr]; - if (auto baseVftInfo = GetCachedVFTInfo(baseVftAddr, baseClassInfo)) + if (auto baseVftInfo = GetCachedVFTInfo(baseVftAddr, classInfo, baseClassInfo)) { - classInfo.baseVft = baseVftInfo.value(); + baseClassInfo.vft = baseVftInfo.value(); break; } } @@ -864,12 +755,12 @@ void MicrosoftRTTIProcessor::ProcessVFT() } if (auto vftInfo = GetCachedVFTInfo(vftAddr, classInfo)) - { classInfo.vft = vftInfo.value(); - } + + m_classInfo[coLocatorAddr] = classInfo; } - + auto end_time = std::chrono::high_resolution_clock::now(); std::chrono::duration<double> elapsed_time = end_time - start_time; m_logger->LogInfo("ProcessVFT took %f seconds", elapsed_time.count()); -} +}
\ No newline at end of file |
