diff options
| author | Mason Reed <mason@vector35.com> | 2025-03-01 00:34:12 -0500 |
|---|---|---|
| committer | Mason Reed <mason@vector35.com> | 2025-03-19 21:17:34 -0400 |
| commit | f719d1d5565f7b7b1237b8e183c2985f006fa3e4 (patch) | |
| tree | 0fbb2c1af1254add49049abdec3f935669d658de /plugins/rtti | |
| parent | 93e92844d77b72f07bd7563211d60dc3595cc2e0 (diff) | |
Refactor and fixup MSVC and Itanium RTTI
Bunch of misc fixes and performance improvements
Diffstat (limited to 'plugins/rtti')
| -rw-r--r-- | plugins/rtti/itanium.cpp | 340 | ||||
| -rw-r--r-- | plugins/rtti/itanium.h | 6 | ||||
| -rw-r--r-- | plugins/rtti/microsoft.cpp | 345 | ||||
| -rw-r--r-- | plugins/rtti/microsoft.h | 24 | ||||
| -rw-r--r-- | plugins/rtti/plugin.cpp | 10 | ||||
| -rw-r--r-- | plugins/rtti/rtti.cpp | 102 | ||||
| -rw-r--r-- | plugins/rtti/rtti.h | 26 |
7 files changed, 489 insertions, 364 deletions
diff --git a/plugins/rtti/itanium.cpp b/plugins/rtti/itanium.cpp index e1af724b..0714d17d 100644 --- a/plugins/rtti/itanium.cpp +++ b/plugins/rtti/itanium.cpp @@ -16,6 +16,8 @@ TypeInfo::TypeInfo(BinaryView *view, uint64_t address) reader.Seek(address); base = reader.ReadPointer(); auto typeNameAddr = reader.ReadPointer(); + if (!view->IsValidOffset(typeNameAddr)) + return; reader.Seek(typeNameAddr); type_name = reader.ReadCString(512); } @@ -36,8 +38,7 @@ BaseClassTypeInfo::BaseClassTypeInfo(BinaryView *view, uint64_t address) reader.Seek(address); base_type = reader.ReadPointer(); offset_flags = reader.Read32(); - // TODO: Test this... - offset_flags_masks = static_cast<OffsetFlagsMasks>(reader.Read32()); + offset_flags_masks = reader.Read32(); } @@ -49,12 +50,11 @@ VMIClassTypeInfo::VMIClassTypeInfo(BinaryView *view, uint64_t address) : ClassTy flags = reader.Read32(); base_count = reader.Read32(); base_info = {}; - for (size_t i = 1; i < base_count; i++) + for (size_t i = 0; i < base_count; i++) { - // TODO: Verify this is correct. uint64_t currentBaseAddr = reader.GetOffset(); - base_info.emplace_back(view, reader.GetOffset()); - reader.Seek(currentBaseAddr + 12); + base_info.emplace_back(view, currentBaseAddr); + reader.Seek(currentBaseAddr + 0x10); } } @@ -183,13 +183,37 @@ Ref<Type> BaseClassTypeInfoType(BinaryView *view) } -Ref<Type> VMIClassTypeInfoType(BinaryView *view, int baseCount) +Ref<Type> VMIFlagsMasksType(BinaryView *view) +{ + auto typeId = Type::GenerateAutoTypeId(TYPE_SOURCE_ITANIUM, QualifiedName("VMIFlagsMasks")); + Ref<Type> typeCache = view->GetTypeById(typeId); + + if (typeCache == nullptr) + { + Ref<Architecture> arch = view->GetDefaultArchitecture(); + Ref<Type> uintType = Type::IntegerType(4, false); + + EnumerationBuilder enumerationBuilder; + enumerationBuilder.AddMemberWithValue("__non_diamond_repeat_mask", 0x1); + enumerationBuilder.AddMemberWithValue("__diamond_shaped_mask", 0x2); + + Ref<Type> enumerationType = TypeBuilder::EnumerationType(arch, enumerationBuilder.Finalize()).Finalize(); + view->DefineType(typeId, QualifiedName("__cxxabiv1::__flags_masks"), enumerationType); + + typeCache = view->GetTypeById(typeId); + } + + return typeCache; +} + + +Ref<Type> VMIClassTypeInfoType(BinaryView *view, uint64_t baseCount) { Ref<Architecture> arch = view->GetDefaultArchitecture(); Ref<Type> uintType = Type::IntegerType(4, false); StructureBuilder structureBuilder; - structureBuilder.AddMemberAtOffset(uintType, "__flags", 0x10); + structureBuilder.AddMemberAtOffset(VMIFlagsMasksType(view), "__flags", 0x10); structureBuilder.AddMemberAtOffset(uintType, "__base_count", 0x14); Ref<Type> baseInfoType = Type::ArrayType(BaseClassTypeInfoType(view), baseCount); structureBuilder.AddMemberAtOffset(baseInfoType, "__base_info", 0x18); @@ -238,9 +262,45 @@ std::optional<TypeInfoVariant> ReadTypeInfoVariant(BinaryView *view, uint64_t ob } +std::optional<BaseClassInfo> ItaniumRTTIProcessor::ProcessVFTBaseClassInfo(uint64_t vftAddr, ClassInfo &classInfo) +{ + BinaryReader reader = BinaryReader(m_view); + // Because we have this we _need_ to have the adjustment stuff. + // NOTE: We assume two 0x4 ints with the first being what we want. + reader.Seek(vftAddr - 0x10); + + auto adjustmentOffset = static_cast<int32_t>(reader.Read32()); + auto baseIdx = static_cast<int32_t>(reader.Read32()); + uint64_t classOffset = std::abs(adjustmentOffset); + + std::optional<BaseClassInfo> selectedBaseClassInfo = std::nullopt; + // Assuming we do not have a baseClassInfo already passed we can deduce it here. + for (auto& baseClass : classInfo.baseClasses) + { + // if (baseClass.offset == 0) + // { + // // If the base class is at offset 0 that means it has yet to be adjusted. + // // NOTE: This should only happen for `TIVSIClass`. If this assigns more than + // // one base class to this offset we are screwed. + // baseClass.offset = classOffset; + // LogInfo("Adjusting base class offset for %llx to %llx", vftAddr, classOffset); + // } + + if (baseClass.offset == classOffset) + { + // Found the appropriate base class for this vtable. + selectedBaseClassInfo = baseClass; + } + } + + // Return the selected base class for use in later processing such as `ProcessVFT`. + return selectedBaseClassInfo; +} + + std::optional<ClassInfo> ItaniumRTTIProcessor::ProcessRTTI(uint64_t objectAddr) { - // TODO: You cant get subobject offsets from rtti, its stored above this ptr in vtable. + // TODO: You cant get sub-object offsets from rtti, its stored above this ptr in vtable. // Get object as type info then check to see if it's valid. auto typeInfoVariant = ReadTypeInfoVariant(m_view, objectAddr); if (!typeInfoVariant.has_value()) @@ -258,35 +318,89 @@ std::optional<ClassInfo> ItaniumRTTIProcessor::ProcessRTTI(uint64_t objectAddr) m_view->UndefineAutoSymbol(typeInfoSymbol); m_view->DefineAutoSymbol(new Symbol{DataSymbol, typeInfoName, objectAddr}); + auto nameFromTypeInfoSymbol = [&](uint64_t addr) -> std::optional<std::string> { + auto sym = m_view->GetSymbolByAddress(addr); + if (sym == nullptr || sym->GetType() != ExternalSymbol) + return std::nullopt; + auto symName = sym->GetShortName(); + // Remove type info prefix. + if (symName.rfind("_typeinfo_for_", 0) != 0) + return std::nullopt; + return symName.substr(14); + }; + if (typeInfoVariant == TIVSIClass) { // Read the base class. auto siClassTypeInfo = SIClassTypeInfo(m_view, objectAddr); auto subTypeInfoVariant = ReadTypeInfoVariant(m_view, siClassTypeInfo.base_type); + std::string subTypeName; if (!subTypeInfoVariant.has_value()) - return std::nullopt; - auto subTypeInfo = TypeInfo(m_view, siClassTypeInfo.base_type); + { + // Allow externals to be used in place of a backed subtype. + // TODO: We should probably warn that vtables will likely be inaccurate. + // TODO: Because we wont know what offsets are valid. + auto externTypeName = nameFromTypeInfoSymbol(siClassTypeInfo.base_type); + if (!externTypeName.has_value()) + return std::nullopt; + m_logger->LogDebug("Non-backed external subtype for %llx", objectAddr); + subTypeName = externTypeName.value(); + } + else + { + auto subTypeInfo = TypeInfo(m_view, siClassTypeInfo.base_type); + subTypeName = subTypeInfo.type_name; + } // Demangle base class name and set - auto baseClassName = DemangleNameItanium(m_view, allowMangledClassNames, subTypeInfo.type_name); + auto baseClassName = DemangleNameItanium(m_view, allowMangledClassNames, subTypeName); if (!baseClassName.has_value()) { m_logger->LogWarn("Skipping base class with mangled name %llx", siClassTypeInfo.base_type); return std::nullopt; } - classInfo.baseClassName = baseClassName; // NOTE: The base class offset is not able to be resolved here. // NOTE: To resolve the base class offset you must go to the vtable. + uint64_t baseClassOffset = 0; + auto subBaseClassInfo = BaseClassInfo {baseClassName.value(), baseClassOffset}; + classInfo.baseClasses.emplace_back(subBaseClassInfo); m_view->DefineDataVariable(objectAddr, Confidence(SIClassTypeInfoType(m_view), 255)); } else if (typeInfoVariant == TIVVMIClass) { - // TODO: Read multiple base classes. auto vmiClassTypeInfo = VMIClassTypeInfo(m_view, objectAddr); m_view->DefineDataVariable(objectAddr, Confidence(VMIClassTypeInfoType(m_view, vmiClassTypeInfo.base_count), 255)); + for (const auto& baseInfo : vmiClassTypeInfo.base_info) + { + // Remove the flags and just get the offset + auto baseTypeInfoVariant = ReadTypeInfoVariant(m_view, baseInfo.base_type); + std::string subTypeName; + if (!baseTypeInfoVariant.has_value()) + { + // Allow externals to be used in place of a backed base type. + auto externTypeName = nameFromTypeInfoSymbol(baseInfo.base_type); + if (!externTypeName.has_value()) + return std::nullopt; + m_logger->LogDebug("Non-backed external subtype for %llx", objectAddr); + subTypeName = externTypeName.value(); + } else + { + auto baseTypeInfo = TypeInfo(m_view, baseInfo.base_type); + subTypeName = baseTypeInfo.type_name; + } + auto baseClassName = DemangleNameItanium(m_view, allowMangledClassNames, subTypeName); + if (!baseClassName.has_value()) + { + m_logger->LogWarn("Skipping base class with mangled name %llx", baseInfo.base_type); + continue; + } + // Shift off the flag bits. + uint64_t offset = baseInfo.offset_flags >> 8; + auto baseClassInfo = BaseClassInfo {baseClassName.value(), offset}; + classInfo.baseClasses.emplace_back(baseClassInfo); + } } else { - // auto classTypeInfo = ClassTypeInfo(m_view, objectAddr); m_view->DefineDataVariable(objectAddr, Confidence(ClassTypeInfoType(m_view), 255)); } @@ -294,13 +408,13 @@ std::optional<ClassInfo> ItaniumRTTIProcessor::ProcessRTTI(uint64_t objectAddr) } -std::optional<VirtualFunctionTableInfo> ItaniumRTTIProcessor::ProcessVFT(uint64_t vftAddr, ClassInfo &classInfo) +std::optional<VirtualFunctionTableInfo> ItaniumRTTIProcessor::ProcessVFT(uint64_t vftAddr, ClassInfo &classInfo, std::optional<BaseClassInfo> baseClassInfo) { VirtualFunctionTableInfo vftInfo = {vftAddr}; BinaryReader reader = BinaryReader(m_view); reader.Seek(vftAddr); // Gather all virtual functions - std::vector<Ref<Function> > virtualFunctions = {}; + std::vector<VirtualFunctionInfo> virtualFunctions = {}; while (true) { uint64_t vFuncAddr = reader.ReadPointer(); @@ -310,17 +424,30 @@ std::optional<VirtualFunctionTableInfo> ItaniumRTTIProcessor::ProcessVFT(uint64_ Ref<Segment> segment = m_view->GetSegmentAt(vFuncAddr); if (segment == nullptr || !(segment->GetFlags() & (SegmentExecutable | SegmentDenyWrite))) { - // Last CompleteObjectLocator or hit the next CompleteObjectLocator - break; + // TODO: Sometimes vFunc idx will be zeroed iirc. + // We allow vfuncs to point to extern functions. + auto vFuncSym = m_view->GetSymbolByAddress(vFuncAddr); + if (!vFuncSym) + break; + DataVariable dv; + bool foundDv = m_view->GetDataVariableAtAddress(vFuncAddr, dv); + // Last virtual function, or hit the next vtable. + if (!foundDv || !dv.type->m_object) + break; + // Void externs are very likely to be a func. + // TODO: Add some sanity checks for this! + if (!dv.type->IsFunction() && !(dv.type->IsVoid() && vFuncSym->GetType() == ExternalSymbol)) + break; + } + else + { + // TODO: Is likely a function check here? + m_logger->LogDebug("Discovered function from virtual function table... %llx", vFuncAddr); + m_view->AddFunctionForAnalysis(m_view->GetDefaultPlatform(), vFuncAddr, true); } - // TODO: Sometimes vFunc idx will be zeroed. - // TODO: Is likely a function check here? - m_logger->LogDebug("Discovered function from virtual function table... %llx", vFuncAddr); - auto vFunc = m_view->AddFunctionForAnalysis(m_view->GetDefaultPlatform(), vFuncAddr, true); - funcs.emplace_back(vFunc); } // Only ever add one function. - virtualFunctions.emplace_back(funcs.front()); + virtualFunctions.emplace_back(VirtualFunctionInfo{vFuncAddr}); } if (virtualFunctions.empty()) @@ -329,29 +456,12 @@ std::optional<VirtualFunctionTableInfo> ItaniumRTTIProcessor::ProcessVFT(uint64_ return std::nullopt; } - // All vft verification has been done, we can write the classOffset now. - if (classInfo.baseClassName.has_value() && !classInfo.classOffset.has_value()) - { - // Because we have this we _need_ to have the adjustment stuff. - // NOTE: We assume two 0x4 ints with the first being what we want. - // NOTE: This is where we actually classOffset is pulled. - reader.Seek(vftAddr - 0x10); - auto adjustmentOffset = static_cast<int32_t>(reader.Read32()); - auto _what = static_cast<int32_t>(reader.Read32()); - uint64_t classOffset = std::abs(adjustmentOffset); - classInfo.classOffset = classOffset; - } - - - for (auto &func: virtualFunctions) - vftInfo.virtualFunctions.emplace_back(VirtualFunctionInfo{func->GetStart()}); - // Create virtual function table type auto vftTypeName = fmt::format("{}::VTable", classInfo.className); - if (classInfo.baseClassName.has_value()) + if (baseClassInfo.has_value()) { // TODO: What is the correct form for the name? - vftTypeName = fmt::format("{}::{}", classInfo.baseClassName.value(), vftTypeName); + 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); @@ -368,17 +478,18 @@ std::optional<VirtualFunctionTableInfo> ItaniumRTTIProcessor::ProcessVFT(uint64_ auto vftSize = virtualFunctions.size() * addrSize; vftBuilder.SetWidth(vftSize); - if (auto baseVft = classInfo.baseVft) + if (baseClassInfo.has_value() && baseClassInfo->vft.has_value()) { - if (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 = 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()); + // TODO: Different type name please lol + 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. @@ -393,20 +504,46 @@ std::optional<VirtualFunctionTableInfo> ItaniumRTTIProcessor::ProcessVFT(uint64_ for (auto &&vFunc: virtualFunctions) { + // NOTE: The analyzed function type might not be available here. + auto vFuncAnalysis = m_view->GetAnalysisFunctionsForAddress(vFunc.funcAddr); + Ref<Type> vFuncType = nullptr; + Ref<Symbol> vFuncSym = nullptr; + if (!vFuncAnalysis.empty()) + { + vFuncType = vFuncAnalysis[0]->GetType(); + vFuncSym = vFuncAnalysis[0]->GetSymbol(); + } else + { + DataVariable dv; + bool foundDv = m_view->GetDataVariableAtAddress(vFunc.funcAddr, dv); + if (!foundDv) + { + m_logger->LogWarn("Skipping vfunc with no type... %llx", vFunc.funcAddr); + return std::nullopt; + } + vFuncType = dv.type.GetValue(); + + vFuncSym = m_view->GetSymbolByAddress(vFunc.funcAddr); + if (vFuncSym == nullptr) + { + m_logger->LogWarn("Skipping vfunc with no symbol... %llx", vFunc.funcAddr); + return std::nullopt; + } + } + auto vFuncName = fmt::format("vFunc_{}", vFuncIdx); // If we have a better name, use it. - auto vFuncSymName = vFunc->GetSymbol()->GetShortName(); + auto vFuncSymName = vFuncSym->GetShortName(); if (vFuncSymName.compare(0, 4, "sub_") != 0) - vFuncName = vFunc->GetSymbol()->GetShortName(); + vFuncName = vFuncSym->GetShortName(); // MyClass::func -> func std::size_t pos = vFuncName.rfind("::"); if (pos != std::string::npos) vFuncName = vFuncName.substr(pos + 2); - // NOTE: The analyzed function type might not be available here. auto vFuncOffset = vFuncIdx * addrSize; vftBuilder.AddMemberAtOffset( - Type::PointerType(addrSize, vFunc->GetType(), true), vFuncName, vFuncOffset); + Type::PointerType(addrSize, vFuncType, true), vFuncName, vFuncOffset); vFuncIdx++; } m_view->DefineType(typeId, vftTypeName, @@ -415,8 +552,8 @@ std::optional<VirtualFunctionTableInfo> ItaniumRTTIProcessor::ProcessVFT(uint64_ auto vftName = fmt::format("_vtable_for_{}", classInfo.className); // TODO: How to display base classes? - if (classInfo.baseClassName.has_value()) - vftName += fmt::format("{{for `{}'}}", classInfo.baseClassName.value()); + if (baseClassInfo.has_value()) + vftName += fmt::format("{{for `{}'}}", baseClassInfo->className); auto vftSymbol = m_view->GetSymbolByAddress(vftAddr); if (vftSymbol != nullptr) m_view->UndefineAutoSymbol(vftSymbol); @@ -459,6 +596,7 @@ void ItaniumRTTIProcessor::ProcessRTTI() } }; + m_view->BeginBulkModifySymbols(); // Scan data sections for rtti. for (const Ref<Section> §ion: m_view->GetSections()) { @@ -468,6 +606,49 @@ void ItaniumRTTIProcessor::ProcessRTTI() scan(section); } } + m_view->EndBulkModifySymbols(); + + // Go through all classes and recurse into the base classes using the base class name + for (auto &[classAddr, classInfo]: m_classInfo) + { + std::set<std::string> visitedBases; + std::deque<BaseClassInfo> baseQueue(classInfo.baseClasses.begin(), classInfo.baseClasses.end()); + + while (!baseQueue.empty()) + { + BaseClassInfo baseClass = baseQueue.front(); + baseQueue.pop_front(); + + if (visitedBases.find(baseClass.className) != visitedBases.end()) + continue; + + visitedBases.insert(baseClass.className); + + auto baseClassIt = std::find_if(m_classInfo.begin(), m_classInfo.end(), + [&](const auto &item) { + return item.second.className == baseClass.className; + }); + + if (baseClassIt != m_classInfo.end()) + { + const ClassInfo &nestedBaseClassInfo = baseClassIt->second; + baseQueue.insert(baseQueue.end(), nestedBaseClassInfo.baseClasses.begin(), + nestedBaseClassInfo.baseClasses.end()); + } + + classInfo.baseClasses.push_back(baseClass); + } + + // Remove duplicates in the baseClasses vector while preserving order + std::sort(classInfo.baseClasses.begin(), classInfo.baseClasses.end(), + [](const BaseClassInfo &a, const BaseClassInfo &b) { return a.className < b.className; }); + + classInfo.baseClasses.erase( + std::unique(classInfo.baseClasses.begin(), classInfo.baseClasses.end(), + [](const BaseClassInfo &a, const BaseClassInfo &b) { return a.className == b.className; }), + classInfo.baseClasses.end() + ); + } auto end_time = std::chrono::high_resolution_clock::now(); std::chrono::duration<double> elapsed_time = end_time - start_time; @@ -477,7 +658,8 @@ void ItaniumRTTIProcessor::ProcessRTTI() void ItaniumRTTIProcessor::ProcessVFT() { - std::map<uint64_t, uint64_t> vftMap = {}; + BinaryReader optReader = BinaryReader(m_view); + std::map<uint64_t, std::set<uint64_t>> vftMap = {}; std::map<uint64_t, std::optional<VirtualFunctionTableInfo>> vftFinishedMap = {}; auto start_time = std::chrono::high_resolution_clock::now(); for (auto &[coLocatorAddr, classInfo]: m_classInfo) @@ -488,15 +670,21 @@ void ItaniumRTTIProcessor::ProcessVFT() DataVariable dv; if (m_view->GetDataVariableAtAddress(ref, dv) && m_classInfo.find(dv.address) != m_classInfo.end()) continue; + // Verify that there is two 4 byte values above the type info pointer + optReader.Seek(ref - 8); + auto beforeTypeInfoRef = optReader.ReadPointer(); + if (m_view->IsValidOffset(beforeTypeInfoRef)) + continue; // TODO: This is not pointing at where it should, remember that the vtable will be inside another structure. auto vftAddr = ref + m_view->GetAddressSize(); - vftMap[coLocatorAddr] = vftAddr; + // Found a vtable reference to colocator + // TODO: Access check here. + vftMap[coLocatorAddr].insert(vftAddr); } } if (virtualFunctionTableSweep) { - BinaryReader optReader = BinaryReader(m_view); auto addrSize = m_view->GetAddressSize(); auto scan = [&](const Ref<Segment> &segment) { uint64_t startAddr = segment->GetStart(); @@ -509,7 +697,7 @@ void ItaniumRTTIProcessor::ProcessVFT() if (coLocator == m_classInfo.end()) continue; // Found a vtable reference to colocator. - vftMap[coLocatorAddr] = vtableAddr + addrSize; + // vftMap[coLocatorAddr] = vtableAddr + addrSize; } }; @@ -536,27 +724,35 @@ void ItaniumRTTIProcessor::ProcessVFT() auto cachedVftInfo = vftFinishedMap.find(vftAddr); if (cachedVftInfo != vftFinishedMap.end()) return cachedVftInfo->second; - auto vftInfo = ProcessVFT(vftAddr, classInfo); + // We need to have base class info available here. + // This works by reading off the adjustment and keying into the bases. + // If there is a base at that adjustment we assume this vtable we are creating is for that. + auto selectedBaseClass = ProcessVFTBaseClassInfo(vftAddr, classInfo); + auto vftInfo = ProcessVFT(vftAddr, classInfo, selectedBaseClass); vftFinishedMap[vftAddr] = vftInfo; return vftInfo; }; - for (const auto &[coLocatorAddr, vftAddr]: vftMap) - { + // Adds the VFT entries in class info and base class info. + // TODO: This is so cursed. + auto populateVftEntries = [&](uint64_t coLocatorAddr, uint64_t vftAddr) { 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, bClassInfo] : m_classInfo) { - if (baseClassInfo.className == classInfo.baseClassName.value()) + if (bClassInfo.className == baseClassInfo.className) { - uint64_t baseVftAddr = vftMap[baseCoLocAddr]; - if (auto baseVftInfo = GetCachedVFTInfo(baseVftAddr, baseClassInfo)) + // Recurse into base class and populate all of its vtables. + for (auto& baseVftAddr : vftMap[baseCoLocAddr]) { - classInfo.baseVft = baseVftInfo.value(); - break; + if (auto vftInfo = GetCachedVFTInfo(baseVftAddr, bClassInfo)) + bClassInfo.vft = vftInfo.value(); } + // Now that we have populated all the vtables for the base class, we can assign its + // root vtable to the base vft. + baseClassInfo.vft = bClassInfo.vft; } } } @@ -565,6 +761,14 @@ void ItaniumRTTIProcessor::ProcessVFT() classInfo.vft = vftInfo.value(); m_classInfo[coLocatorAddr] = classInfo; + }; + + for (const auto &[coLocatorAddr, vftAddrs]: vftMap) + { + for (const auto& vftAddr: vftAddrs) + { + populateVftEntries(coLocatorAddr, vftAddr); + } } auto end_time = std::chrono::high_resolution_clock::now(); diff --git a/plugins/rtti/itanium.h b/plugins/rtti/itanium.h index 98fe7052..7fdd23e5 100644 --- a/plugins/rtti/itanium.h +++ b/plugins/rtti/itanium.h @@ -58,7 +58,7 @@ namespace BinaryNinja::RTTI::Itanium { { uint64_t base_type; uint64_t offset_flags; - OffsetFlagsMasks offset_flags_masks; + uint64_t offset_flags_masks; BaseClassTypeInfo(BinaryView *view, uint64_t address); }; @@ -114,9 +114,11 @@ namespace BinaryNinja::RTTI::Itanium { bool checkWritableRData; bool virtualFunctionTableSweep; + std::optional<BaseClassInfo> ProcessVFTBaseClassInfo(uint64_t vftAddr, ClassInfo &classInfo); + std::optional<ClassInfo> ProcessRTTI(uint64_t objectAddr) override; - std::optional<VirtualFunctionTableInfo> ProcessVFT(uint64_t vftAddr, ClassInfo &classInfo) override; + std::optional<VirtualFunctionTableInfo> ProcessVFT(uint64_t vftAddr, ClassInfo &classInfo, std::optional<BaseClassInfo> baseClassInfo) override; public: explicit ItaniumRTTIProcessor(const Ref<BinaryView> &view, bool useMangled = true, bool checkRData = true, bool vttSweep = true); 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 diff --git a/plugins/rtti/microsoft.h b/plugins/rtti/microsoft.h index b67431ec..b92bf7f3 100644 --- a/plugins/rtti/microsoft.h +++ b/plugins/rtti/microsoft.h @@ -57,32 +57,24 @@ namespace BinaryNinja::RTTI::Microsoft { CompleteObjectLocator(BinaryView *view, uint64_t address); }; - class MicrosoftRTTIProcessor + class MicrosoftRTTIProcessor : public RTTIProcessor { - Ref<BinaryView> m_view; - Ref<Logger> m_logger; bool allowMangledClassNames; bool allowAnonymousClassNames; bool checkWritableRData; bool virtualFunctionTableSweep; - std::map<uint64_t, ClassInfo> m_classInfo; + // This will process a CHD and store all `BaseClassInfo` in `classInfo`. + std::vector<BaseClassInfo> ProcessClassHierarchyDescriptor(uint64_t address, CompleteObjectLocator &coLocator, const ClassInfo &classInfo); - std::set<uint64_t> m_visitedClassHierarchyDescAddrs; - - void DeserializedMetadata(const Ref<Metadata> &metadata); - - std::optional<ClassInfo> ProcessRTTI(uint64_t coLocatorAddr); - - std::optional<VirtualFunctionTableInfo> ProcessVFT(uint64_t vftAddr, const ClassInfo &classInfo); + std::optional<ClassInfo> ProcessRTTI(uint64_t objectAddr) override; + std::optional<VirtualFunctionTableInfo> ProcessVFT(uint64_t vftAddr, ClassInfo &classInfo, std::optional<BaseClassInfo> baseClassInfo) override; public: - MicrosoftRTTIProcessor(const Ref<BinaryView> &view, bool useMangled = true, bool checkRData = true, bool vftSweep = true, bool allowAnonymous = true); - - Ref<Metadata> SerializedMetadata(); + explicit MicrosoftRTTIProcessor(const Ref<BinaryView> &view, bool useMangled = true, bool checkRData = true, bool vftSweep = true, bool allowAnonymous = true); - void ProcessRTTI(); + void ProcessRTTI() override; - void ProcessVFT(); + void ProcessVFT() override; }; }
\ No newline at end of file diff --git a/plugins/rtti/plugin.cpp b/plugins/rtti/plugin.cpp index 82b51f78..faa741f1 100644 --- a/plugins/rtti/plugin.cpp +++ b/plugins/rtti/plugin.cpp @@ -58,12 +58,12 @@ extern "C" { // TODO: 2. Identify if the function is unique to a class, renaming and retyping if true // TODO: 3. Identify functions which address a VFT and are probably a constructor (alloc use), retyping if true // TODO: 4. Identify functions which address a VFT and are probably a deconstructor (free use), retyping if true - Ref<Workflow> rttiMetaWorkflow = Workflow::Instance("core.module.metaAnalysis")->Clone("core.module.metaAnalysis"); + Ref<Workflow> rttiMetaWorkflow = Workflow::Instance("core.module.metaAnalysis")->Clone(); // Add RTTI analysis. rttiMetaWorkflow->RegisterActivity(R"~({ "title": "RTTI Analysis", - "name": "plugin.rtti.rttiAnalysis", + "name": "analysis.rtti.rttiAnalysis", "role": "action", "description": "This analysis step attempts to parse and symbolize rtti information.", "eligibility": { @@ -74,7 +74,7 @@ extern "C" { // Add Virtual Function Table analysis. rttiMetaWorkflow->RegisterActivity(R"~({ "title": "VFT Analysis", - "name": "plugin.rtti.vftAnalysis", + "name": "analysis.rtti.vftAnalysis", "role": "action", "description": "This analysis step attempts to parse and symbolize virtual function table information.", "eligibility": { @@ -84,9 +84,9 @@ extern "C" { })~", &VFTAnalysis); // Run rtti before debug info is applied. - rttiMetaWorkflow->Insert("core.module.loadDebugInfo", "plugin.rtti.rttiAnalysis"); + rttiMetaWorkflow->Insert("core.module.loadDebugInfo", "analysis.rtti.rttiAnalysis"); // Run vft after functions have analyzed (so that the virtual functions have analyzed) - rttiMetaWorkflow->Insert("core.module.notifyCompletion", "plugin.rtti.vftAnalysis"); + rttiMetaWorkflow->Insert("core.module.deleteUnusedAutoFunctions", "analysis.rtti.vftAnalysis"); Workflow::RegisterWorkflow(rttiMetaWorkflow); return true; diff --git a/plugins/rtti/rtti.cpp b/plugins/rtti/rtti.cpp index d212edbc..91d3481c 100644 --- a/plugins/rtti/rtti.cpp +++ b/plugins/rtti/rtti.cpp @@ -14,19 +14,12 @@ std::optional<std::string> RTTI::DemangleNameMS(BinaryView* view, bool allowMang } -std::string RemoveItaniumPrefix(const std::string& name) { - // Remove class prefixes. - // 1 and 7 is class_type - // 9 is si_class_type - // 1..4 is vmi_class_type - if (name.rfind('1', 0) == 0) - return name.substr(1); - if (name.rfind('7', 0) == 0) - return name.substr(1); - if (name.rfind('9', 0) == 0) - return name.substr(1); - if (name.rfind("4", 0) == 0) - return name.substr(2); +std::string RemoveItaniumPrefix(std::string &name) +{ + // Remove numerical prefixes. + // TODO: We might want to use the numbers for figuring out the class info. + while (!name.empty() && std::isdigit(name[0])) + name = name.substr(1); return name; } @@ -56,54 +49,85 @@ std::optional<std::string> RTTI::DemangleNameLLVM(bool allowMangled, const std:: } -Ref<Metadata> ClassInfo::SerializedMetadata() +Ref<Metadata> BaseClassInfo::SerializedMetadata() const +{ + std::map<std::string, Ref<Metadata>> baseClassMeta; + baseClassMeta["className"] = new Metadata(className); + baseClassMeta["classOffset"] = new Metadata(offset); + // NOTE: We omit base vft functions as it can be resolved manually and just bloats the size. + if (vft.has_value()) + baseClassMeta["vft"] = vft->SerializedMetadata(false); + return new Metadata(baseClassMeta); +} + +BaseClassInfo BaseClassInfo::DeserializedMetadata(const Ref<Metadata> &metadata) { - std::map<std::string, Ref<Metadata> > classInfoMeta; + std::map<std::string, Ref<Metadata>> baseClassMeta = metadata->GetKeyValueStore(); + std::string className = baseClassMeta["className"]->GetString(); + uint64_t offset = baseClassMeta["classOffset"]->GetUnsignedInteger(); + BaseClassInfo baseClassInfo = {className, offset}; + if (baseClassMeta.find("vft") != baseClassMeta.end()) + baseClassInfo.vft = VirtualFunctionTableInfo::DeserializedMetadata(baseClassMeta["vft"]); + return baseClassInfo; +} + + +Ref<Metadata> ClassInfo::SerializedMetadata() const +{ + std::map<std::string, Ref<Metadata>> classInfoMeta; classInfoMeta["processor"] = new Metadata(static_cast<uint64_t>(processor)); 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 (!baseClasses.empty()) + { + std::vector<Ref<Metadata> > basesMeta; + basesMeta.reserve(baseClasses.size()); + for (const auto& baseClass : baseClasses) + basesMeta.emplace_back(baseClass.SerializedMetadata()); + classInfoMeta["bases"] = new Metadata(basesMeta); + } 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(); + std::map<std::string, Ref<Metadata>> classInfoMeta = metadata->GetKeyValueStore(); std::string className = classInfoMeta["className"]->GetString(); RTTIProcessorType processor = static_cast<RTTIProcessorType>(classInfoMeta["processor"]->GetUnsignedInteger()); ClassInfo info = {processor, className}; - 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("bases") != classInfoMeta.end()) + { + for (auto &entry: classInfoMeta["bases"]->GetArray()) + info.baseClasses.emplace_back(BaseClassInfo::DeserializedMetadata(entry)); + } if (classInfoMeta.find("vft") != classInfoMeta.end()) info.vft = VirtualFunctionTableInfo::DeserializedMetadata(classInfoMeta["vft"]); return info; } -Ref<Metadata> VirtualFunctionTableInfo::SerializedMetadata() +Ref<Metadata> VirtualFunctionTableInfo::SerializedMetadata(const bool serializeFunctions) const { - 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; + std::map<std::string, Ref<Metadata>> vftMeta; vftMeta["address"] = new Metadata(address); - vftMeta["functions"] = new Metadata(funcsMeta); + // NOTE: We allow omitting baseVft functions as it can be resolved manually and just bloats the size. + if (serializeFunctions && !virtualFunctions.empty()) + { + std::vector<Ref<Metadata> > funcsMeta; + funcsMeta.reserve(virtualFunctions.size()); + for (auto &vFunc: virtualFunctions) + funcsMeta.emplace_back(vFunc.SerializedMetadata()); + vftMeta["functions"] = new Metadata(funcsMeta); + } return new Metadata(vftMeta); } VirtualFunctionTableInfo VirtualFunctionTableInfo::DeserializedMetadata(const Ref<Metadata> &metadata) { - std::map<std::string, Ref<Metadata> > vftMeta = metadata->GetKeyValueStore(); + std::map<std::string, Ref<Metadata>> vftMeta = metadata->GetKeyValueStore(); VirtualFunctionTableInfo vftInfo = {vftMeta["address"]->GetUnsignedInteger()}; if (vftMeta.find("functions") != vftMeta.end()) { @@ -114,9 +138,9 @@ VirtualFunctionTableInfo VirtualFunctionTableInfo::DeserializedMetadata(const Re } -Ref<Metadata> VirtualFunctionInfo::SerializedMetadata() +Ref<Metadata> VirtualFunctionInfo::SerializedMetadata() const { - std::map<std::string, Ref<Metadata> > vFuncMeta; + std::map<std::string, Ref<Metadata>> vFuncMeta; vFuncMeta["address"] = new Metadata(funcAddr); return new Metadata(vFuncMeta); } @@ -124,7 +148,7 @@ Ref<Metadata> VirtualFunctionInfo::SerializedMetadata() VirtualFunctionInfo VirtualFunctionInfo::DeserializedMetadata(const Ref<Metadata> &metadata) { - std::map<std::string, Ref<Metadata> > vFuncMeta = metadata->GetKeyValueStore(); + std::map<std::string, Ref<Metadata>> vFuncMeta = metadata->GetKeyValueStore(); VirtualFunctionInfo vFuncInfo = {vFuncMeta["address"]->GetUnsignedInteger()}; return vFuncInfo; } @@ -132,7 +156,7 @@ VirtualFunctionInfo VirtualFunctionInfo::DeserializedMetadata(const Ref<Metadata Ref<Metadata> RTTIProcessor::SerializedMetadata() { - std::map<std::string, Ref<Metadata> > classesMeta; + std::map<std::string, Ref<Metadata>> classesMeta; for (auto &[objectAddr, classInfo]: m_classInfo) { auto addrStr = std::to_string(objectAddr); @@ -147,7 +171,7 @@ Ref<Metadata> RTTIProcessor::SerializedMetadata() classesMeta[addrStr] = classInfo.SerializedMetadata(); } - std::map<std::string, Ref<Metadata> > itaniumMeta; + std::map<std::string, Ref<Metadata>> itaniumMeta; itaniumMeta["classes"] = new Metadata(classesMeta); return new Metadata(itaniumMeta); } @@ -155,7 +179,7 @@ Ref<Metadata> RTTIProcessor::SerializedMetadata() void RTTIProcessor::DeserializedMetadata(RTTIProcessorType type, const Ref<Metadata> &metadata) { - std::map<std::string, Ref<Metadata> > msvcMeta = metadata->GetKeyValueStore(); + std::map<std::string, Ref<Metadata>> msvcMeta = metadata->GetKeyValueStore(); if (msvcMeta.find("classes") != msvcMeta.end()) { for (auto &[objectAddrStr, classInfoMeta]: msvcMeta["classes"]->GetKeyValueStore()) diff --git a/plugins/rtti/rtti.h b/plugins/rtti/rtti.h index 474253ec..e2913db4 100644 --- a/plugins/rtti/rtti.h +++ b/plugins/rtti/rtti.h @@ -16,7 +16,7 @@ namespace BinaryNinja::RTTI { { uint64_t funcAddr; - Ref<Metadata> SerializedMetadata(); + Ref<Metadata> SerializedMetadata() const; static VirtualFunctionInfo DeserializedMetadata(const Ref<Metadata> &metadata); }; @@ -26,7 +26,7 @@ namespace BinaryNinja::RTTI { uint64_t address; std::vector<VirtualFunctionInfo> virtualFunctions; - Ref<Metadata> SerializedMetadata(); + Ref<Metadata> SerializedMetadata(bool serializeFunctions = true) const; static VirtualFunctionTableInfo DeserializedMetadata(const Ref<Metadata> &metadata); }; @@ -37,17 +37,29 @@ namespace BinaryNinja::RTTI { Itanium = 1, }; + struct BaseClassInfo + { + std::string className; + // TODO: This has to be optional, as we might need to resolve this at a later stage. + // TODO: The offset also might literally not exist. + uint64_t offset; + std::optional<VirtualFunctionTableInfo> vft; + + Ref<Metadata> SerializedMetadata() const; + + static BaseClassInfo DeserializedMetadata(const Ref<Metadata> &metadata); + }; + // TODO: This needs to have some flags. Virtual, pure iirc. struct ClassInfo { RTTIProcessorType processor; std::string className; - std::optional<std::string> baseClassName; - std::optional<uint64_t> classOffset; + std::optional<VirtualFunctionTableInfo> vft; - std::optional<VirtualFunctionTableInfo> baseVft; + std::vector<BaseClassInfo> baseClasses; - Ref<Metadata> SerializedMetadata(); + Ref<Metadata> SerializedMetadata() const; static ClassInfo DeserializedMetadata(const Ref<Metadata> &metadata); }; @@ -63,7 +75,7 @@ namespace BinaryNinja::RTTI { virtual std::optional<ClassInfo> ProcessRTTI(uint64_t objectAddr) = 0; - virtual std::optional<VirtualFunctionTableInfo> ProcessVFT(uint64_t vftAddr, ClassInfo &classInfo) = 0; + virtual std::optional<VirtualFunctionTableInfo> ProcessVFT(uint64_t vftAddr, ClassInfo &classInfo, std::optional<BaseClassInfo> baseClassInfo) = 0; public: virtual ~RTTIProcessor() = default; |
