"""Live diagnostic: find POE2 camera/fog in PID 11644. User-mode RPM only."""
from __future__ import annotations

import ctypes
import ctypes.wintypes as w
import struct
import sys
from collections import namedtuple

kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
psapi = ctypes.WinDLL("psapi", use_last_error=True)

PROCESS_VM_READ = 0x0010
PROCESS_QUERY_INFORMATION = 0x0400
MEM_COMMIT = 0x1000
PAGE_NOACCESS = 0x01
PAGE_GUARD = 0x100

kernel32.OpenProcess.restype = w.HANDLE
kernel32.ReadProcessMemory.argtypes = [w.HANDLE, ctypes.c_void_p, ctypes.c_void_p, ctypes.c_size_t, ctypes.POINTER(ctypes.c_size_t)]
kernel32.ReadProcessMemory.restype = w.BOOL
kernel32.VirtualQueryEx.argtypes = [w.HANDLE, ctypes.c_void_p, ctypes.c_void_p, ctypes.c_size_t]
kernel32.VirtualQueryEx.restype = ctypes.c_size_t

class MEMORY_BASIC_INFORMATION(ctypes.Structure):
    _fields_ = [
        ("BaseAddress", ctypes.c_void_p),
        ("AllocationBase", ctypes.c_void_p),
        ("AllocationProtect", w.DWORD),
        ("PartitionId", w.WORD),
        ("RegionSize", ctypes.c_size_t),
        ("State", w.DWORD),
        ("Protect", w.DWORD),
        ("Type", w.DWORD),
    ]

class MODULEINFO(ctypes.Structure):
    _fields_ = [
        ("lpBaseOfDll", ctypes.c_void_p),
        ("SizeOfImage", w.DWORD),
        ("EntryPoint", ctypes.c_void_p),
    ]


def open_pid(pid: int):
    h = kernel32.OpenProcess(PROCESS_VM_READ | PROCESS_QUERY_INFORMATION, False, pid)
    if not h:
        raise OSError(ctypes.get_last_error(), "OpenProcess")
    return h


def rpm(h, addr, n):
    buf = (ctypes.c_ubyte * n)()
    got = ctypes.c_size_t(0)
    ok = kernel32.ReadProcessMemory(h, ctypes.c_void_p(addr), buf, n, ctypes.byref(got))
    if not ok:
        return None
    return bytes(buf[: got.value])


def main_module(h):
    mods = (ctypes.c_uint64 * 1024)()
    needed = w.DWORD()
    if not psapi.EnumProcessModulesEx(h, ctypes.byref(mods), ctypes.sizeof(mods), ctypes.byref(needed), 0x03):
        return None
    n = needed.value // 8
    name = ctypes.create_unicode_buffer(260)
    psapi.GetModuleBaseNameW.argtypes = [w.HANDLE, ctypes.c_void_p, w.LPWSTR, w.DWORD]
    psapi.GetModuleInformation.argtypes = [w.HANDLE, ctypes.c_void_p, ctypes.c_void_p, w.DWORD]
    for i in range(n):
        base_h = mods[i]
        psapi.GetModuleBaseNameW(h, base_h, name, 260)
        if name.value.lower() in ("pathofexile.exe", "pathofexilesteam.exe", "pathofexile2.exe"):
            mi = MODULEINFO()
            psapi.GetModuleInformation(h, ctypes.c_void_p(base_h), ctypes.byref(mi := mi if False else MODULEINFO()), ctypes.sizeof(MODULEINFO))
            # retry clean
            mi = MODULEINFO()
            psapi.GetModuleInformation(h, ctypes.c_void_p(base_h), ctypes.byref(mi), ctypes.sizeof(MODULEINFO))
            return name.value, int(mi.lpBaseOfDll or 0), int(mi.SizeOfImage)
    return None


def aob_find(data: bytes, pattern: str):
    parts = pattern.split()
    needle = []
    mask = []
    for p in parts:
        if p == "?":
            needle.append(0)
            mask.append(False)
        else:
            needle.append(int(p, 16))
            mask.append(True)
    n = len(needle)
    first = needle[0]
    first_m = mask[0]
    hits = []
    for i in range(0, len(data) - n + 1):
        if first_m and data[i] != first:
            continue
        ok = True
        for j in range(n):
            if mask[j] and data[i + j] != needle[j]:
                ok = False
                break
        if ok:
            hits.append(i)
            if len(hits) >= 8:
                break
    return hits


def scan_aob_in_module(h, base, size, pattern, chunk=2 * 1024 * 1024):
    hits = []
    off = 0
    plen = len(pattern.split())
    while off < size:
        n = min(chunk, size - off)
        blob = rpm(h, base + off, n)
        if blob:
            for rel in aob_find(blob, pattern):
                hits.append(base + off + rel)
        off += max(n - plen, 1) if blob else n
        if len(hits) >= 6:
            break
    return hits


def rip_target(h, match, disp_off):
    raw = rpm(h, match + disp_off, 4)
    if not raw or len(raw) < 4:
        return 0
    disp = struct.unpack("<i", raw)[0]
    return match + disp_off + 4 + disp


def f32(b, o):
    return struct.unpack_from("<f", b, o)[0]


def plausible_cam(block, i):
    try:
        cur, mn, mx, fov, zn, zf = (f32(block, i + o) for o in (0x2C, 0x30, 0x34, 0x38, 0x3C, 0x40))
    except Exception:
        return None
    if not all(map(lambda x: x == x and abs(x) < 1e8, (cur, mn, mx, fov, zn, zf))):
        return None
    if mn < 1 or mn > 80:
        return None
    if mx < 8 or mx > 400:
        return None
    if not ((0.2 <= fov <= 2.0) or (30 <= fov <= 120)):
        return None
    if zn < 0.001 or zn > 20:
        return None
    if zf < 20 or zf > 20000:
        return None
    return (cur, mn, mx, fov, zn, zf)


def scan_structs(h, limit_regions=400):
    mbi = MEMORY_BASIC_INFORMATION()
    addr = 0x10000
    found = []
    regions = 0
    scanned = 0
    while addr < 0x7FFFFFFFFFFF and regions < 8000 and len(found) < 25:
        q = kernel32.VirtualQueryEx(h, ctypes.c_void_p(addr), ctypes.byref(mbi), ctypes.sizeof(mbi))
        if q == 0:
            break
        base = int(mbi.BaseAddress or 0)
        size = int(mbi.RegionSize)
        prot = mbi.Protect
        readable = (mbi.State == MEM_COMMIT) and not (prot & (PAGE_NOACCESS | PAGE_GUARD)) and prot != 0
        nxt = base + size
        if readable and 0x80 <= size <= (64 << 20):
            regions += 1
            step = min(size, 2 << 20)
            off = 0
            while off + 0x50 < size and len(found) < 25:
                n = min(step, size - off)
                blob = rpm(h, base + off, n)
                if blob:
                    scanned += len(blob)
                    for i in range(0, len(blob) - 0x50, 4):
                        cam = plausible_cam(blob, i)
                        if cam:
                            found.append((base + off + i, cam, mbi.Type, size))
                            if len(found) >= 25:
                                break
                off += n
        if nxt <= addr:
            break
        addr = nxt
    return found, scanned


def main():
    pid = int(sys.argv[1]) if len(sys.argv) > 1 else 11644
    h = open_pid(pid)
    info = main_module(h)
    print("module", info)
    if not info:
        print("NO MAIN MODULE")
        return
    name, base, size = info
    print(f"{name} base={base:#x} size={size:#x}")
    peek = rpm(h, base, 64)
    print("PE peek", None if peek is None else peek[:16].hex())

    patterns = {
        "cam_old_lea": "48 8D 0D ? ? ? ? F3 0F 10 41 2C F3 0F 10 49 34 0F 2F C1",
        "cam_p20_mov": "48 8B 05 ? ? ? ? 48 85 C0 74 ? F3 0F 10 40 34",
        "cam_p20_cmp": "F3 0F 10 41 2C F3 0F 10 49 34 0F 2F C1 76 ? F3 0F 11 49 2C",
        "fog_old": "48 8B 0D ? ? ? ? F3 0F 11 05 ? ? ? ? 48 85 C9 74 ? E8",
        "fog_p20": "F3 0F 11 05 ? ? ? ? E8 ? ? ? ? 48 8D 4D ? E8",
        "movss_40_34": "F3 0F 10 40 34",
        "movss_41_2C": "F3 0F 10 41 2C",
    }
    for key, pat in patterns.items():
        hits = scan_aob_in_module(h, base, min(size, 0x3000000), pat)
        print(f"AOB {key}: {len(hits)} hits")
        for hit in hits[:4]:
            extra = ""
            if "48 8B 05" in pat or "48 8D 0D" in pat or "48 8B 0D" in pat:
                tgt = rip_target(h, hit, 3)
                extra = f" rip={tgt:#x}"
                if tgt:
                    d8 = rpm(h, tgt, 8)
                    extra += f" *[{d8.hex() if d8 else '?'}]"
            if "F3 0F 11 05" in pat:
                tgt = rip_target(h, hit, 4)
                extra += f" store={tgt:#x}"
            print(f"  {hit:#x}{extra}")

    found, scanned = scan_structs(h)
    print(f"struct candidates {len(found)} scanned_bytes={scanned}")
    for addr, cam, typ, rsz in found[:20]:
        cur, mn, mx, fov, zn, zf = cam
        print(f"  {addr:#x} type={typ:#x} rsz={rsz:#x} cur={cur:.2f} min={mn:.2f} max={mx:.2f} fov={fov:.3f} zNear={zn:.3f} zFar={zf:.1f}")


if __name__ == "__main__":
    main()
