"""RE Phase 2 - Tim Pointer Base danh sach NPC (tight filters, 2-run validation).
ONLY ReadProcessMemory (VM_READ), never write. Requires Admin.
Usage: python scan_phase2.py [pid]   (no args = scan all vggame)
Output: logs/re_npc_list/phase2_pid<PID>.txt
"""
from __future__ import annotations

import ctypes
import os
import subprocess
import sys
import time
from collections import Counter
from ctypes import wintypes

import numpy as np

LOG_DIR = r"C:\Projects\JX\logs\re_npc_list"
os.makedirs(LOG_DIR, exist_ok=True)

k32 = ctypes.WinDLL("kernel32", use_last_error=True)
PROCESS_VM_READ = 0x0010
PROCESS_QUERY_INFORMATION = 0x0400

k32.OpenProcess.restype = wintypes.HANDLE
k32.OpenProcess.argtypes = [wintypes.DWORD, wintypes.BOOL, wintypes.DWORD]
k32.ReadProcessMemory.restype = wintypes.BOOL
k32.ReadProcessMemory.argtypes = [wintypes.HANDLE, wintypes.LPCVOID, wintypes.LPVOID,
                                  ctypes.c_size_t, ctypes.POINTER(ctypes.c_size_t)]
k32.VirtualQueryEx.restype = ctypes.c_size_t
k32.VirtualQueryEx.argtypes = [wintypes.HANDLE, wintypes.LPCVOID, ctypes.c_void_p, ctypes.c_size_t]
k32.CloseHandle.argtypes = [wintypes.HANDLE]

STATIC_NAME = 0x006DEC44
STATIC_LEVEL = 0x006DEC66
STATIC_MAP = 0x006966CC

MEM_COMMIT = 0x1000
MEM_IMAGE = 0x1000000
MEM_MAPPED = 0x40000
PAGE_GUARD = 0x100
PAGE_NOACCESS = 0x01

CHUNK = 8 * 1024 * 1024
STEP = CHUNK - 0x400
COORD_OFF = 0x270


class MBI(ctypes.Structure):
    _fields_ = [
        ("BaseAddress", ctypes.c_void_p),
        ("AllocationBase", ctypes.c_void_p),
        ("AllocationProtect", wintypes.DWORD),
        ("RegionSize", ctypes.c_size_t),
        ("State", wintypes.DWORD),
        ("Protect", wintypes.DWORD),
        ("Type", wintypes.DWORD),
    ]


def vggame_pids():
    out = subprocess.check_output(["tasklist", "/FI", "IMAGENAME eq vggame.exe", "/FO", "CSV", "/NH"],
                                  text=True, errors="replace")
    pids = []
    for line in out.splitlines():
        parts = [p.strip('"') for p in line.split('","')]
        if len(parts) >= 2 and parts[0].lower() == "vggame.exe":
            try:
                pids.append(int(parts[1]))
            except ValueError:
                pass
    return pids


def read_mem(h, addr, size):
    buf = (ctypes.c_char * size)()
    br = ctypes.c_size_t(0)
    if not k32.ReadProcessMemory(h, ctypes.c_void_p(addr), buf, size, ctypes.byref(br)):
        return None
    return bytes(buf)[: br.value]


def enum_regions(h):
    regions = []
    addr = 0
    mbi = MBI()
    sz = ctypes.sizeof(mbi)
    maxa = 0x7FFF0000
    while addr < maxa:
        if k32.VirtualQueryEx(h, ctypes.c_void_p(addr), ctypes.byref(mbi), sz) == 0:
            break
        base = int(mbi.BaseAddress or 0)
        size = int(mbi.RegionSize)
        if size <= 0:
            size = 0x1000
        readable = (mbi.State == MEM_COMMIT and (mbi.Protect & 0xEE) != 0
                    and (mbi.Protect & PAGE_NOACCESS) == 0 and (mbi.Protect & PAGE_GUARD) == 0)
        if readable and not (mbi.Type == MEM_MAPPED and size > 16 * 1024 * 1024):
            regions.append((base, size, int(mbi.Type)))
        addr = base + size
    return regions


def as_i32(buf):
    n = len(buf) & ~3
    return np.frombuffer(buf[:n], dtype=np.int32)


def as_u32(buf):
    n = len(buf) & ~3
    return np.frombuffer(buf[:n], dtype=np.uint32)

def read_i32(h, addr):
    b = read_mem(h, addr, 4)
    if not b or len(b) < 4:
        return None
    import struct
    return struct.unpack("<i", b)[0]


def read_fields(h, S):
    """Doc khoi S: hp(S-24), s1(S-20), s2(S-4), sta(S), ms(S+4), cx(S+0x270), cy(S+0x274)."""
    raw = read_mem(h, S - 0x28, 0x28 + 0x280)
    if not raw or len(raw) < 0x28 + 0x280:
        return None
    import struct

    def at(rel):
        return struct.unpack("<i", raw[0x28 + rel:0x28 + rel + 4])[0]

    return {"S": S, "hp": at(-24), "s1": at(-20), "v100": at(-16), "s2": at(-4),
            "sta": at(0), "ms": at(4), "x": at(COORD_OFF), "y": at(COORD_OFF + 4)}


def hud(x, y):
    return (round(x / 256.0), round(y / 512.0))


def load_prev_anchor(pid):
    """Lay (S, x, y) cua lan phase1 truoc tu log, de neo/tolerance."""
    import re
    p = os.path.join(LOG_DIR, f"phase1_pid{pid}.txt")
    try:
        txt = open(p, encoding="utf-8", errors="replace").read()
    except OSError:
        return None
    ms = re.findall(r"=> PlayerS=0x([0-9A-Fa-f]+) coords=\((\d+),(\d+)\)", txt)
    if not ms:
        return None
    s, x, y = ms[-1]
    return (int(s, 16), int(x), int(y))


def dump_compact(h, addr, size, log, label):
    raw = read_mem(h, addr, size)
    if not raw:
        log(f"   dump fail @0x{addr:X}")
        return
    log(f"   --- {label} @0x{addr:X} ({len(raw)}B) ---")
    for off in range(0, len(raw), 32):
        chunk = raw[off:off + 32]
        hexs = " ".join(f"{b:02X}" for b in chunk)
        asc = "".join(chr(b) if 32 <= b < 127 else "." for b in chunk)
        log(f"   {addr + off:08X}: {hexs}  {asc}")


def iter_window(h, lo, hi):
    for base, size, _typ in enum_regions(h):
        if base + size <= lo or base >= hi:
            continue
        s0 = max(base, lo)
        s1 = min(base + size, hi)
        off = 0
        total = s1 - s0
        while off < total:
            n = min(CHUNK, total - off)
            data = read_mem(h, s0 + off, n)
            if data:
                yield s0 + off, data
            off += n


def pass_A(h, log, anchor):
    """Strict sig [sig1,100,0,0,sig2,sta,ms] trong cac cua so heap; chon player (uu tien anchor)."""
    ranges = [(0x0F000000, 0x14000000), (0x30000000, 0x38000000), (0x58000000, 0x68000000)]
    found = []
    for (lo, hi) in ranges:
        for base, data in iter_window(h, lo, hi):
            a = as_i32(data)
            n = a.size
            if n < 16:
                continue
            i = np.arange(1, n - 8)
            m = ((a[i] >= 0) & (a[i] < 200) & (a[i + 1] == 100) & (a[i + 2] == 0) &
                 (a[i + 3] == 0) & (a[i + 4] >= 0) & (a[i + 4] < 200) &
                 (a[i - 1] > 0) & (a[i - 1] < 50_000_000) & (a[i + 5] > 0) &
                 (a[i + 5] < 100_000) & (a[i + 6] >= a[i + 5]) & (a[i + 6] < 200_000))
            for k0 in np.nonzero(m)[0]:
                k = int(k0) + 1
                found.append(base + (k + 5) * 4)
        if found:
            break
    valid = []
    for S in found[:60]:
        f = read_fields(h, S)
        if not f:
            continue
        hx, hy = hud(f["x"], f["y"])
        ok = f["x"] >= 16000 and f["y"] >= 16000 and f["sta"] > 0
        log(f"   sig S=0x{S:X} hp={f['hp']} sta={f['sta']}/{f['ms']} coords=({f['x']},{f['y']}) HUD={hx, hy} {'VALID' if ok else ''}")
        if ok:
            valid.append(f)
    player = None
    if anchor and valid:
        for f in valid:
            if f["S"] == anchor[0]:
                dx, dy = abs(f["x"] - anchor[1]), abs(f["y"] - anchor[2])
                tag = "STABLE" if dx <= 60 and dy <= 60 else f"MOVED(dx={dx},dy={dy})"
                log(f"   anchor 0x{anchor[0]:X} trung khop player ({tag})")
                player = f
                break
    if player is None and valid:
        player = valid[0]
        log(f"   chon player S=0x{player['S']:X} (khong co anchor hoac anchor lech)")
    if player is None:
        log("   [A] KHONG tim duoc player hop le")
    return found, valid, player


def pass_B(h, log, player):
    """Quet ung vien object: property + coords hop le, x>=16000, gan player. Tra ve cands sap xep theo khoang cach."""
    cands = []
    for base, size, _typ in enum_regions(h):
        off = 0
        while off < size:
            n = min(CHUNK, size - off)
            m0 = max(0, off - 0x300)
            m1 = min(size, off + n + 0x300)
            data = read_mem(h, base + m0, m1 - m0)
            if data:
                a = as_i32(data)
                cnt = a.size
                off_in = (off - m0) // 4
                j0 = max(6, off_in)
                j1 = min(cnt - 158, off_in + n // 4)
                if j1 > j0:
                    j = np.arange(j0, j1)
                    m = ((a[j - 6] > 0) & (a[j - 6] < 50_000_000) &
                         (a[j - 5] > -1000) & (a[j - 5] < 1000) &
                         (a[j - 1] > -1000) & (a[j - 1] < 1000) &
                         (a[j] >= 0) & (a[j] < 200_000) & (a[j + 1] >= a[j]) & (a[j + 1] < 1_000_000) &
                         (a[j + 156] >= 16000) & (a[j + 156] < 0x800000) &
                         (a[j + 157] >= 16000) & (a[j + 157] < 0x800000))
                    for k0 in np.nonzero(m)[0]:
                        k = int(k0) + j0
                        S = base + m0 + k * 4
                        cands.append((S, int(a[k - 6]), int(a[k - 5]), int(a[k - 1]),
                                      int(a[k]), int(a[k + 1]), int(a[k + 156]), int(a[k + 157])))
            off += n
    cands.sort()
    out = []
    for c in cands:
        if out and c[0] - out[-1][0] < 0x40:
            continue
        out.append(c)
    log(f"  [B] {len(out)} ung vien (dedup tu {len(cands)})")
    px, py = player["x"], player["y"]
    scored = []
    for c in out:
        dxhud = abs(round(c[6] / 256.0) - round(px / 256.0))
        dyhud = abs(round(c[7] / 512.0) - round(py / 512.0))
        snapx = c[6] % 256 == 0
        snapy = c[7] % 512 == 0
        scored.append((dxhud <= 160 and dyhud <= 160 and not (snapx and snapy), c, dxhud, dyhud))
    near = [s for s in scored if s[0]]
    near.sort(key=lambda t: (t[2] ** 2 + t[3] ** 2))
    log(f"  [B] {len(near)} ung vien trong ban kinh 160 HUD (loai snap 256/512)")
    for keep, c, dxhud, dyhud in near[:20]:
        log(f"     S=0x{c[0]:X} hp={c[1]} sig=({c[2]},{c[3]}) sta={c[4]}/{c[5]} "
            f"coords=({c[6]},{c[7]}) HUD=({round(c[6]/256)},{round(c[7]/512)}) d=({dxhud},{dyhud})")
    return [c for _k, c, _d1, _d2 in near]


def merged_intervals(targets, radius):
    ints = sorted((max(0, t - radius), t + radius) for t in targets)
    merged = []
    for s, e in ints:
        if merged and s <= merged[-1][1] + 4:
            merged[-1] = (merged[-1][0], max(merged[-1][1], e))
        else:
            merged.append((s, e))
    return merged


def pointer_hits(h, targets, radius, static_only=False):
    merged = merged_intervals(targets, radius)
    if not merged:
        return []
    starts = np.array([s for s, _ in merged], dtype=np.uint32)
    ends = np.array([e for _, e in merged], dtype=np.uint32)
    hits = []
    for base, size, _typ in enum_regions(h):
        if static_only and base >= 0x10000000:
            continue
        off = 0
        while off < size:
            n = min(CHUNK, size - off)
            data = read_mem(h, base + off, n)
            if data:
                a = as_u32(data)
                pos = np.searchsorted(starts, a, side="right") - 1
                vmask = pos >= 0
                if vmask.any():
                    idx = np.nonzero(vmask)[0]
                    pv = pos[idx]
                    ok = a[idx] <= ends[pv]
                    for k in idx[ok]:
                        k = int(k)
                        hits.append((base + off + k * 4, int(a[k])))
            off += n
    hits.sort()
    return hits


def find_runs(hits, min_len=3):
    runs = []
    cur = []
    for hh in hits:
        if cur and hh[0] - cur[-1][0] == 4:
            cur.append(hh)
        else:
            if len(cur) >= min_len:
                runs.append(cur)
            cur = [hh]
    if len(cur) >= min_len:
        runs.append(cur)
    return runs


def slope_stride_report(addrs, log, label):
    from collections import Counter as _C
    a = sorted(addrs)
    if len(a) < 2:
        log(f"   [stride/{label}] it du lieu")
        return
    diffs = _C(b - c for b, c in zip(a[1:], a))
    log(f"   [stride/{label}] top diffs: {diffs.most_common(8)}")


def verify_stability(h, cands, log, pause_s=6):
    time.sleep(pause_s)
    stable = []
    for c in cands:
        S = c[0]
        raw = read_mem(h, S - 0x28, 0x28 + 0x280)
        if not raw or len(raw) < 0x28 + 0x280:
            continue
        import struct
        hp = struct.unpack("<i", raw[0x28 - 24:0x28 - 20])[0]
        st = struct.unpack("<i", raw[0x28:0x28 + 4])[0]
        cx = struct.unpack("<i", raw[0x28 + COORD_OFF:0x28 + COORD_OFF + 4])[0]
        cy = struct.unpack("<i", raw[0x28 + COORD_OFF + 4:0x28 + COORD_OFF + 8])[0]
        if hp == c[1] and cx == c[6] and cy == c[7] and st == c[4]:
            stable.append(c)
    log(f"  [stab] {len(stable)}/{len(cands)} on dinh")
    return stable


def module_ranges(h):
    mods = []
    for base, size, typ in enum_regions(h):
        if base < 0x10000000 and typ == MEM_IMAGE:
            mods.append((base, size))
    return mods

# ==== P4MAIN2 ====
def main():
    only = [int(x) for x in sys.argv[1:]] or None
    pids = [p for p in vggame_pids() if only is None or p in only]
    print(f"vggame PIDs: {pids}", flush=True)
    for pid in pids:
        log_path = os.path.join(LOG_DIR, f"phase2_pid{pid}.txt")
        lines = []

        def log(msg, _lines=lines, _path=log_path):
            s = f"[{time.strftime('%H:%M:%S')}] {msg}"
            print(s, flush=True)
            _lines.append(s)
            with open(_path, "w", encoding="utf-8") as f:
                f.write("\n".join(_lines))

        log(f"===== PHASE2 PID {pid} =====")
        h = k32.OpenProcess(PROCESS_QUERY_INFORMATION | PROCESS_VM_READ, False, pid)
        if not h:
            log(f"OpenProcess FAILED err={ctypes.get_last_error()}")
            continue
        try:
            nmB = read_mem(h, STATIC_NAME, 32)
            lvB = read_mem(h, STATIC_LEVEL, 4)
            mpB = read_mem(h, STATIC_MAP, 8)
            name = nmB.split(b"\x00", 1)[0].decode("latin1", "replace") if nmB else "?"
            log(f"name='{name}' level={int.from_bytes(lvB,'little') if lvB else -1} "
                f"mapId={int.from_bytes(mpB[:4],'little') if mpB else -1}")
            anchor = load_prev_anchor(pid)
            if anchor:
                log(f"anchor lan truoc: S=0x{anchor[0]:X} coords=({anchor[1]},{anchor[2]})")
            found, valid, player = pass_A(h, log, anchor)
            if player is None:
                continue
            S_p = player["S"]
            log(f"=> Player S=0x{S_p:X} hp={player['hp']} sta={player['sta']}/{player['ms']} "
                f"coords=({player['x']},{player['y']}) HUD={hud(player['x'], player['y'])}")
            dump_compact(h, S_p - 0x40, 0x180, log, f"playerS=0x{S_p:X}")
            near = pass_B(h, log, player)
            if not near:
                log("  khong co ung vien gan -> dung")
                continue
            slope_stride_report([c[0] for c in near], log, "near-cands")
            stable = verify_stability(h, near[:60], log)
            keep = stable if len(stable) >= 4 else near[:60]
            if len(stable) < 4:
                log("  stab qua it -> dung top near (GHI CHU: chua xac minh on dinh)")
            else:
                log(f"  dung {len(keep)} ung vien on dinh")
            tset = sorted(set(c[0] for c in keep) | {S_p})
            hits = [t for t in pointer_hits(h, tset, 0x800) if (t[1] & 3) == 0]
            import bisect
            log(f"  [C] refs (aligned, ±0x800): {len(hits)}")
            runs = find_runs(hits, min_len=3)
            log(f"  [C] runs >=3: {len(runs)}")
            cmin, cmax = min(tset), max(tset)
            scored = []
            for run in runs:
                vals = [v for _, v in run]
                m = sum(1 for v in vals if cmin - 0x800 <= v <= cmax + 0x800)
                scored.append((m, len(run), run[0][0], vals[:12]))
            scored.sort(reverse=True)
            for mm, ln, base_addr, sample in scored[:25]:
                log(f"     RUN base=0x{base_addr:X} len={ln} inrange={mm} "
                    f"vals={[hex(v) for v in sample]}")
            mods = module_ranges(h)
            log(f"  [mods] MEM_IMAGE <0x10000000: {[(hex(b), hex(s)) for b, s in mods][:12]}")
            top_bases = [s[2] for s in scored[:10] if s[0] > 0]
            if top_bases:
                hitsD = pointer_hits(h, top_bases, 0x40, static_only=True)
                log(f"  [D] static refs toi run bases: {len(hitsD)} (PENDING: chua verify)")
                for hit, val in hitsD[:20]:
                    log(f"     STATIC? 0x{hit:X} -> 0x{val:X}")
        except Exception as ex:
            import traceback
            log(f"EXCEPTION: {ex}")
            log(traceback.format_exc())
        finally:
            k32.CloseHandle(h)
            log(f"===== PID {pid} done =====")
    print("ALL DONE", flush=True)


if __name__ == "__main__":
    main()


