Files
2026-07-10 11:56:50 -07:00

312 lines
11 KiB
C++

#include "aes_crypto.hpp"
#include "../resolve/api_resolver.hpp"
#include "../common/hashes.hpp"
#include "../common/nt_types.hpp"
#include "../common/log.hpp"
#include <bcrypt.h>
#include <fstream>
#include <cstring>
namespace AesCrypto {
namespace {
using BCryptOpenAlgorithmProvider_t = NTSTATUS(WINAPI*)(BCRYPT_ALG_HANDLE*, LPCWSTR, LPCWSTR, ULONG);
using BCryptCloseAlgorithmProvider_t = NTSTATUS(WINAPI*)(BCRYPT_ALG_HANDLE, ULONG);
using BCryptGenerateSymmetricKey_t = NTSTATUS(WINAPI*)(BCRYPT_ALG_HANDLE, BCRYPT_KEY_HANDLE*, PUCHAR, ULONG, PUCHAR, ULONG, ULONG);
using BCryptDestroyKey_t = NTSTATUS(WINAPI*)(BCRYPT_KEY_HANDLE);
using BCryptEncrypt_t = NTSTATUS(WINAPI*)(BCRYPT_KEY_HANDLE, PUCHAR, ULONG, VOID*, PUCHAR, ULONG, PUCHAR, ULONG, ULONG*, ULONG);
using BCryptDecrypt_t = NTSTATUS(WINAPI*)(BCRYPT_KEY_HANDLE, PUCHAR, ULONG, VOID*, PUCHAR, ULONG, PUCHAR, ULONG, ULONG*, ULONG);
using BCryptSetProperty_t = NTSTATUS(WINAPI*)(BCRYPT_HANDLE, LPCWSTR, PUCHAR, ULONG, ULONG);
using BCryptGenRandom_t = NTSTATUS(WINAPI*)(BCRYPT_ALG_HANDLE, PUCHAR, ULONG, ULONG);
struct BcryptApi {
BCryptOpenAlgorithmProvider_t Open = nullptr;
BCryptCloseAlgorithmProvider_t Close = nullptr;
BCryptGenerateSymmetricKey_t GenKey = nullptr;
BCryptDestroyKey_t DestroyKey = nullptr;
BCryptEncrypt_t Encrypt = nullptr;
BCryptDecrypt_t Decrypt = nullptr;
BCryptSetProperty_t SetProperty = nullptr;
BCryptGenRandom_t GenRandom = nullptr;
bool loaded = false;
};
bool LoadBcrypt(BcryptApi& api) {
if (api.loaded) {
return true;
}
api.Open = reinterpret_cast<BCryptOpenAlgorithmProvider_t>(
ApiResolver::ResolveBcrypt(Hashes::H_BCryptOpenAlgorithmProvider));
api.Close = reinterpret_cast<BCryptCloseAlgorithmProvider_t>(
ApiResolver::ResolveBcrypt(Hashes::H_BCryptCloseAlgorithmProvider));
api.GenKey = reinterpret_cast<BCryptGenerateSymmetricKey_t>(
ApiResolver::ResolveBcrypt(Hashes::H_BCryptGenerateSymmetricKey));
api.DestroyKey = reinterpret_cast<BCryptDestroyKey_t>(
ApiResolver::ResolveBcrypt(Hashes::H_BCryptDestroyKey));
api.Encrypt = reinterpret_cast<BCryptEncrypt_t>(
ApiResolver::ResolveBcrypt(Hashes::H_BCryptEncrypt));
api.Decrypt = reinterpret_cast<BCryptDecrypt_t>(
ApiResolver::ResolveBcrypt(Hashes::H_BCryptDecrypt));
api.SetProperty = reinterpret_cast<BCryptSetProperty_t>(
ApiResolver::ResolveBcrypt(Hashes::H_BCryptSetProperty));
// BCryptGenRandom is in bcrypt.dll but not in our hash table — resolve by adding hash
api.GenRandom = reinterpret_cast<BCryptGenRandom_t>(
ApiResolver::ResolveBcrypt(Hashes::H_BCryptGenRandom));
api.loaded = api.Open && api.Close && api.GenKey && api.DestroyKey &&
api.Encrypt && api.Decrypt && api.SetProperty && api.GenRandom;
return api.loaded;
}
} // namespace
bool GenerateKeyMaterial(KeyMaterial& out) {
Log::Dbg("AES: GenerateKeyMaterial (BCryptGenRandom system RNG)");
BcryptApi api{};
if (!LoadBcrypt(api)) {
Log::Dbg("AES: bcrypt API resolve failed");
return false;
}
out.key.resize(32);
out.iv.resize(16);
// BCryptGenRandom with NULL handle requires BCRYPT_USE_SYSTEM_PREFERRED_RNG on Win8+.
constexpr ULONG kSystemRng = 0x00000002;
if (!NT_SUCCESS(api.GenRandom(nullptr, out.key.data(), static_cast<ULONG>(out.key.size()), kSystemRng)) ||
!NT_SUCCESS(api.GenRandom(nullptr, out.iv.data(), static_cast<ULONG>(out.iv.size()), kSystemRng))) {
Log::Dbg("AES: GenRandom failed");
return false;
}
Log::Dbg("AES: key=%zu bytes iv=%zu bytes", out.key.size(), out.iv.size());
return true;
}
bool Encrypt(const std::vector<BYTE>& plaintext, const KeyMaterial& km, std::vector<BYTE>& ciphertext) {
Log::Dbg("AES Encrypt: plain=%zu key=%zu iv=%zu mode=AES-256-CBC",
plaintext.size(), km.key.size(), km.iv.size());
if (km.key.size() != 32 || km.iv.size() != 16 || plaintext.empty()) {
Log::Dbg("AES Encrypt: bad args");
return false;
}
BcryptApi api{};
if (!LoadBcrypt(api)) {
Log::Dbg("AES Encrypt: bcrypt API resolve failed");
return false;
}
BCRYPT_ALG_HANDLE hAlg = nullptr;
if (!NT_SUCCESS(api.Open(&hAlg, BCRYPT_AES_ALGORITHM, nullptr, 0))) {
return false;
}
api.SetProperty(hAlg, BCRYPT_CHAINING_MODE, (PUCHAR)BCRYPT_CHAIN_MODE_CBC,
sizeof(BCRYPT_CHAIN_MODE_CBC), 0);
BCRYPT_KEY_HANDLE hKey = nullptr;
if (!NT_SUCCESS(api.GenKey(hAlg, &hKey, nullptr, 0,
const_cast<PUCHAR>(km.key.data()), static_cast<ULONG>(km.key.size()), 0))) {
api.Close(hAlg, 0);
return false;
}
std::vector<BYTE> ivWork = km.iv;
ULONG cipherLen = 0;
api.Encrypt(hKey, const_cast<PUCHAR>(plaintext.data()), static_cast<ULONG>(plaintext.size()),
nullptr, ivWork.data(), static_cast<ULONG>(ivWork.size()),
nullptr, 0, &cipherLen, BCRYPT_BLOCK_PADDING);
ciphertext.resize(cipherLen);
ULONG written = 0;
NTSTATUS st = api.Encrypt(
hKey,
const_cast<PUCHAR>(plaintext.data()), static_cast<ULONG>(plaintext.size()),
nullptr,
ivWork.data(), static_cast<ULONG>(ivWork.size()),
ciphertext.data(), static_cast<ULONG>(ciphertext.size()),
&written, BCRYPT_BLOCK_PADDING);
api.DestroyKey(hKey);
api.Close(hAlg, 0);
if (!NT_SUCCESS(st)) {
Log::Dbg("AES Encrypt: BCryptEncrypt failed NTSTATUS=0x%08lX",
static_cast<unsigned long>(st));
return false;
}
ciphertext.resize(written);
Log::Dbg("AES Encrypt OK: cipher=%zu bytes", ciphertext.size());
return true;
}
bool Decrypt(const std::vector<BYTE>& ciphertext, const KeyMaterial& km, std::vector<BYTE>& plaintext) {
Log::Dbg("AES Decrypt: cipher=%zu key=%zu iv=%zu mode=AES-256-CBC",
ciphertext.size(), km.key.size(), km.iv.size());
if (km.key.size() != 32 || km.iv.size() != 16 || ciphertext.empty()) {
Log::Dbg("AES Decrypt: bad args");
return false;
}
BcryptApi api{};
if (!LoadBcrypt(api)) {
Log::Dbg("AES Decrypt: bcrypt API resolve failed");
return false;
}
BCRYPT_ALG_HANDLE hAlg = nullptr;
if (!NT_SUCCESS(api.Open(&hAlg, BCRYPT_AES_ALGORITHM, nullptr, 0))) {
return false;
}
api.SetProperty(hAlg, BCRYPT_CHAINING_MODE, (PUCHAR)BCRYPT_CHAIN_MODE_CBC,
sizeof(BCRYPT_CHAIN_MODE_CBC), 0);
BCRYPT_KEY_HANDLE hKey = nullptr;
if (!NT_SUCCESS(api.GenKey(hAlg, &hKey, nullptr, 0,
const_cast<PUCHAR>(km.key.data()), static_cast<ULONG>(km.key.size()), 0))) {
api.Close(hAlg, 0);
return false;
}
std::vector<BYTE> ivWork = km.iv;
ULONG plainLen = 0;
api.Decrypt(hKey, const_cast<PUCHAR>(ciphertext.data()), static_cast<ULONG>(ciphertext.size()),
nullptr, ivWork.data(), static_cast<ULONG>(ivWork.size()),
nullptr, 0, &plainLen, BCRYPT_BLOCK_PADDING);
plaintext.resize(plainLen);
ULONG written = 0;
NTSTATUS st = api.Decrypt(
hKey,
const_cast<PUCHAR>(ciphertext.data()), static_cast<ULONG>(ciphertext.size()),
nullptr,
ivWork.data(), static_cast<ULONG>(ivWork.size()),
plaintext.data(), static_cast<ULONG>(plaintext.size()),
&written, BCRYPT_BLOCK_PADDING);
api.DestroyKey(hKey);
api.Close(hAlg, 0);
if (!NT_SUCCESS(st)) {
Log::Dbg("AES Decrypt: BCryptDecrypt failed NTSTATUS=0x%08lX",
static_cast<unsigned long>(st));
return false;
}
plaintext.resize(written);
Log::Dbg("AES Decrypt OK: plain=%zu bytes", plaintext.size());
return true;
}
std::vector<BYTE> PackKeyMaterial(const KeyMaterial& km) {
std::vector<BYTE> blob;
auto appendU32 = [&](DWORD v) {
blob.push_back(static_cast<BYTE>(v & 0xFF));
blob.push_back(static_cast<BYTE>((v >> 8) & 0xFF));
blob.push_back(static_cast<BYTE>((v >> 16) & 0xFF));
blob.push_back(static_cast<BYTE>((v >> 24) & 0xFF));
};
appendU32(static_cast<DWORD>(km.key.size()));
blob.insert(blob.end(), km.key.begin(), km.key.end());
appendU32(static_cast<DWORD>(km.iv.size()));
blob.insert(blob.end(), km.iv.begin(), km.iv.end());
return blob;
}
bool UnpackKeyMaterial(const std::vector<BYTE>& blob, KeyMaterial& out) {
if (blob.size() < 8) {
return false;
}
size_t off = 0;
auto readU32 = [&]() -> DWORD {
DWORD v = blob[off] | (blob[off + 1] << 8) | (blob[off + 2] << 16) | (blob[off + 3] << 24);
off += 4;
return v;
};
const DWORD keyLen = readU32();
if (off + keyLen > blob.size()) {
return false;
}
out.key.assign(blob.begin() + off, blob.begin() + off + keyLen);
off += keyLen;
if (off + 4 > blob.size()) {
return false;
}
const DWORD ivLen = readU32();
if (off + ivLen > blob.size()) {
return false;
}
out.iv.assign(blob.begin() + off, blob.begin() + off + ivLen);
return out.key.size() == 32 && out.iv.size() == 16;
}
bool EncryptFileToDisk(const std::wstring& inputPath, std::wstring& outEncPath, std::wstring& outKeyPath) {
std::ifstream in(inputPath, std::ios::binary | std::ios::ate);
if (!in.is_open()) {
return false;
}
const auto size = in.tellg();
in.seekg(0);
std::vector<BYTE> plain(static_cast<size_t>(size));
if (!in.read(reinterpret_cast<char*>(plain.data()), size)) {
return false;
}
in.close();
KeyMaterial km{};
if (!GenerateKeyMaterial(km)) {
return false;
}
std::vector<BYTE> cipher;
if (!Encrypt(plain, km, cipher)) {
return false;
}
outEncPath = inputPath + L".enc";
outKeyPath = inputPath + L".key";
std::ofstream enc(outEncPath, std::ios::binary);
std::ofstream key(outKeyPath, std::ios::binary);
if (!enc.is_open() || !key.is_open()) {
return false;
}
enc.write(reinterpret_cast<const char*>(cipher.data()), cipher.size());
auto packed = PackKeyMaterial(km);
key.write(reinterpret_cast<const char*>(packed.data()), packed.size());
return true;
}
bool ReadEncryptedPayload(const std::wstring& encPath, const std::wstring& keyPath,
std::vector<BYTE>& plaintext) {
std::ifstream enc(encPath, std::ios::binary | std::ios::ate);
std::ifstream key(keyPath, std::ios::binary | std::ios::ate);
if (!enc.is_open() || !key.is_open()) {
return false;
}
std::vector<BYTE> cipher(static_cast<size_t>(enc.tellg()));
std::vector<BYTE> packed(static_cast<size_t>(key.tellg()));
enc.seekg(0);
key.seekg(0);
enc.read(reinterpret_cast<char*>(cipher.data()), cipher.size());
key.read(reinterpret_cast<char*>(packed.data()), packed.size());
KeyMaterial km{};
if (!UnpackKeyMaterial(packed, km)) {
return false;
}
return Decrypt(cipher, km, plaintext);
}
} // namespace AesCrypto