VoltMod
C++23 framework for CS2 server plugins
Loading...
Searching...
No Matches
VtableLookup.windows.cpp
Go to the documentation of this file.
2
3#include <windows.h>
4
5#include <array>
6#include <cstdint>
7#include <cstring>
8#include <format>
9#include <span>
10#include <string>
11#include <string_view>
12#include <utility>
13
14namespace VoltMod
15{
16
17// MSVC x64 RTTI records store offsets; RVAs are relative to the module base.
18static constexpr size_t TypeDescriptorName = 0x10;
19static constexpr size_t LocatorOffset = 0x04;
20static constexpr size_t LocatorTypeDescriptor = 0x0C;
21static constexpr size_t LocatorHierarchy = 0x10;
22static constexpr size_t LocatorSelf = 0x14;
23static constexpr size_t LocatorSize = 0x18;
24static constexpr size_t HierarchyBaseCount = 0x08;
25static constexpr size_t HierarchyBaseList = 0x0C;
26static constexpr size_t HierarchySize = 0x10;
27static constexpr size_t BaseMemberOffset = 0x08;
28static constexpr size_t BaseVirtualOffset = 0x0C;
29static constexpr size_t BaseSize = 0x18;
30static constexpr uint32_t MaxBases = 1024;
31
32static const uint8_t* FindValue(const uint8_t* begin, const uint8_t* end, const void* needle, size_t len, size_t stride)
33{
34 if (!begin || !end || len == 0 || static_cast<size_t>(end - begin) < len)
35 {
36 return nullptr;
37 }
38
39 for (const uint8_t* at = begin; at + len <= end; at += stride)
40 {
41 if (std::memcmp(at, needle, len) == 0)
42 {
43 return at;
44 }
45 }
46 return nullptr;
47}
48
49template <typename T>
50static T ReadAt(const uint8_t* address)
51{
52 T value{};
53 std::memcpy(&value, address, sizeof(T));
54 return value;
55}
56
57static const uint8_t* End(const ScanRange& range)
58{
59 return range.Base ? range.Base + range.Size : nullptr;
60}
61
62static const uint8_t* AtRva(const PeRtti& rtti, int64_t rva, size_t bytes)
63{
64 if (rva < 0 || static_cast<uint64_t>(rva) > rtti.Size || rtti.Size - static_cast<size_t>(rva) < bytes)
65 {
66 return nullptr;
67 }
68 return rtti.Base + rva;
69}
70
71static std::array<std::string, 2> ClassAndStructNames(std::string_view name)
72{
73 return {std::format(".?AV{}@@", name), std::format(".?AU{}@@", name)};
74}
75
76static uint32_t FindTypeDescriptor(const PeRtti& rtti, std::string_view className)
77{
78 for (const std::string& mangled : ClassAndStructNames(className))
79 {
80 // Include the terminator to reject longer names.
81 const uint8_t* name = FindValue(rtti.Data.Base, End(rtti.Data), mangled.c_str(), mangled.size() + 1, 1);
82 if (name && static_cast<size_t>(name - rtti.Base) >= TypeDescriptorName)
83 {
84 return static_cast<uint32_t>(name - rtti.Base - TypeDescriptorName);
85 }
86 }
87 return 0;
88}
89
90static bool TypeDescriptorMatches(const PeRtti& rtti, int32_t rva, std::span<const std::string> names)
91{
92 for (const std::string& mangled : names)
93 {
94 const uint8_t* name = AtRva(rtti, int64_t{rva} + TypeDescriptorName, mangled.size() + 1);
95 if (name && std::memcmp(name, mangled.c_str(), mangled.size() + 1) == 0)
96 {
97 return true;
98 }
99 }
100 return false;
101}
102
104{
105 const uint8_t* begin = rtti.ReadOnlyData.Base;
106 const uint8_t* end = End(rtti.ReadOnlyData);
107 for (const uint8_t* ref = FindValue(begin, end, &typeDescriptor, sizeof(uint32_t), 4); ref;
108 ref = FindValue(ref + 4, end, &typeDescriptor, sizeof(uint32_t), 4))
109 {
110 if (static_cast<size_t>(ref - begin) < LocatorTypeDescriptor ||
111 static_cast<size_t>(end - ref) < LocatorSize - LocatorTypeDescriptor)
112 {
113 continue;
114 }
115
118 ReadAt<uint32_t>(locator + LocatorSelf) == static_cast<uint32_t>(locator - rtti.Base))
119 {
120 return locator;
121 }
122 }
123 return nullptr;
124}
125
126static void* TableAfter(const PeRtti& rtti, const uint8_t* locator)
127{
128 const uint8_t* word =
129 FindValue(rtti.ReadOnlyData.Base, End(rtti.ReadOnlyData), &locator, sizeof(void*), sizeof(void*));
130 return word ? const_cast<uint8_t*>(word + sizeof(void*)) : nullptr;
131}
132
133static std::pair<uint32_t, const uint8_t*> PrimaryLocator(const PeRtti& rtti, std::string_view className)
134{
137}
138
139void* FindVirtualTableInRtti(const PeRtti& rtti, std::string_view className)
140{
141 const uint8_t* locator = PrimaryLocator(rtti, className).second;
142 return locator ? TableAfter(rtti, locator) : nullptr;
143}
144
145Result<BaseSubobject> FindBaseInRtti(const PeRtti& rtti, std::string_view className, std::string_view baseName)
146{
148 if (!locator)
149 {
150 return std::unexpected(Error::NotFound(std::format("no RTTI for '{}'", className)));
151 }
152
155 const uint8_t* bases = hierarchy && count <= MaxBases
157 : nullptr;
158 if (!bases)
159 {
160 return std::unexpected(Error::Invalid(std::format("the RTTI base list of '{}' is unreadable", className)));
161 }
162
163 const std::array<std::string, 2> wanted = ClassAndStructNames(baseName);
164 std::vector<int32_t> offsets;
165 size_t virtualBases = 0;
166 // The first entry is the class itself; remaining entries are bases and offsets.
167 for (uint32_t i = 1; i < count; ++i)
168 {
169 const uint8_t* base = AtRva(rtti, ReadAt<int32_t>(bases + i * sizeof(int32_t)), BaseSize);
170 if (!base || !TypeDescriptorMatches(rtti, ReadAt<int32_t>(base), wanted))
171 {
172 continue;
173 }
174
175 if (ReadAt<int32_t>(base + BaseVirtualOffset) != -1)
176 {
177 ++virtualBases;
178 }
179 else
180 {
181 offsets.push_back(ReadAt<int32_t>(base + BaseMemberOffset));
182 }
183 }
184
185 if (offsets.size() + virtualBases > 1)
186 {
187 return std::unexpected(
188 Error::Invalid(std::format("'{}' is a base {} times", baseName, offsets.size() + virtualBases)));
189 }
190 if (virtualBases)
191 {
192 return std::unexpected(Error::Unsupported(std::format("'{}' is a virtual base", baseName)));
193 }
194 if (offsets.empty())
195 {
196 return std::unexpected(Error::NotFound(std::format("'{}' is not a base", baseName)));
197 }
198
199 const int32_t offset = offsets.front();
200 const uint8_t* baseLocator =
201 offset == 0 ? locator : FindLocator(rtti, typeDescriptor, static_cast<uint32_t>(offset));
202 return BaseSubobject{.Offset = offset, .Table = baseLocator ? TableAfter(rtti, baseLocator) : nullptr};
203}
204
205static ScanRange FindSection(const Image& module, std::string_view name)
206{
207 const auto* dos = reinterpret_cast<const IMAGE_DOS_HEADER*>(module.Base);
208 if (module.Size < sizeof(IMAGE_DOS_HEADER) || dos->e_magic != IMAGE_DOS_SIGNATURE)
209 {
210 return {};
211 }
212
213 const auto* nt = reinterpret_cast<const IMAGE_NT_HEADERS64*>(module.Base + dos->e_lfanew);
214 if (static_cast<size_t>(dos->e_lfanew) + sizeof(IMAGE_NT_HEADERS64) > module.Size ||
215 nt->Signature != IMAGE_NT_SIGNATURE)
216 {
217 return {};
218 }
219
221 for (WORD i = 0; i < nt->FileHeader.NumberOfSections; ++i)
222 {
223 // Section names are eight bytes, padded with NULs.
224 std::string_view actual(reinterpret_cast<const char*>(sections[i].Name), IMAGE_SIZEOF_SHORT_NAME);
225 if (const size_t padding = actual.find('\0'); padding != std::string_view::npos)
226 {
227 actual = actual.substr(0, padding);
228 }
229 if (actual != name)
230 {
231 continue;
232 }
233
234 // Use virtual size for the mapped image; raw size covers a zero virtual size.
235 const DWORD size = sections[i].Misc.VirtualSize ? sections[i].Misc.VirtualSize : sections[i].SizeOfRawData;
236 if (sections[i].VirtualAddress + static_cast<size_t>(size) > module.Size)
237 {
238 return {};
239 }
240
241 return {module.Base + sections[i].VirtualAddress, size};
242 }
243 return {};
244}
245
246static PeRtti RttiOf(const Image& module)
247{
248 return {.Base = module.Base,
249 .Size = module.Size,
250 .Data = FindSection(module, ".data"),
251 .ReadOnlyData = FindSection(module, ".rdata")};
252}
253
254void* FindVirtualTableIn(const Image& module, std::string_view className)
255{
256 const PeRtti rtti = RttiOf(module);
257 if (!rtti.Data.Base || !rtti.ReadOnlyData.Base)
258 {
259 return nullptr;
260 }
262}
263
264Result<BaseSubobject> FindBaseIn(const Image& module, std::string_view className, std::string_view baseName)
265{
266 const PeRtti rtti = RttiOf(module);
267 if (!rtti.Data.Base || !rtti.ReadOnlyData.Base)
268 {
269 return std::unexpected(Error::NotFound("the module has no RTTI sections"));
270 }
272}
273
274} // namespace VoltMod
static constexpr size_t BaseVirtualOffset
Result< BaseSubobject > FindBaseInRtti(const PeRtti &rtti, std::string_view className, std::string_view baseName)
static const uint8_t * AtRva(const PeRtti &rtti, int64_t rva, size_t bytes)
static const uint8_t * FindLocator(const PeRtti &rtti, uint32_t typeDescriptor, uint32_t offset)
static uint32_t FindTypeDescriptor(const PeRtti &rtti, std::string_view className)
static const uint8_t * End(const ScanRange &range)
static constexpr size_t LocatorTypeDescriptor
static bool TypeDescriptorMatches(const PeRtti &rtti, int32_t rva, std::span< const std::string > names)
static std::string ReadFile(const std::filesystem::path &path)
Definition Loader.cpp:56
static constexpr size_t LocatorHierarchy
static std::pair< uint32_t, const uint8_t * > PrimaryLocator(const PeRtti &rtti, std::string_view className)
static constexpr size_t BaseMemberOffset
Result< BaseSubobject > FindBaseIn(const Image &module, std::string_view className, std::string_view baseName)
static constexpr size_t HierarchyBaseList
static PeRtti RttiOf(const Image &module)
void * FindVirtualTableInRtti(const PeRtti &rtti, std::string_view className)
static constexpr size_t HierarchyBaseCount
static void * TableAfter(const PeRtti &rtti, const uint8_t *locator)
static constexpr uint32_t MaxBases
static ScanRange FindSection(const Image &module, std::string_view name)
static constexpr size_t LocatorSelf
static constexpr size_t TypeDescriptorName
static constexpr size_t HierarchySize
static constexpr size_t LocatorOffset
static constexpr size_t LocatorSize
std::expected< T, Error > Result
Definition Result.hpp:64
T ReadAt(void *base, std::ptrdiff_t offset) noexcept
static const uint8_t * FindValue(const uint8_t *begin, const uint8_t *end, const void *needle, size_t len, size_t stride)
static std::array< std::string, 2 > ClassAndStructNames(std::string_view name)
static constexpr size_t BaseSize
void * FindVirtualTableIn(const Image &module, std::string_view className)
static Error Invalid(std::string detail)
Definition Result.hpp:51
static Error Unsupported(std::string detail)
Definition Result.hpp:57
static Error NotFound(std::string detail)
Definition Result.hpp:49