#include "crypto.h"

#include <Windows.h>
#include <Wincrypt.h>

#include "utils.h"
#include "misc.h"

// 
// Enables RtlDecompressBufferEx implementation in theory supporting XPRESS and XPRESS Huffman
// compression algorithms. This implementation however currently doesn't work, so it was disabled.
//
//#define _SUPPORT_FOR_XPRESS_COMPRESSION


struct KeyBLOB16 {
	BLOBHEADER hdr;
	DWORD dwKeySize;
	BYTE keyData[16];
};

struct KeyBLOB32 {
	BLOBHEADER hdr;
	DWORD dwKeySize;
	BYTE keyData[32];
};

struct KeyBLOB {
	BLOBHEADER hdr;
	DWORD dwKeySize;
	BYTE keyData[1];
};

bool aes128_encrypt_cbc(
	const std::vector<uint8_t>& data,
	const std::vector<uint8_t>& key,
	std::vector<uint8_t>* output,
	const std::vector<uint8_t>& iv
)
{
	return aes_encrypt(data, key, output, CALG_AES_128, CRYPT_MODE_CBC, true, iv);
}

bool aes128_decrypt_cbc(
	const std::vector<uint8_t>& data,
	const std::vector<uint8_t>& key,
	std::vector<uint8_t>* output,
	const std::vector<uint8_t>& iv
)
{
	return aes_encrypt(data, key, output, CALG_AES_128, CRYPT_MODE_CBC, false, iv);
}

bool aes256_encrypt_cbc(
	const std::vector<uint8_t>& data,
	const std::vector<uint8_t>& key,
	std::vector<uint8_t>* output,
	const std::vector<uint8_t>& iv
)
{
	return aes_encrypt(data, key, output, CALG_AES_256, CRYPT_MODE_CBC, true, iv);
}

bool aes256_decrypt_cbc(
	const std::vector<uint8_t>& data,
	const std::vector<uint8_t>& key,
	std::vector<uint8_t>* output,
	const std::vector<uint8_t>& iv
)
{
	return aes_encrypt(data, key, output, CALG_AES_256, CRYPT_MODE_CBC, false, iv);
}

bool aes128_encrypt_ecb(
	const std::vector<uint8_t>& data,
	const std::vector<uint8_t>& key,
	std::vector<uint8_t>* output,
	const std::vector<uint8_t>& iv
)
{
	return aes_encrypt(data, key, output, CALG_AES_128, CRYPT_MODE_ECB, true, iv);
}

bool aes128_decrypt_ecb(
	const std::vector<uint8_t>& data,
	const std::vector<uint8_t>& key,
	std::vector<uint8_t>* output,
	const std::vector<uint8_t>& iv
)
{
	return aes_encrypt(data, key, output, CALG_AES_128, CRYPT_MODE_ECB, false, iv);
}

bool aes256_encrypt_ecb(
	const std::vector<uint8_t>& data,
	const std::vector<uint8_t>& key,
	std::vector<uint8_t>* output,
	const std::vector<uint8_t>& iv
)
{
	return aes_encrypt(data, key, output, CALG_AES_256, CRYPT_MODE_ECB, true, iv);
}

bool aes256_decrypt_ecb(
	const std::vector<uint8_t>& data,
	const std::vector<uint8_t>& key,
	std::vector<uint8_t>* output,
	const std::vector<uint8_t>& iv
)
{
	return aes_encrypt(data, key, output, CALG_AES_256, CRYPT_MODE_ECB, false, iv);
}

bool aes_encrypt(
	const std::vector<uint8_t>& data,
	const std::vector<uint8_t>& key,
	std::vector<uint8_t>* output,
	ALG_ID algoId,
	DWORD mode,
	bool encrypt,
	const std::vector<uint8_t>& iv
)
{
	HCRYPTPROV hProvider = NULL;
	HCRYPTKEY hKey;

	RESOLVE(Advapi32, CryptAcquireContextW);
	RESOLVE(Advapi32, CryptDestroyKey);
	RESOLVE(Advapi32, CryptReleaseContext);

	LPCWSTR prov = nullptr;
	if (algoId == CALG_AES_256) prov = OBFI(MS_ENH_RSA_AES_PROV_W);

	if (!_CryptAcquireContextW(
		&hProvider,
		NULL,  // pszContainer = no named container
		prov,
		PROV_RSA_AES,
		CRYPT_VERIFYCONTEXT
	)) {
		info(OBF(L"[!] Could not acquire Crypto context! Error: "), std::hex, GetLastError());
		return false;
	}

	KeyBLOB *kb = nullptr;
	DWORD kbSize = 0;

	BLOBHEADER hdr = { 0 };
	hdr.bType = PLAINTEXTKEYBLOB;
	hdr.bVersion = CUR_BLOB_VERSION;
	hdr.reserved = 0;
	hdr.aiKeyAlg = algoId;

	if (algoId == CALG_AES_128)
	{
		KeyBLOB16 kb16 = { 0 };
		kb16.dwKeySize = sizeof(kb16.keyData);
		memcpy(&kb16.hdr, &hdr, sizeof(hdr));
		kb = reinterpret_cast<KeyBLOB*>(&kb16);
		kbSize = sizeof(kb16);
	}
	else if (algoId == CALG_AES_256)
	{
		KeyBLOB32 kb32 = { 0 };
		kb32.dwKeySize = sizeof(kb32.keyData);
		memcpy(&kb32.hdr, &hdr, sizeof(hdr));
		kb = reinterpret_cast<KeyBLOB*>(&kb32);
		kbSize = sizeof(kb32);
	}
	else 
	{
		info(OBF(L"[!] Algorithm not supported by this implementation."));
		return false;
	}

	memcpy(
		kb->keyData, 
		key.data(), 
		min(key.size(), kb->dwKeySize)
	);

	RESOLVE(Advapi32, CryptImportKey);
	if (!_CryptImportKey(
		hProvider,
		(BYTE*)kb,
		kbSize,
		0,
		CRYPT_EXPORTABLE,
		&hKey
	))
	{
		info(OBF(L"[!] Could not import crypto key structure! Error: "), std::hex, GetLastError());
		_CryptReleaseContext(hProvider, 0);
		return false;
	}

	BYTE _iv[16] = { 0 };
	if (iv.size() > 0)
	{
		memcpy(_iv, iv.data(), min(sizeof(_iv), iv.size()));
	}

	RESOLVE(Advapi32, CryptSetKeyParam);
	if (!_CryptSetKeyParam(
		hKey,
		KP_IV,
		_iv,
		0
	) || 
		!_CryptSetKeyParam(
		hKey,
		KP_MODE, 
		(BYTE*)&mode, 
		0
	))
	{
		info(OBF(L"[!] Unable to change algorithm's encryption/decryption mode or IV! Error: "), std::hex, GetLastError());
		_CryptDestroyKey(hKey);
		_CryptReleaseContext(hProvider, 0);
		return false;
	}


	// The CryptEncrypt method uses the *same* buffer for both the input and
	// output (!), so we copy the data to be encrypted into the output array.
	// Also, for some reason, the AES-128 block cipher on Windows requires twice
	// the block size in the output buffer. So we resize it to that length and
	// then chop off the excess after we are done.
	output->clear();
	output->resize(data.size() * 2);
	memcpy(output->data(), &data[0], data.size());

	// This acts as both the length of bytes to be encoded (on input) and the
	// number of bytes used in the resulting encrypted data (on output).
	DWORD length = static_cast<DWORD>(data.size());

	if (encrypt)
	{
		RESOLVE(Advapi32, CryptEncrypt);
		if (!_CryptEncrypt(
			hKey,
			NULL,  // hHash = no hash
			true,  // Final
			0,     // dwFlags
			reinterpret_cast<BYTE*>(output->data()),
			&length,
			static_cast<DWORD>(output->size())
		)) {
			info(OBF(L"[!] AES Encryption failed. Error: "), std::hex, GetLastError());
			_CryptDestroyKey(hKey);
			_CryptReleaseContext(hProvider, 0);
			return false;
		}
	}
	else
	{
		RESOLVE(Advapi32, CryptDecrypt);
		if (!_CryptDecrypt(
			hKey,
			NULL,  // hHash = no hash
			true,  // Final
			0,     // dwFlags
			reinterpret_cast<BYTE*>(output->data()),
			&length
		)) {
			info(OBF(L"[!] AES Decryption failed. Error: "), std::hex, GetLastError());
			_CryptDestroyKey(hKey);
			_CryptReleaseContext(hProvider, 0);
			return false;
		}
	}

	output->resize(length);

	_CryptDestroyKey(hKey);
	_CryptReleaseContext(hProvider, 0);
	
	return true;
}

//
// More infos about shellcode compression can be found here:
//		https://modexp.wordpress.com/2019/12/08/shellcode-compression/
//
bool decompressBuffer(
	std::vector<uint8_t>& buffer, 
	std::vector<uint8_t>& output, 
	PayloadCompression algo
)
{
	if (algo == PayloadCompression::NoCompression)
	{
		if(output.empty()) output.resize(buffer.size());
		memcpy(output.data(), buffer.data(), buffer.size());
		return true;
	}

	const size_t Multiplier = 15;
	if (output.empty()) output.resize(Multiplier * buffer.size());

	USHORT algorithm = 0;
	std::wstring algostring;
	ULONG finalUncompressedSize = 0;

#ifdef _SUPPORT_FOR_XPRESS_COMPRESSION
	if (algo == PayloadCompression::Xpress)
	{
		algorithm = COMPRESSION_FORMAT_XPRESS | COMPRESSION_ENGINE_MAXIMUM;
		//algorithm = COMPRESS_ALGORITHM_XPRESS;
		algostring = OBF_WSTR(L"XPress");
		//output.resize(Multiplier * buffer.size());
	}
	else if (algo == PayloadCompression::XpressHuffman)
	{
		algorithm = COMPRESSION_FORMAT_XPRESS_HUFF | COMPRESSION_ENGINE_MAXIMUM;
		//algorithm = COMPRESS_ALGORITHM_XPRESS_HUFF;
		algostring = OBF_WSTR(L"XPress with Huffman");
		//output.resize(Multiplier * buffer.size());
	}
#endif

	if (algo == PayloadCompression::Lznt1)
	{
		algorithm = COMPRESSION_FORMAT_LZNT1 | COMPRESSION_ENGINE_MAXIMUM;
		algostring = OBF_WSTR(L"LZnt1");
	}

	NTSTATUS out = 0;

#ifdef _SUPPORT_FOR_XPRESS_COMPRESSION
	if (algo == PayloadCompression::XpressHuffman || algo == PayloadCompression::Xpress)
	{
		ULONG compressBufferWorkSpaceSize = 0,
			compressBufferFragmentWorkSpaceSize = 0;

		RESOLVE(ntdll, RtlGetCompressionWorkSpaceSize);
		RESOLVE(ntdll, RtlDecompressBufferEx);

		out = _RtlGetCompressionWorkSpaceSize(
			algorithm,
			&compressBufferWorkSpaceSize,
			&compressBufferFragmentWorkSpaceSize
		);

		if (out != STATUS_SUCCESS)
		{
			info(OBF(L"[!] Could not get compression workspace size! NTSTATUS: "), std::hex, out);
			return false;
		}

		uint8_t* workspace = new uint8_t[compressBufferWorkSpaceSize + 1];
		memset(workspace, 0, compressBufferWorkSpaceSize + 1);

		out = _RtlDecompressBufferEx(
			algorithm,
			reinterpret_cast<PUCHAR>(output.data()),
			static_cast<ULONG>(output.size()),
			reinterpret_cast<PUCHAR>(buffer.data()),
			static_cast<ULONG>(buffer.size()),
			&finalUncompressedSize,
			workspace
		);

		delete [] workspace;

		/*
		DECOMPRESSOR_HANDLE dh = NULL;

		RESOLVE(cabinet, CreateDecompressor);
		RESOLVE(cabinet, CloseDecompressor);

		//algorithm |= COMPRESS_RAW;

		if (!_CreateDecompressor(
			algorithm,
			NULL,
			&dh
		))
		{
			info(OBF(L"[!] Could not create decompressor. Error: "), std::hex, GetLastError());
			unloadModule(OBFI(L"cabinet"));
			return false;
		}

		size_t uncompressedRequiredBytes = 0;

		RESOLVE(cabinet, Decompress);
		_Decompress(
			dh, 
			buffer.data(), 
			buffer.size(), 
			NULL, 
			0, 
			&uncompressedRequiredBytes
		);

		auto err = GetLastError();
		if(err != ERROR_INSUFFICIENT_BUFFER)
		{
			info(OBF(L"[!] Decompressor failed unexpectedly. Error: "), std::hex, GetLastError());
			_CloseDecompressor(dh);
			unloadModule(OBFI(L"cabinet"));
			return false;
		}

		output.resize(uncompressedRequiredBytes);

		if(!_Decompress(
			dh,
			buffer.data(),
			buffer.size(),
			output.data(),
			output.size(),
			reinterpret_cast<PSIZE_T>(&finalUncompressedSize)
		))
		{
			info(OBF(L"[!] Decompressor was not able to decompressor our buffer. Error: "), std::hex, GetLastError());
			_CloseDecompressor(dh);
			unloadModule(OBFI(L"cabinet"));
			return false;
		}

		_CloseDecompressor(dh);
		unloadModule(OBFI(L"cabinet"));
		*/
	}
	else
#endif
	{
		RESOLVE(ntdll, RtlDecompressBuffer);
		out = _RtlDecompressBuffer(
			algorithm,
			reinterpret_cast<PUCHAR>(output.data()),
			static_cast<ULONG>(output.size()),
			reinterpret_cast<PUCHAR>(buffer.data()),
			static_cast<ULONG>(buffer.size()),
			&finalUncompressedSize
		);
	}

	if (out != STATUS_SUCCESS)
	{
		info(OBF(L"[!] Could not decompress the buffer! NTSTATUS: "), std::hex, out);
		return false;
	}

	if (!finalUncompressedSize)
	{
		return false;
	}

	output.resize(finalUncompressedSize);
	verbose(OBF(L"[+] Decompressed buffer with "), algostring, OBF(L". Initial size: "), buffer.size(),
		OBF(L", after decompression: "), finalUncompressedSize);

	return output.size() >= (0.01 * buffer.size());
}

void xor8(
	uint8_t* buf,
	size_t bufSize,
	uint32_t xorKey
)
{
	for (size_t i = 0; i < bufSize; i++)
	{
		buf[i] ^= static_cast<uint8_t>(xorKey & 0xff);
	}
}

void xor32(
	uint8_t* buf, 
	size_t bufSize, 
	uint32_t xorKey
)
{
	uint32_t* buf32 = reinterpret_cast<uint32_t*>(buf);

	auto bufSizeRounded = (bufSize - (bufSize % sizeof(uint32_t))) / 4;
	for (size_t i = 0; i < bufSizeRounded; i++)
	{
		buf32[i] ^= xorKey;
	}

	for (size_t i = 4 * bufSizeRounded; i < bufSize; i++)
	{
		buf[i] ^= static_cast<uint8_t>(xorKey & 0xff);
	}
}

void RC4Encoder::swap(uint8_t* a, uint8_t* b)
{
    uint8_t tmp = *a;
    *a = *b;
    *b = tmp;
}

void RC4Encoder::KSA(const bytes_t& key, uint8_t* S)
{
    size_t len = key.size();
    size_t j = 0;

    for (size_t i = 0; i < N; i++)
    {
        S[i] = (uint8_t)i;
    }

    for (size_t i = 0; i < N; i++) {
        j = (j + S[i] + key[i % len]) % N;

        swap(&S[i], &S[j]);
    }
}

void RC4Encoder::PRGA(uint8_t* S, const bytes_t& plaintext, bytes_t& ciphertext)
{
    size_t i = 0;
    size_t j = 0;

    for (size_t n = 0, len = plaintext.size(); n < len; n++)
    {
        i = (i + 1) % N;
        j = (j + S[i]) % N;

        swap(&S[i], &S[j]);
        size_t rnd = S[((size_t)S[i] + (size_t)S[j]) % N];

        ciphertext[n] = uint8_t((uint8_t)rnd ^ (uint8_t)plaintext[n]);
    }
}

RC4Encoder::bytes_t RC4Encoder::RC4(const bytes_t& key, const bytes_t& plaintext)
{
    uint8_t S[N];
    bytes_t ciphertext(plaintext.size() + 1, 0);

    KSA(key, S);
    PRGA(S, plaintext, ciphertext);

    return ciphertext;
}