Files

187 lines
5.5 KiB
C++

#include "util.h"
namespace util {
bool HasCetEnabled(int process_id) {
PROCESS_MITIGATION_USER_SHADOW_STACK_POLICY cet{};
bool result = true;
auto hProcess = OpenProcess(PROCESS_QUERY_LIMITED_INFORMATION, false, process_id);
if (hProcess == INVALID_HANDLE_VALUE)
return false;
if (!GetProcessMitigationPolicy(hProcess, ProcessUserShadowStackPolicy, &cet, sizeof(PROCESS_MITIGATION_USER_SHADOW_STACK_POLICY)))
result = false;
if (cet.EnableUserShadowStack)
result = true;
CloseHandle(hProcess);
return result;
}
std::unordered_map<unsigned int, Process> GetCETProcesses() {
std::unordered_map<unsigned int, Process> processes_;
std::unordered_set<unsigned int> rejected_;
THREADENTRY32 thread_entry{};
thread_entry.dwSize = sizeof(THREADENTRY32);
auto snapshot = CreateToolhelp32Snapshot(TH32CS_SNAPTHREAD, 0);
if (snapshot == INVALID_HANDLE_VALUE)
throw std::runtime_error(std::format("Failed to get snapshot: {}", GetLastError()));
if (!Thread32First(snapshot, &thread_entry)) {
CloseHandle(snapshot);
throw std::runtime_error(std::format("Failed to get first thread: {}", GetLastError()));
}
do {
auto pid = thread_entry.th32OwnerProcessID;
auto tid = thread_entry.th32ThreadID;
if (rejected_.contains(pid))
continue;
auto [it, inserted] = processes_.try_emplace(pid, Process{ .processId = pid });
if (inserted && !HasCetEnabled(pid)) {
processes_.erase(it);
rejected_.insert(pid);
continue;
}
it->second.threads.push_back(Thread{ .processId = pid, .threadId = tid });
} while (Thread32Next(snapshot, &thread_entry));
CloseHandle(snapshot);
return processes_;
}
ShadowStackResult GetShadowStackFrames(HANDLE thread_handle, HANDLE process) {
ShadowStackResult result;
result.HasShadowStack = false;
auto features = GetEnabledXStateFeatures();
if (!(features & XSTATE_MASK_CET_U))
throw std::runtime_error("Required CET features not available.");
DWORD ctx_size = 0;
InitializeContext2(nullptr, CONTEXT_FULL | CONTEXT_XSTATE, nullptr, &ctx_size, XSTATE_MASK_CET_U);
auto ctx_buf = std::vector<uint8_t>(ctx_size);
CONTEXT* ctx = nullptr;
if (!InitializeContext2(ctx_buf.data(), CONTEXT_FULL | CONTEXT_XSTATE, &ctx, &ctx_size, XSTATE_MASK_CET_U)) {
//ERROR
return result;
}
if (!SetXStateFeaturesMask(ctx, XSTATE_MASK_CET_U)) {
//error
return result;
}
if (!GetThreadContext(thread_handle, ctx)) {
return result;
}
uint64_t mask = 0;
GetXStateFeaturesMask(ctx, &mask);
if (!(mask & XSTATE_MASK_CET_U)) {
// CET_U wasn't populated - thread has no active shadow stack
return result;
}
DWORD feature_len = 0;
auto* cet_state = static_cast<uint64_t*>(LocateXStateFeature(ctx, XSTATE_CET_U, &feature_len));
if (!cet_state || feature_len < sizeof(uint64_t)) {
return result;
}
auto ssp = cet_state[1];
if (ssp == 0) {
return result;
}
result.HasShadowStack = true;
result.frames.reserve(64);
MEMORY_BASIC_INFORMATION mbi{};
VirtualQueryEx(process, reinterpret_cast<void*>(ssp), &mbi, sizeof(mbi));
for (size_t i = 0; i < 64; ++i) {
auto addr = ssp + (i * sizeof(uint64_t));
uint64_t ret = 0;
size_t bytesread = 0;
if (!ReadProcessMemory(process, reinterpret_cast<void*>(addr), &ret, sizeof(ret), &bytesread))
break;
if (ret == 0)
break;
// check if it's a restore token.
if ( (ret & 1) && (((ret & ~1ULL) >= reinterpret_cast<uint64_t>(mbi.AllocationBase) )
&& (ret & ~1ULL) < (reinterpret_cast<uint64_t>(mbi.AllocationBase) + mbi.RegionSize)) )
continue;
result.frames.push_back(ret);
}
return result;
}
std::vector<uintptr_t> GetNormalFrames(HANDLE thread_handle, HANDLE process) {
std::vector<uintptr_t> frames;
CONTEXT ctx{};
ctx.ContextFlags = CONTEXT_FULL;
if (!GetThreadContext(thread_handle, &ctx)) {
return {};
}
STACKFRAME_EX sf{};
sf.StackFrameSize = sizeof(sf);
sf.AddrPC.Offset = ctx.Rip;
sf.AddrPC.Mode = AddrModeFlat;
sf.AddrFrame.Offset = ctx.Rbp;
sf.AddrFrame.Mode = AddrModeFlat;
sf.AddrStack.Offset = ctx.Rsp;
sf.AddrStack.Mode = AddrModeFlat;
for (size_t i = 0; i < 64; ++i) {
auto result = StackWalkEx(
IMAGE_FILE_MACHINE_AMD64,
process,
thread_handle,
&sf,
&ctx,
nullptr,
SymFunctionTableAccess64,
SymGetModuleBase64,
nullptr,
SYM_STKWALK_DEFAULT
);
if (!sf.AddrPC.Offset || sf.AddrPC.Offset == UINT64_MAX)
break;
if (SymGetModuleBase64(process, sf.AddrPC.Offset) == 0)
break;
frames.push_back(sf.AddrPC.Offset);
}
return frames;
}
}