summaryrefslogtreecommitdiff
path: root/plugins/msvc_rtti/rtti.cpp
diff options
context:
space:
mode:
authorMason Reed <mason@vector35.com>2024-10-23 22:10:17 -0400
committerMason Reed <mason@vector35.com>2024-10-24 10:53:22 -0400
commit93e0a64e77169c29960a1cbd9bdedadfeb4a5f7e (patch)
treed2bad13887e0d7f45287690e02236ddffed9a336 /plugins/msvc_rtti/rtti.cpp
parent21488f76c5d33485323bccda36ef714453e1aac7 (diff)
Add MSVC RTTI plugin
Adds two commands that must be run manually "MSVC\\Find RTTI" and "MSVC\\Find VFTs" both of which will apply their respective data to the view AND store metadata for scripts to use under the "msvc" key.
Diffstat (limited to 'plugins/msvc_rtti/rtti.cpp')
-rw-r--r--plugins/msvc_rtti/rtti.cpp672
1 files changed, 672 insertions, 0 deletions
diff --git a/plugins/msvc_rtti/rtti.cpp b/plugins/msvc_rtti/rtti.cpp
new file mode 100644
index 00000000..fe89e739
--- /dev/null
+++ b/plugins/msvc_rtti/rtti.cpp
@@ -0,0 +1,672 @@
+#include "rtti.h"
+
+using namespace BinaryNinja;
+
+constexpr int COL_SIG_REV0 = 0;
+constexpr int COL_SIG_REV1 = 1;
+constexpr int RTTI_CONFIDENCE = 100;
+
+
+ClassHierarchyDescriptor::ClassHierarchyDescriptor(BinaryView *view, uint64_t address)
+{
+ BinaryReader reader = BinaryReader(view);
+ reader.Seek(address);
+ signature = reader.Read32();
+ attributes = reader.Read32();
+ numBaseClasses = reader.Read32();
+ pBaseClassArray = static_cast<int32_t>(reader.Read32());
+}
+
+
+BaseClassDescriptor::BaseClassDescriptor(BinaryView *view, uint64_t address)
+{
+ BinaryReader reader = BinaryReader(view);
+ reader.Seek(address);
+ pTypeDescriptor = static_cast<int32_t>(reader.Read32());
+ numContainedBases = reader.Read32();
+ where_mdisp = static_cast<int32_t>(reader.Read32());
+ where_pdisp = static_cast<int32_t>(reader.Read32());
+ where_vdisp = static_cast<int32_t>(reader.Read32());
+ attributes = reader.Read32();
+ pClassHierarchyDescriptor = static_cast<int32_t>(reader.Read32());
+}
+
+
+BaseClassArray::BaseClassArray(BinaryView *view, uint64_t address, uint32_t length) : length(length)
+{
+ BinaryReader reader = BinaryReader(view);
+ reader.Seek(address);
+ descriptors = {};
+ for (size_t i = 0; i < length; i++)
+ descriptors.emplace_back(reader.Read32());
+}
+
+
+TypeDescriptor::TypeDescriptor(BinaryView *view, uint64_t address)
+{
+ BinaryReader reader = BinaryReader(view);
+ reader.Seek(address);
+ pVFTable = reader.ReadPointer();
+ spare = reader.ReadPointer();
+ name = reader.ReadCString(512);
+}
+
+
+CompleteObjectLocator::CompleteObjectLocator(BinaryView *view, uint64_t address)
+{
+ BinaryReader reader = BinaryReader(view);
+ reader.Seek(address);
+ signature = reader.Read32();
+ offset = reader.Read32();
+ cdOffset = reader.Read32();
+ pTypeDescriptor = static_cast<int32_t>(reader.Read32());
+ pClassHeirarchyDescriptor = static_cast<int32_t>(reader.Read32());
+ if (signature == COL_SIG_REV1)
+ {
+ pSelf = static_cast<int32_t>(reader.Read32());
+ } else
+ {
+ pSelf = 0;
+ }
+}
+
+
+std::optional<CompleteObjectLocator> ReadCompleteObjectorLocator(BinaryView *view, uint64_t address)
+{
+ auto coLocator = CompleteObjectLocator(view, address);
+ uint64_t startAddr = view->GetStart();
+ uint64_t endAddr = view->GetEnd();
+
+ if (coLocator.signature > 1)
+ return std::nullopt;
+
+ if (coLocator.signature == COL_SIG_REV1)
+ {
+ if (coLocator.pSelf + startAddr != address)
+ return std::nullopt;
+
+ // Relative addrs
+ if (coLocator.pTypeDescriptor + startAddr > endAddr)
+ return std::nullopt;
+
+ if (coLocator.pClassHeirarchyDescriptor + startAddr > endAddr)
+ return std::nullopt;
+ } else
+ {
+ // Absolute addrs
+ if (coLocator.pTypeDescriptor < startAddr || coLocator.pTypeDescriptor > endAddr)
+ return std::nullopt;
+
+ if (coLocator.pClassHeirarchyDescriptor < startAddr || coLocator.pClassHeirarchyDescriptor > endAddr)
+ return std::nullopt;
+ }
+
+ return coLocator;
+}
+
+
+Ref<Type> GetPMDType(BinaryView *view)
+{
+ auto typeId = Type::GenerateAutoTypeId("msvc_rtti", QualifiedName("PMD"));
+ Ref<Type> typeCache = view->GetTypeById(typeId);
+
+ if (typeCache == nullptr)
+ {
+ Ref<Type> intType = Type::IntegerType(4, true);
+
+ StructureBuilder pmdBuilder;
+ pmdBuilder.AddMember(intType, "mdisp");
+ pmdBuilder.AddMember(intType, "pdisp");
+ pmdBuilder.AddMember(intType, "vdisp");
+
+ view->DefineType(typeId, QualifiedName("_PMD"), TypeBuilder::StructureType(pmdBuilder.Finalize()).Finalize());
+ typeCache = view->GetTypeById(typeId);
+ }
+
+ return typeCache;
+}
+
+
+Ref<Type> ClassHierarchyDescriptorType(BinaryView *view, BNPointerBaseType ptrBaseTy);
+
+Ref<Type> BaseClassDescriptorType(BinaryView *view, BNPointerBaseType ptrBaseTy)
+{
+ auto typeId = Type::GenerateAutoTypeId("msvc_rtti", QualifiedName("RTTIBaseClassDescriptor"));
+ Ref<Type> typeCache = view->GetTypeById(typeId);
+
+ if (typeCache == nullptr)
+ {
+ Ref<Type> uintType = Type::IntegerType(4, false);
+
+ StructureBuilder baseClassDescriptorBuilder;
+ // Would require creating a new type for every type descriptor length. Instead just use void*
+ Ref<Type> pTypeDescType = TypeBuilder::PointerType(4, Type::VoidType())
+ .SetPointerBase(ptrBaseTy, 0)
+ .Finalize();
+ baseClassDescriptorBuilder.AddMember(pTypeDescType, "pTypeDescriptor");
+ baseClassDescriptorBuilder.AddMember(uintType, "numContainedBases");
+ baseClassDescriptorBuilder.AddMember(GetPMDType(view), "where");
+ baseClassDescriptorBuilder.AddMember(uintType, "attributes");
+ Ref<Type> pClassDescType = TypeBuilder::PointerType(4, ClassHierarchyDescriptorType(view, ptrBaseTy))
+ .SetPointerBase(ptrBaseTy, 0)
+ .Finalize();
+ baseClassDescriptorBuilder.AddMember(pClassDescType, "pClassDescriptor");
+
+ view->DefineType(typeId, QualifiedName("_RTTIBaseClassDescriptor"),
+ TypeBuilder::StructureType(baseClassDescriptorBuilder.Finalize()).Finalize());
+ typeCache = view->GetTypeById(typeId);
+ }
+
+ return typeCache;
+}
+
+
+Ref<Type> BaseClassArrayType(BinaryView *view, const uint64_t length, BNPointerBaseType ptrBaseTy)
+{
+ StructureBuilder baseClassArrayBuilder;
+ Ref<Type> pBaseClassDescType = TypeBuilder::PointerType(4, BaseClassDescriptorType(view, ptrBaseTy))
+ .SetPointerBase(ptrBaseTy, 0)
+ .Finalize();
+ baseClassArrayBuilder.AddMember(
+ Type::ArrayType(pBaseClassDescType, length), "arrayOfBaseClassDescriptors");
+ return TypeBuilder::StructureType(baseClassArrayBuilder.Finalize()).Finalize();
+}
+
+
+Ref<Type> ClassHierarchyDescriptorType(BinaryView *view, BNPointerBaseType ptrBaseTy)
+{
+ auto typeId = Type::GenerateAutoTypeId("msvc_rtti", QualifiedName("RTTIClassHierarchyDescriptor"));
+ Ref<Type> typeCache = view->GetTypeById(typeId);
+
+ if (typeCache == nullptr)
+ {
+ Ref<Type> uintType = Type::IntegerType(4, false);
+
+ StructureBuilder classHierarchyDescriptorBuilder;
+ classHierarchyDescriptorBuilder.AddMember(uintType, "signature");
+ classHierarchyDescriptorBuilder.AddMember(uintType, "attributes");
+ classHierarchyDescriptorBuilder.AddMember(uintType, "numBaseClasses");
+ Ref<Type> pBaseClassArrayType = TypeBuilder::PointerType(4, Type::VoidType())
+ .SetPointerBase(ptrBaseTy, 0)
+ .Finalize();
+ classHierarchyDescriptorBuilder.AddMember(pBaseClassArrayType, "pBaseClassArray");
+
+ view->DefineType(typeId, QualifiedName("_RTTIClassHierarchyDescriptor"),
+ TypeBuilder::StructureType(classHierarchyDescriptorBuilder.Finalize()).Finalize());
+
+ typeCache = view->GetTypeById(typeId);
+ }
+
+ return typeCache;
+}
+
+
+Ref<Type> CompleteObjectLocator64Type(BinaryView *view)
+{
+ auto typeId = Type::GenerateAutoTypeId("msvc_rtti", QualifiedName("RTTICompleteObjectLocator64"));
+ Ref<Type> typeCache = view->GetTypeById(typeId);
+
+ if (typeCache == nullptr)
+ {
+ Ref<Architecture> arch = view->GetDefaultArchitecture();
+ Ref<Type> uintType = Type::IntegerType(4, false);
+
+ StructureBuilder completeObjectLocatorBuilder;
+ Ref<Enumeration> sigEnum = EnumerationBuilder()
+ .AddMemberWithValue("COL_SIG_REV0", 0)
+ .AddMemberWithValue("COL_SIG_REV1", 1)
+ .Finalize();
+ Ref<Type> sigType = Type::EnumerationType(arch, sigEnum, 4);
+ completeObjectLocatorBuilder.AddMember(sigType, "signature");
+ completeObjectLocatorBuilder.AddMember(uintType, "offset");
+ completeObjectLocatorBuilder.AddMember(uintType, "cdOffset");
+ Ref<Type> pTypeDescType = TypeBuilder::PointerType(4, Type::VoidType())
+ .SetPointerBase(RelativeToBinaryStartPointerBaseType, 0)
+ .Finalize();
+ completeObjectLocatorBuilder.AddMember(pTypeDescType, "pTypeDescriptor");
+ Ref<Type> pClassHierarchyDescType = TypeBuilder::PointerType(
+ 4, ClassHierarchyDescriptorType(view, RelativeToBinaryStartPointerBaseType))
+ .SetPointerBase(RelativeToBinaryStartPointerBaseType, 0)
+ .Finalize();
+ completeObjectLocatorBuilder.AddMember(pClassHierarchyDescType, "pClassHierarchyDescriptor");
+ Ref<Type> pSelfType = TypeBuilder::PointerType(4, Type::NamedType(view, typeId))
+ .SetPointerBase(RelativeToBinaryStartPointerBaseType, 0)
+ .Finalize();
+ completeObjectLocatorBuilder.AddMember(pSelfType, "pSelf");
+
+ view->DefineType(typeId, QualifiedName("_RTTICompleteObjectLocator"),
+ TypeBuilder::StructureType(completeObjectLocatorBuilder.Finalize()).Finalize());
+
+ typeCache = view->GetTypeById(typeId);
+ }
+
+ return typeCache;
+}
+
+
+Ref<Type> CompleteObjectLocator32Type(BinaryView *view)
+{
+ auto typeId = Type::GenerateAutoTypeId("msvc_rtti", QualifiedName("RTTICompleteObjectLocator32"));
+ Ref<Type> typeCache = view->GetTypeById(typeId);
+
+ if (typeCache == nullptr)
+ {
+ Ref<Architecture> arch = view->GetDefaultArchitecture();
+ Ref<Type> uintType = Type::IntegerType(4, false);
+
+ StructureBuilder completeObjectLocatorBuilder;
+ Ref<Enumeration> sigEnum = EnumerationBuilder()
+ .AddMemberWithValue("COL_SIG_REV0", 0)
+ .AddMemberWithValue("COL_SIG_REV1", 1)
+ .Finalize();
+ Ref<Type> sigType = Type::EnumerationType(arch, sigEnum, 4);
+ completeObjectLocatorBuilder.AddMember(sigType, "signature");
+ completeObjectLocatorBuilder.AddMember(uintType, "offset");
+ completeObjectLocatorBuilder.AddMember(uintType, "cdOffset");
+ Ref<Type> pTypeDescType = TypeBuilder::PointerType(4, Type::VoidType())
+ .Finalize();
+ completeObjectLocatorBuilder.AddMember(pTypeDescType, "pTypeDescriptor");
+ Ref<Type> pClassHierarchyDescType = TypeBuilder::PointerType(
+ 4, ClassHierarchyDescriptorType(view, AbsolutePointerBaseType))
+ .Finalize();
+ completeObjectLocatorBuilder.AddMember(pClassHierarchyDescType, "pClassHierarchyDescriptor");
+
+ view->DefineType(typeId, QualifiedName("_RTTICompleteObjectLocator"),
+ TypeBuilder::StructureType(completeObjectLocatorBuilder.Finalize()).Finalize());
+
+ typeCache = view->GetTypeById(typeId);
+ }
+
+ return typeCache;
+}
+
+
+Ref<Type> TypeDescriptorType(BinaryView *view, uint64_t length)
+{
+ size_t addrSize = view->GetAddressSize();
+ StructureBuilder typeDescriptorBuilder;
+ typeDescriptorBuilder.AddMember(Type::PointerType(addrSize, Type::VoidType(), true), "pVFTable");
+ typeDescriptorBuilder.AddMember(Type::PointerType(addrSize, Type::VoidType()), "spare");
+ // Char array needs to be individually resized.
+ typeDescriptorBuilder.AddMember(Type::ArrayType(Type::IntegerType(1, true, "char"), length), "name");
+ return TypeBuilder::StructureType(typeDescriptorBuilder.Finalize()).Finalize();
+}
+
+
+Ref<Metadata> ClassInfo::SerializedMetadata()
+{
+ 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();
+ 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;
+}
+
+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);
+}
+
+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())
+ {
+ 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;
+}
+
+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();
+ }
+
+ std::map<std::string, Ref<Metadata> > msvcMeta;
+ msvcMeta["classes"] = new Metadata(classesMeta);
+ return new Metadata(msvcMeta);
+}
+
+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())
+ {
+ uint64_t coLocatorAddr = std::stoull(coLocatorAddrStr);
+ m_classInfo[coLocatorAddr] = ClassInfo::DeserializedMetadata(classInfoMeta);
+ }
+ }
+}
+
+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();
+}
+
+
+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;
+
+ auto startAddr = m_view->GetStart();
+ auto resolveAddr = [&](const uint64_t relAddr) {
+ return coLocator->signature == COL_SIG_REV1 ? startAddr + relAddr : relAddr;
+ };
+
+ auto ptrBaseTy = coLocator->signature ? RelativeToBinaryStartPointerBaseType : AbsolutePointerBaseType;
+
+ // 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);
+ if (!className.has_value())
+ return std::nullopt;
+
+ 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 classHierarchyDescAddr = resolveAddr(coLocator->pClassHeirarchyDescriptor);
+ 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));
+
+ 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;
+ }
+
+ if (baseClassDesc.where_mdisp == coLocator->offset && classInfo.className != baseClassName.value())
+ classInfo.baseClassName = baseClassName;
+
+ auto baseClassDescName = fmt::format("{}::`RTTI Base Class Descriptor at ({},{},{},{})", baseClassName.value(),
+ 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 coLocatorName = fmt::format("{}::`RTTI Complete Object Locator'", className.value());
+ if (classInfo.baseClassName.has_value())
+ coLocatorName += fmt::format("{{for `{}'}}", classInfo.baseClassName.value());
+ m_view->DefineAutoSymbol(new Symbol{DataSymbol, coLocatorName, coLocatorAddr});
+ if (coLocator->signature == COL_SIG_REV1)
+ m_view->DefineDataVariable(coLocatorAddr, Confidence(CompleteObjectLocator64Type(m_view), RTTI_CONFIDENCE));
+ else
+ m_view->DefineDataVariable(coLocatorAddr, Confidence(CompleteObjectLocator32Type(m_view), RTTI_CONFIDENCE));
+
+ return classInfo;
+}
+
+
+std::optional<VirtualFunctionTableInfo> MicrosoftRTTIProcessor::ProcessVFT(uint64_t vftAddr, const ClassInfo &classInfo)
+{
+ VirtualFunctionTableInfo vftInfo = {vftAddr};
+ // Gather all virtual functions
+ BinaryReader reader = BinaryReader(m_view);
+ reader.Seek(vftAddr);
+ std::vector<Ref<Function> > virtualFunctions = {};
+ while (true)
+ {
+ uint64_t vFuncAddr = reader.ReadPointer();
+ auto funcs = m_view->GetAnalysisFunctionsForAddress(vFuncAddr);
+ if (funcs.empty())
+ {
+ Ref<Segment> segment = m_view->GetSegmentAt(vFuncAddr);
+ if (segment == nullptr || !(segment->GetFlags() & (SegmentExecutable | SegmentDenyWrite)))
+ {
+ // Last CompleteObjectLocator or hit the next CompleteObjectLocator
+ break;
+ }
+ // 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());
+ }
+
+ if (virtualFunctions.empty())
+ {
+ m_logger->LogDebug("Skipping empty virtual function table... %llx", vftAddr);
+ return std::nullopt;
+ }
+
+ 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())
+ {
+ vftTypeName = fmt::format("{}::{}", classInfo.baseClassName.value(), vftTypeName);
+ // TODO: What is the correct form for the name?
+ }
+ // 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);
+ Ref<Type> vftType = m_view->GetTypeById(typeId);
+
+ // TODO we need to inherit vtables in cases where the backing functions are coupled to a no member field vtable.
+ // TOOD: Inheriting vtables is a terrible idea usually btw.
+
+ if (vftType == nullptr)
+ {
+ size_t addrSize = m_view->GetAddressSize();
+ StructureBuilder vftBuilder = {};
+ vftBuilder.SetPropagateDataVariableReferences(true);
+ size_t vFuncIdx = 0;
+ for (auto &&vFunc: virtualFunctions)
+ {
+ // TODO: Identify when the functions name can be used instead of the vFunc_* placeholder.
+ auto vFuncName = fmt::format("vFunc_{}", vFuncIdx);
+ // NOTE: The analyzed function type might not be available here.
+ vftBuilder.AddMember(
+ Type::PointerType(addrSize, vFunc->GetType(), true), vFuncName);
+ vFuncIdx++;
+ }
+ m_view->DefineType(typeId, vftTypeName,
+ Confidence(TypeBuilder::StructureType(vftBuilder.Finalize()).Finalize(), RTTI_CONFIDENCE));
+ vftType = m_view->GetTypeById(typeId);
+ }
+
+ auto vftName = fmt::format("{}::`vftable'", classInfo.className);
+ if (classInfo.baseClassName.has_value())
+ vftName += fmt::format("{{for `{}'}}", classInfo.baseClassName.value());
+ m_view->DefineAutoSymbol(new Symbol{DataSymbol, vftName, vftAddr});
+ m_view->DefineDataVariable(vftAddr, Confidence(vftType, RTTI_CONFIDENCE));
+ return vftInfo;
+}
+
+
+MicrosoftRTTIProcessor::MicrosoftRTTIProcessor(const Ref<BinaryView> &view, bool useMangled, bool checkRData) : m_view(
+ view)
+{
+ m_logger = new Logger("Microsoft RTTI");
+ allowMangledClassNames = useMangled;
+ checkWritableRData = checkRData;
+ m_classInfo = {};
+ auto metadata = view->QueryMetadata(VIEW_METADATA_MSVC);
+ if (metadata != nullptr)
+ {
+ // Load in metadata to the processor.
+ DeserializedMetadata(metadata);
+ }
+}
+
+
+void MicrosoftRTTIProcessor::ProcessRTTI()
+{
+ auto start_time = std::chrono::high_resolution_clock::now();
+ uint64_t startAddr = m_view->GetStart();
+ uint64_t endAddr = m_view->GetEnd();
+ BinaryReader optReader = BinaryReader(m_view);
+ auto addrSize = m_view->GetAddressSize();
+
+ auto scan = [&](const Ref<Segment> &segment) {
+ for (uint64_t coLocatorAddr = segment->GetStart(); coLocatorAddr < segment->GetEnd() - 0x18;
+ coLocatorAddr += addrSize)
+ {
+ optReader.Seek(coLocatorAddr);
+ uint32_t sigVal = optReader.Read32();
+ if (sigVal == COL_SIG_REV1)
+ {
+ // Check for self reference
+ optReader.SeekRelative(16);
+ if (optReader.Read32() == coLocatorAddr - startAddr)
+ {
+ if (auto classInfo = ProcessRTTI(coLocatorAddr))
+ m_classInfo[coLocatorAddr] = classInfo.value();
+ }
+ } else if (sigVal == COL_SIG_REV0)
+ {
+ // Check ?AV
+ optReader.SeekRelative(8);
+ uint64_t typeDescNameAddr = optReader.Read32() + 8;
+ if (typeDescNameAddr > startAddr && typeDescNameAddr < endAddr)
+ {
+ // Make sure we do not read across segment boundary.
+ auto typeDescSegment = m_view->GetSegmentAt(typeDescNameAddr);
+ if (typeDescSegment != nullptr && typeDescSegment->GetEnd() - typeDescNameAddr > 4)
+ {
+ optReader.Seek(typeDescNameAddr);
+ auto typeDescNameStart = optReader.ReadString(4);
+ if (typeDescNameStart == ".?AV" || typeDescNameStart == ".?AU" || typeDescNameStart == ".?AW")
+ {
+ if (auto classInfo = ProcessRTTI(coLocatorAddr))
+ m_classInfo[coLocatorAddr] = classInfo.value();
+ }
+ }
+ }
+ }
+ }
+ };
+
+ // Scan data sections for colocators.
+ auto rdataSection = m_view->GetSectionByName(".rdata");
+ for (const Ref<Segment> &segment: m_view->GetSegments())
+ {
+ if (segment->GetFlags() == (SegmentReadable | SegmentContainsData))
+ {
+ m_logger->LogDebug("Attempting to find VirtualFunctionTables 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",
+ segment->GetStart());
+ scan(segment);
+ }
+ }
+
+ auto end_time = std::chrono::high_resolution_clock::now();
+ std::chrono::duration<double> elapsed_time = end_time - start_time;
+ m_logger->LogInfo("ProcessRTTI took %f seconds", elapsed_time.count());
+}
+
+void MicrosoftRTTIProcessor::ProcessVFT()
+{
+ auto start_time = std::chrono::high_resolution_clock::now();
+ for (auto &[coLocatorAddr, classInfo]: m_classInfo)
+ {
+ for (auto &ref: m_view->GetDataReferences(coLocatorAddr))
+ {
+ auto vftAddr = ref + m_view->GetAddressSize();
+ if (auto vftInfo = ProcessVFT(vftAddr, classInfo))
+ m_classInfo[coLocatorAddr].vft = vftInfo.value();
+ }
+ }
+
+ 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