#include "memory/memory_reader.hpp"
#include <iostream>

MemoryReader::~MemoryReader() {
    Detach();
}

bool MemoryReader::Attach(uint32_t pid) {
    Detach();

    m_handle = OpenProcess(
        PROCESS_QUERY_LIMITED_INFORMATION | PROCESS_VM_READ | PROCESS_VM_WRITE | PROCESS_VM_OPERATION,
        FALSE,
        pid);

    if (!m_handle) {
        m_handle = OpenProcess(
            PROCESS_QUERY_LIMITED_INFORMATION | PROCESS_VM_READ,
            FALSE,
            pid);
    }

    if (!m_handle) {
        std::cerr << "[Memory] OpenProcess(" << pid << ") that bai voi loi "
                  << GetLastError() << " (co the do UAC/High Integrity Token)." << std::endl;
        return false;
    }

    m_pid = pid;
    std::cout << "[Memory] Da gan thanh cong vao PID " << pid
              << " qua " << BackendName() << std::endl;
    return true;
}

void MemoryReader::Detach() {
    if (m_handle) {
        CloseHandle(m_handle);
        m_handle = nullptr;
    }
    m_pid = 0;
}

bool MemoryReader::Read(uintptr_t address, void* buffer, size_t size) {
    if (!m_handle || !buffer || size == 0) return false;

    SIZE_T bytesRead = 0;
    if (!ReadProcessMemory(m_handle,
                           reinterpret_cast<LPCVOID>(address),
                           buffer,
                           size,
                           &bytesRead)) {
        return false;
    }
    return bytesRead == size;
}

bool MemoryReader::Write(uintptr_t address, const void* buffer, size_t size) {
    if (!m_handle || !buffer || size == 0) return false;
    if (address < 0x10000 || address > 0x7FFFFFFFFFFF) return false;

    SIZE_T bytesWritten = 0;
    if (!WriteProcessMemory(m_handle,
                            reinterpret_cast<LPVOID>(address),
                            buffer,
                            size,
                            &bytesWritten)) {
        DWORD oldProtect = 0;
        if (VirtualProtectEx(m_handle, reinterpret_cast<LPVOID>(address), size, PAGE_EXECUTE_READWRITE, &oldProtect)) {
            BOOL ok = WriteProcessMemory(m_handle,
                                         reinterpret_cast<LPVOID>(address),
                                         buffer,
                                         size,
                                         &bytesWritten);
            DWORD temp = 0;
            VirtualProtectEx(m_handle, reinterpret_cast<LPVOID>(address), size, oldProtect, &temp);
            FlushInstructionCache(m_handle, reinterpret_cast<LPCVOID>(address), size);
            return (ok && bytesWritten == size);
        }
        return false;
    }
    return bytesWritten == size;
}

bool MemoryReader::ForEachReadableRegion(const RegionFn& fn) {
    if (!m_handle) return false;

    uintptr_t addr = 0x10000;
    MEMORY_BASIC_INFORMATION mbi{};
    while (VirtualQueryEx(m_handle, reinterpret_cast<LPCVOID>(addr),
                          &mbi, sizeof(mbi)) != 0) {
        const uintptr_t regionBase = reinterpret_cast<uintptr_t>(mbi.BaseAddress);
        const bool readable = (mbi.State == MEM_COMMIT)
            && !(mbi.Protect & (PAGE_NOACCESS | PAGE_GUARD))
            && mbi.Protect != 0;
        if (readable && mbi.RegionSize <= (512u << 20)) {
            if (!fn(regionBase, static_cast<size_t>(mbi.RegionSize))) return true;
        }
        const uintptr_t next = regionBase + mbi.RegionSize;
        if (next <= regionBase) break;
        addr = next;
    }
    return true;
}

bool MemoryReader::ForEachHeapRegion(const RegionFn& fn) {
    if (!m_handle) return false;

    uintptr_t addr = 0x10000ULL;
    MEMORY_BASIC_INFORMATION mbi{};
    while (VirtualQueryEx(m_handle, reinterpret_cast<LPCVOID>(addr),
                          &mbi, sizeof(mbi)) != 0) {
        const uintptr_t regionBase = reinterpret_cast<uintptr_t>(mbi.BaseAddress);
        const bool match = (mbi.State == MEM_COMMIT)
            && (mbi.Type == MEM_PRIVATE)
            && ((mbi.Protect & PAGE_READWRITE) != 0)
            && (mbi.RegionSize >= (4u << 10))
            && (mbi.RegionSize <= (256u << 20));
        if (match) {
            if (!fn(regionBase, static_cast<size_t>(mbi.RegionSize))) return true;
        }
        const uintptr_t next = regionBase + mbi.RegionSize;
        if (next <= regionBase) break;
        addr = next;
    }
    return true;
}
