#include "memory/aob_scanner.hpp"

#include <psapi.h>
#include <algorithm>
#include <cctype>
#include <cstring>

intptr_t AobScanner::Find(const uint8_t* data, size_t dataSize, const AobPattern& pattern) {
    if (!pattern.Valid() || !data || dataSize < pattern.Size()) return -1;

    const size_t pSize = pattern.Size();
    const uint8_t firstByte = pattern.bytes[0];
    const bool firstWild = pattern.wildcard[0];

    for (size_t i = 0; i + pSize <= dataSize; ++i) {
        if (!firstWild && data[i] != firstByte) continue;

        bool matched = true;
        for (size_t j = 0; j < pSize; ++j) {
            if (!pattern.wildcard[j] && data[i + j] != pattern.bytes[j]) {
                matched = false;
                break;
            }
        }
        if (matched) return static_cast<intptr_t>(i);
    }
    return -1;
}

namespace {

bool FillModuleInfo(const wchar_t* wideName, uintptr_t base, size_t size, ModuleInfo& outInfo) {
    char nameBuf[256] = {};
    WideCharToMultiByte(CP_UTF8, 0, wideName, -1, nameBuf, sizeof(nameBuf), nullptr, nullptr);
    outInfo.name = nameBuf;
    outInfo.base = base;
    outInfo.size = size;
    return outInfo.base != 0 && outInfo.size != 0;
}

bool FindModulePsapi(uint32_t pid, const std::wstring& moduleName, ModuleInfo& outInfo) {
    HANDLE process = OpenProcess(PROCESS_QUERY_INFORMATION | PROCESS_VM_READ, FALSE, pid);
    if (!process) return false;

    HMODULE mods[1024] = {};
    DWORD needed = 0;
    bool found = false;
    if (EnumProcessModulesEx(process, mods, sizeof(mods), &needed, LIST_MODULES_ALL)) {
        const size_t count = needed / sizeof(HMODULE);
        wchar_t name[MAX_PATH] = {};
        for (size_t i = 0; i < count && i < 1024; ++i) {
            if (GetModuleBaseNameW(process, mods[i], name, MAX_PATH) == 0) continue;
            if (_wcsicmp(name, moduleName.c_str()) != 0) continue;
            MODULEINFO mi{};
            if (!GetModuleInformation(process, mods[i], &mi, sizeof(mi))) continue;
            found = FillModuleInfo(name, reinterpret_cast<uintptr_t>(mi.lpBaseOfDll),
                                    static_cast<size_t>(mi.SizeOfImage), outInfo);
            break;
        }
    }
    CloseHandle(process);
    return found;
}

} // namespace

bool AobScanner::FindModule(uint32_t pid, const std::wstring& moduleName, ModuleInfo& outInfo) {
    HANDLE snapshot = CreateToolhelp32Snapshot(TH32CS_SNAPMODULE | TH32CS_SNAPMODULE32, pid);
    if (snapshot != INVALID_HANDLE_VALUE) {
        MODULEENTRY32W entry{};
        entry.dwSize = sizeof(entry);
        if (Module32FirstW(snapshot, &entry)) {
            do {
                if (_wcsicmp(entry.szModule, moduleName.c_str()) == 0) {
                    FillModuleInfo(entry.szModule,
                                   reinterpret_cast<uintptr_t>(entry.modBaseAddr),
                                   static_cast<size_t>(entry.modBaseSize), outInfo);
                    CloseHandle(snapshot);
                    return outInfo.base != 0 && outInfo.size != 0;
                }
            } while (Module32NextW(snapshot, &entry));
        }
        CloseHandle(snapshot);
    }
    return FindModulePsapi(pid, moduleName, outInfo);
}

std::vector<uintptr_t> AobScanner::ScanModuleHits(IMemoryReader& reader, const ModuleInfo& module,
                                                  const AobPattern& pattern, size_t chunkSize,
                                                  size_t maxHits) {
    std::vector<uintptr_t> hits;
    if (!reader.IsAttached() || !pattern.Valid() || module.size == 0 || maxHits == 0) {
        return hits;
    }
    if (chunkSize < pattern.Size() * 2) chunkSize = pattern.Size() * 2;

    const size_t overlap = pattern.Size() - 1;
    std::vector<uint8_t> buffer(chunkSize);
    size_t offset = 0;
    while (offset < module.size && hits.size() < maxHits) {
        const size_t toRead = (std::min)(chunkSize, module.size - offset);
        if (!reader.Read(module.base + offset, buffer.data(), toRead)) {
            offset += (std::max)(toRead, static_cast<size_t>(0x1000));
            continue;
        }

        size_t searchFrom = 0;
        while (hits.size() < maxHits && searchFrom < toRead) {
            const intptr_t idx = Find(buffer.data() + searchFrom, toRead - searchFrom, pattern);
            if (idx < 0) break;
            hits.push_back(module.base + offset + searchFrom + static_cast<size_t>(idx));
            searchFrom += static_cast<size_t>(idx) + 1;
        }

        if (toRead <= overlap) break;
        offset += toRead - overlap;
    }
    return hits;
}

uintptr_t AobScanner::ScanModule(IMemoryReader& reader, const ModuleInfo& module,
                                 const AobPattern& pattern, size_t chunkSize) {
    const auto hits = ScanModuleHits(reader, module, pattern, chunkSize, 1);
    return hits.empty() ? 0 : hits[0];
}

uintptr_t AobScanner::ResolveRipRelative(uintptr_t matchAddress, size_t ripOffset) {
    return matchAddress + ripOffset;
}

uintptr_t AobScanner::ResolveRipTarget(IMemoryReader& reader, uintptr_t matchAddress, size_t ripOffset) {
    if (matchAddress == 0) return 0;
    int32_t disp = 0;
    if (!reader.ReadValue<int32_t>(matchAddress + ripOffset, disp)) {
        return 0;
    }
    return matchAddress + ripOffset + 4 + disp;
}
